Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4d2e13cf2b | ||
|
|
5134e5ecfc | ||
|
|
caadfdec0b | ||
|
|
9af700cc4f | ||
|
|
0c7908b9f2 | ||
|
|
06c169f73f | ||
|
|
29d0113b37 | ||
|
|
6e3020b942 | ||
|
|
832fc3af84 | ||
|
|
bc9e77384e | ||
|
|
3dc526475d | ||
|
|
89709f5f80 | ||
|
|
09a1c3a948 | ||
|
|
403392b41b | ||
|
|
c33fadc266 | ||
|
|
0443ab3a61 | ||
|
|
22a44e67a8 | ||
|
|
24b8619f64 | ||
|
|
3319b6410e | ||
|
|
37d45fdee3 | ||
|
|
55cb98ff56 | ||
|
|
517cd8d102 | ||
|
|
7ea7680f56 | ||
|
|
2c804b0ac4 | ||
|
|
589b62b529 | ||
|
|
21ac7e95a3 | ||
|
|
fb27716186 | ||
|
|
37a9da50df | ||
|
|
db9977926c | ||
|
|
c0c6c2181a | ||
|
|
ae5d23f226 | ||
|
|
c584a4270c | ||
|
|
91aea7fe8c | ||
|
|
b4073f6378 | ||
|
|
bb6b2db88b | ||
|
|
248315de14 | ||
|
|
75db531c12 | ||
|
|
c89fd237b8 | ||
|
|
6fd8c599c1 | ||
|
|
877221c118 | ||
|
|
2856def6c0 | ||
|
|
d6cda4a04b | ||
|
|
fe3300bd65 | ||
|
|
783205a965 | ||
|
|
dc1bc41d2e | ||
|
|
655afbe90b | ||
|
|
6f8221df58 | ||
|
|
ff5cec43bd | ||
|
|
10558173fb | ||
|
|
754787f43d | ||
|
|
27c97bfe96 | ||
|
|
c6ec1a3484 | ||
|
|
40b655e99e | ||
|
|
b696c5deff | ||
|
|
7572283517 | ||
|
|
61cee42ded | ||
|
|
815446d5bb | ||
|
|
a146e17bdc | ||
|
|
2c4e1fce8f | ||
|
|
81e245548d | ||
|
|
4ed45ce843 | ||
|
|
2d3035a112 | ||
|
|
39837e0a3a | ||
|
|
8c7428122b | ||
|
|
51246bcb31 | ||
|
|
c7e634776d | ||
|
|
a5c9459401 | ||
|
|
5055fb85aa | ||
|
|
7ed7e81e84 | ||
|
|
303c426c3f | ||
|
|
f9c3ccd869 | ||
|
|
eb53281c9a | ||
|
|
a285a390c1 | ||
|
|
75df948f34 | ||
|
|
260f3c3a22 | ||
|
|
b0487dd6dd | ||
|
|
70e4ffcc65 | ||
|
|
0883638027 | ||
|
|
ee5de69e37 | ||
|
|
0cc331d1c6 | ||
|
|
41f256321b | ||
|
|
44b9463498 | ||
|
|
e6d35fc4cc | ||
|
|
396d9ac181 | ||
|
|
67a7b23b85 | ||
|
|
edf2c6c8f7 | ||
|
|
0eba3df119 | ||
|
|
aa851d93c6 | ||
|
|
fa76764c3b | ||
|
|
83ec36cd38 | ||
|
|
cdd7b88bec | ||
|
|
8927c9bb3d | ||
|
|
422a4768ea | ||
|
|
c65b29ec0f | ||
|
|
0be069c165 | ||
|
|
5ffd4e53c3 | ||
|
|
4fa3a74827 | ||
|
|
416baef813 | ||
|
|
45fea34bd0 | ||
|
|
953432b5fe | ||
|
|
e5b5e5917b | ||
|
|
67c9de8efd | ||
|
|
a3b487422d | ||
|
|
f53ec857c0 | ||
|
|
b617d56c60 | ||
|
|
718b226177 | ||
|
|
f8ec63203c | ||
|
|
fd7a59d37a | ||
|
|
ea6d02da0f | ||
|
|
ab84bbf08c | ||
|
|
b58b0ea7ca | ||
|
|
d3676b4f71 | ||
|
|
c4688b958d | ||
|
|
33b91bd8ae | ||
|
|
4c05abbe59 | ||
|
|
0fc630b34b | ||
|
|
ee11069ef2 | ||
|
|
b05be8a907 | ||
|
|
a66477b710 | ||
|
|
7b1aa749eb | ||
|
|
958237473f | ||
|
|
a70a6589af | ||
|
|
4bc4630721 | ||
|
|
a8e5f0a54d | ||
|
|
8bc4ac2641 | ||
|
|
58f9170319 | ||
|
|
e98730b20d | ||
|
|
9802b0d135 | ||
|
|
ed2d7d4acd | ||
|
|
f85e906dec | ||
|
|
70b89d01c2 | ||
|
|
3fd0384ffc | ||
|
|
78d276b4ff | ||
|
|
5796d44363 | ||
|
|
6e14e446cb | ||
|
|
c93d4f04aa | ||
|
|
33cd199e6d | ||
|
|
4712544d5e | ||
|
|
958fdbdc88 | ||
|
|
390e200f76 | ||
|
|
462b66b807 | ||
|
|
388f62f8a0 | ||
|
|
7be009649a | ||
|
|
534206095f | ||
|
|
525c115a3a | ||
|
|
b34d6c836e | ||
|
|
6fdf9b4340 | ||
|
|
18e6a10778 | ||
|
|
36d08fa2a7 | ||
|
|
2bdd2ab94e | ||
|
|
2414dfca70 | ||
|
|
6ea591491e | ||
|
|
d4d9786434 | ||
|
|
ac3449cac9 | ||
|
|
71d6212ab8 | ||
|
|
2308b59f13 | ||
|
|
61a2672215 | ||
|
|
386ac95814 | ||
|
|
914039ac81 | ||
|
|
ff25ccca65 | ||
|
|
01198eaeef | ||
|
|
15d96b1f2a | ||
|
|
342539f1e1 | ||
|
|
c31694af09 | ||
|
|
964a098a4b | ||
|
|
c3394288bb | ||
|
|
8a016931f1 | ||
|
|
e69ce6e1c6 | ||
|
|
6050a94d77 | ||
|
|
5fd26b7549 | ||
|
|
2a9a023172 | ||
|
|
b295a20b9d | ||
|
|
f923edcaaa | ||
|
|
c8260745a6 | ||
|
|
3c67774eb3 | ||
|
|
ce4a323f43 | ||
|
|
b7934e9182 | ||
|
|
46c1d6591b | ||
|
|
3730a9eaac | ||
|
|
a4f7ec1fb3 | ||
|
|
677e164f29 | ||
|
|
da0fd0da0d | ||
|
|
1da3b7f7e8 | ||
|
|
8c57cfa645 | ||
|
|
00924fbf79 | ||
|
|
9df25b6932 | ||
|
|
452954ff1e | ||
|
|
ec8e20af35 | ||
|
|
89629b8f03 | ||
|
|
7240517807 | ||
|
|
0502494e9a | ||
|
|
1f6336fd98 | ||
|
|
368b4a5b22 | ||
|
|
b308391527 | ||
|
|
b854eb09b1 | ||
|
|
62b153749a | ||
|
|
7292cee868 | ||
|
|
bc70696f4f | ||
|
|
dbdcfd8c60 | ||
|
|
7e13fd7ad1 | ||
|
|
124c7a3283 | ||
|
|
cfb49c4c18 | ||
|
|
2560533c1a | ||
|
|
5b1c42e81a | ||
|
|
6f5f263244 | ||
|
|
5922727402 | ||
|
|
03a8363583 | ||
|
|
97901220f2 | ||
|
|
8977a10a2b | ||
|
|
cd1ec31957 | ||
|
|
ef8c9c063c | ||
|
|
dd4f43bfdb | ||
|
|
3a232f5e9a | ||
|
|
23d03d6aae | ||
|
|
d99ac7d3f8 | ||
|
|
c7be66626f | ||
|
|
464e703e47 | ||
|
|
518702caae | ||
|
|
62ae206918 | ||
|
|
516051304e | ||
|
|
0130b49514 | ||
|
|
5aeb1ca708 | ||
|
|
df634bb64f | ||
|
|
6729e64f30 | ||
|
|
ea3f5f22d2 | ||
|
|
0b0910bee2 | ||
|
|
7c3802a55e | ||
|
|
08f64f7908 | ||
|
|
e3ba698453 | ||
|
|
741b64edb6 | ||
|
|
7c0b0e42f5 | ||
|
|
8f890f0b43 | ||
|
|
ede39d82de | ||
|
|
1a8e1a9939 | ||
|
|
7b55a63fc7 | ||
|
|
7d9b249671 | ||
|
|
e124c2656a | ||
|
|
5576e6ed8a | ||
|
|
7453968678 | ||
|
|
47a1bfdd15 | ||
|
|
b5c43968db | ||
|
|
35f8bf97e3 | ||
|
|
1457f2dec8 | ||
|
|
1111a3a222 | ||
|
|
f812072215 | ||
|
|
8934bfb04b | ||
|
|
fd56086e79 | ||
|
|
95391221df | ||
|
|
19db873603 | ||
|
|
7f08376f0c | ||
|
|
15c7e37438 | ||
|
|
b1c2536ed2 | ||
|
|
91762ed807 | ||
|
|
a0c2ec3d2c | ||
|
|
7e8153e889 | ||
|
|
223f484ded | ||
|
|
88901bfa04 | ||
|
|
928eb015bd | ||
|
|
a54878b14f | ||
|
|
8b9e28b503 | ||
|
|
3f0c0e0a0d | ||
|
|
8958b64b5a | ||
|
|
21f9e5295b | ||
|
|
0ffc04797f | ||
|
|
b2809e6293 | ||
|
|
beb9bf60e4 | ||
|
|
2f9b28a57d | ||
|
|
d501e3d6b5 | ||
|
|
819ad1d904 | ||
|
|
9fe3a00dba | ||
|
|
4584adf900 | ||
|
|
dfdb76cc46 | ||
|
|
17df026492 | ||
|
|
e038bab66d | ||
|
|
c39be0e2d6 | ||
|
|
6cd0ba0b6b | ||
|
|
e473ab1231 | ||
|
|
4dbb2f94a6 | ||
|
|
dee07d8a30 | ||
|
|
e031fadc35 | ||
|
|
b4aef82401 | ||
|
|
5cdcdbaeec | ||
|
|
e8d55c0a8b | ||
|
|
57a5e43696 | ||
|
|
7b29834d42 | ||
|
|
2b9b956dad | ||
|
|
25090dbf17 | ||
|
|
bb2663aad1 | ||
|
|
3571db34e2 | ||
|
|
7d1f941580 | ||
|
|
ed4cb358a0 | ||
|
|
caedcbae49 | ||
|
|
9ccda6715c | ||
|
|
edf3ae9209 | ||
|
|
0726db7217 | ||
|
|
9c6c375dfe | ||
|
|
3e3c5b6d78 | ||
|
|
2bc91e8f52 | ||
|
|
e5ed45fb20 | ||
|
|
232421f40b | ||
|
|
6b2d962cd6 | ||
|
|
7ee75a0c04 | ||
|
|
fc9c2ea191 | ||
|
|
6f2e97aa58 | ||
|
|
5f3a628a8d | ||
|
|
8b35ce924b | ||
|
|
966b3fdb57 | ||
|
|
ee3a49a88d | ||
|
|
22f2fe1ffb | ||
|
|
993e749121 | ||
|
|
4c06b392da | ||
|
|
fbcdcf146b | ||
|
|
ec86ce5cf7 | ||
|
|
087878ce84 | ||
|
|
4f69c33de0 | ||
|
|
c205d3a353 | ||
|
|
4210cae68e | ||
|
|
de8ea08f5c | ||
|
|
d07e4154fe | ||
|
|
bb1419328b | ||
|
|
40c09167cd | ||
|
|
56ae99e96a | ||
|
|
1eecbc1ac2 | ||
|
|
78a5015846 | ||
|
|
b93d560788 | ||
|
|
b7626f05fb | ||
|
|
1fc0e3ade7 | ||
|
|
84e4110538 | ||
|
|
19a176fd36 | ||
|
|
bb66d435b7 | ||
|
|
dc4924b66e | ||
|
|
920b655f46 | ||
|
|
05098d25a5 | ||
|
|
3266a8c9eb | ||
|
|
ecdb6f353a | ||
|
|
084d040e22 | ||
|
|
76854d1424 | ||
|
|
45fcf272ef | ||
|
|
c783fd30f2 | ||
|
|
d65ac445a4 | ||
|
|
38920c0ed1 | ||
|
|
5019af79a0 | ||
|
|
f85cb27ef8 | ||
|
|
4602abe5c6 | ||
|
|
b1d40f3409 | ||
|
|
02dc3e689c | ||
|
|
b868da6bcd | ||
|
|
90f4b4fcda | ||
|
+36 |
b3e9d16b6f | ||
|
|
3660bc00fd | ||
|
|
f51d2b026f | ||
|
+11 |
adc9076d17 | ||
|
+2 |
8dae237a0b | ||
|
|
0a8a620fb6 | ||
|
|
f31768e20e | ||
|
|
9bd84258d0 | ||
|
|
4d058a125b | ||
|
|
e4e69a10ec | ||
|
|
6c159a97b7 | ||
|
|
947dcd34bd | ||
|
|
79f0437980 | ||
|
|
6137f7cb7e | ||
|
|
9c9a18d6d4 | ||
|
|
1ac3dd4a89 | ||
|
|
2ed3055c42 | ||
|
|
b8112d72b9 | ||
|
|
4225791313 | ||
|
|
7c7fe44328 | ||
|
|
883f1dda0f | ||
|
|
7a7a25766c | ||
|
|
2b26355002 | ||
|
|
f2a360cb87 | ||
|
|
f9b0534e0c | ||
|
|
6adde203cd | ||
|
|
a7271532f8 | ||
|
|
d95f533214 | ||
|
|
6f1486ffd0 | ||
|
|
140605e660 | ||
|
|
9899293f05 | ||
|
|
e3faec62c5 | ||
|
|
fc05e0a6c5 | ||
|
|
fe6783c166 |
@@ -1 +1,5 @@
|
||||
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.
|
||||
|
||||
@@ -7,10 +7,10 @@ name: Python CI
|
||||
on:
|
||||
push:
|
||||
branches: [main, dev]
|
||||
paths: ['backend/**', 'pyproject.toml', 'uv.lock']
|
||||
paths: ['backend/**', 'pyproject.toml', 'uv.lock', '.github/workflows/backend.yaml']
|
||||
pull_request:
|
||||
branches: [main, dev]
|
||||
paths: ['backend/**', 'pyproject.toml', 'uv.lock']
|
||||
paths: ['backend/**', 'pyproject.toml', 'uv.lock', '.github/workflows/backend.yaml']
|
||||
|
||||
concurrency:
|
||||
group: backend-${{ github.ref }}
|
||||
@@ -38,3 +38,6 @@ jobs:
|
||||
|
||||
- name: Verify formatting
|
||||
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 .
|
||||
|
||||
@@ -231,6 +231,70 @@ jobs:
|
||||
run: |
|
||||
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:
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ !cancelled() && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')) }}
|
||||
|
||||
@@ -310,3 +310,4 @@ dist
|
||||
cypress/videos
|
||||
cypress/screenshots
|
||||
.vscode/settings.json
|
||||
.cptr
|
||||
|
||||
+216
@@ -5,6 +5,222 @@ All notable changes to this project will be documented in this file.
|
||||
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
|
||||
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
|
||||
|
||||
## [0.10.0] - 2026-06-29
|
||||
|
||||
### Added
|
||||
|
||||
- 🤝 **Share folders with your team.** You can now share a folder and the chats inside it with specific users, groups, or everyone, with read or write access; people you share with see shared folders in their sidebar and open the chats in a read-only view when they are not the owner, and administrators control who is allowed to share folders with a new "Folders Sharing" permission that is off by default. [Commit](https://github.com/open-webui/open-webui/commit/5019af79a0c45743ede8c9ff37d68f768e7f6174), [Commit](https://github.com/open-webui/open-webui/commit/38920c0ed1f6ad5fe3bb9d12898fa968ead3634a), [Commit](https://github.com/open-webui/open-webui/commit/d65ac445a43348c5f0323d54c37397ae7f483cb8), [Commit](https://github.com/open-webui/open-webui/commit/c783fd30f20d6be5028cf337bc5e5c2f9afbd3f8), [Commit](https://github.com/open-webui/open-webui/commit/45fcf272ef51c84cb01c1454589da3f98e4adc2c), [Commit](https://github.com/open-webui/open-webui/commit/76854d14246660af8222a5302513020e1f36c4f3), [Commit](https://github.com/open-webui/open-webui/commit/084d040e220ee39f62757d928d839646e813fb25), [Commit](https://github.com/open-webui/open-webui/commit/10558173fb155c63403aa8d80f16f8a3ccfa72a6)
|
||||
- 🗜️ **Automatic context compaction for long chats.** Conversations that grow past a configurable token threshold can now be summarized automatically so they stay within a model's context window, with a notification shown while it happens; administrators can enable it, set the threshold, customize the summarization prompt, and lower the threshold per model. It is off by default. [Commit](https://github.com/open-webui/open-webui/commit/3f0c0e0a0ddff841b015f96f9649c6999a435c73), [Commit](https://github.com/open-webui/open-webui/commit/7f08376f0c06e1a7fba983a2fa93deb8dfbe7cb0), [Commit](https://github.com/open-webui/open-webui/commit/8934bfb04bf366aece872028944e280c25e36d3e), [#19594](https://github.com/open-webui/open-webui/issues/19594)
|
||||
- 🖥️ **Open WebUI Computer agent support.** Open WebUI can now connect to Open WebUI Computer through its OpenAI-compatible gateway, letting chats run full agent sessions on your own machine with file, terminal, git, and web access. [GitHub](https://github.com/open-webui/computer)
|
||||
- 🚀 **Much faster hybrid search on large knowledge bases.** Hybrid search now runs natively in the database on pgvector setups instead of loading an entire collection into memory, so querying large knowledge bases is dramatically faster. [Commit](https://github.com/open-webui/open-webui/commit/223f484ded01d092979693341dc03351a9fa17fa), [#20737](https://github.com/open-webui/open-webui/discussions/20737)
|
||||
- 🗂️ **External knowledge bases.** Knowledge bases can now be backed by an external retrieval source through configurable external knowledge connections, so you can search an existing external system from chat instead of only Open WebUI's built-in store. [Commit](https://github.com/open-webui/open-webui/commit/15c7e374384488effc3d6059d09b3a8aa79c618d)
|
||||
- 🧠 **Reworked memory system.** Memory has been overhauled with distinct memory types — long-lived personal memories and per-conversation context — managed through a structured add, update, and delete flow, giving models a more reliable way to remember and apply what they've learned about you. [Commit](https://github.com/open-webui/open-webui/commit/dbdcfd8c6080c284024052a482e589de226dbf05), [Commit](https://github.com/open-webui/open-webui/commit/7e13fd7ad19c28ba34502bde6c709665f7a6808c), [Commit](https://github.com/open-webui/open-webui/commit/2560533c1a8a2b2a0f7b47becf3703031053ad30), [Commit](https://github.com/open-webui/open-webui/commit/8977a10a2b1e150393635dc5c24e59660d0f2da9), [Commit](https://github.com/open-webui/open-webui/commit/260f3c3a22c55f15ca8f06c5314b23f2a9eb1739), [Commit](https://github.com/open-webui/open-webui/commit/b0487dd6dd942a757828a25aba92e9feef685275), [Commit](https://github.com/open-webui/open-webui/commit/a285a390c12e27e614d1ba9ffb92d1a0e49e7dfa), [Commit](https://github.com/open-webui/open-webui/commit/70e4ffcc6526c1bc90dbfdc0287283574b12c18b), [Commit](https://github.com/open-webui/open-webui/commit/2c4e1fce8f40b0cb5028f1afcb184b6e58c33041), [Commit](https://github.com/open-webui/open-webui/commit/c7e634776d7e77556d149b6cc884ed64363648a9)
|
||||
- 🧩 **New plugin primitive: the Event function.** Where pipe, filter, and action functions all run inside a conversation, the new Event function is the first primitive that hooks into the system itself: it runs your own Python in response to events emitted across the whole application — sign-ups, configuration changes, file uploads, role changes, deletions, startup and shutdown, and more. That makes a new class of behavior possible directly inside Open WebUI, from onboarding and access control to auditing, lifecycle automation, and external integrations. Comes with starter boilerplate in the function editor. [Commit](https://github.com/open-webui/open-webui/commit/e124c2656a4c2092b070e570e35b2f0fb7f584de), [Docs](https://docs.openwebui.com/features/extensibility/plugin/functions/event)
|
||||
- 🔔 **New event system with webhooks.** Open WebUI now emits events for a wide range of system activity — sign-ins, configuration changes, startup, and actions across chats, knowledge, files, and more. Administrators can send these as outbound webhooks, route them to specific users or groups, and manage which events go where from a new event settings admin page. [Commit](https://github.com/open-webui/open-webui/commit/b5c43968db0ea1556b228d143ae5946dc4e944ba), [Commit](https://github.com/open-webui/open-webui/commit/745396867888718289a2dfcf0809b3e162e00629), [Commit](https://github.com/open-webui/open-webui/commit/5576e6ed8a80a4032b7c6cb3ee0cda0254019355), [Commit](https://github.com/open-webui/open-webui/commit/7b55a63fc7ee323e9114713ce1d2f3f688aa37e6), [Commit](https://github.com/open-webui/open-webui/commit/1a8e1a993928a28b9d814c77b2aef0361b630f27), [Commit](https://github.com/open-webui/open-webui/commit/ede39d82de05eeb7591329679c8cab97753e5ff0), [Commit](https://github.com/open-webui/open-webui/commit/8f890f0b43aed3d42b9e3d954e25e54e37d526d0), [Commit](https://github.com/open-webui/open-webui/commit/741b64edb6c2ff04c4b787528cfca1c5b66a1a27), [Commit](https://github.com/open-webui/open-webui/commit/303c426c3fffafc2369205021a82659ee8715a85), [#1240](https://github.com/open-webui/open-webui/issues/1240), [#16426](https://github.com/open-webui/open-webui/pull/16426)
|
||||
- 🔐 **Configure authentication from the admin panel.** LDAP and OAuth/OIDC settings now have a dedicated Authentication settings page, so providers can be configured from the admin interface. [Commit](https://github.com/open-webui/open-webui/commit/5cdcdbaeec9fc8156721c38c33ec37956962871c), [#12945](https://github.com/open-webui/open-webui/pull/12945)
|
||||
- 🏷️ **More custom header variables.** Custom request headers now support "{{USER_MESSAGE_ID}}", "{{USER_MESSAGE_PARENT_ID}}", and "{{TASK}}", letting connected services tell apart real user messages from automated background requests like title, tag, and follow-up generation. [Commit](https://github.com/open-webui/open-webui/commit/f85cb27ef835aa76aff7de6176bf2159ba392061)
|
||||
- 📄 **File details forwarded to external document extractors.** External custom document-extraction servers now receive the file's ID, name, and content type, and these are also available as custom header variables, so extraction can be tailored per file. [Commit](https://github.com/open-webui/open-webui/commit/b1c2536ed2f8639efade04618018e6de9b332df2), [#26259](https://github.com/open-webui/open-webui/issues/26259)
|
||||
- 🎰 **Last model pre-selected for new slots.** When you add another model to a multi-model chat, the slot now defaults to the model you last picked instead of starting empty. [#25974](https://github.com/open-webui/open-webui/pull/25974)
|
||||
- ⚡ **Faster model overview.** The admin model overview now loads its feedback history and tags through batched queries, so it opens noticeably faster on instances with many chats. [Commit](https://github.com/open-webui/open-webui/commit/40c09167cd6de1c853a5dd03c88b4fdcb279dfe1)
|
||||
- 🏎️ **Lighter channel profile previews.** Profile previews in channels now load a person's details only when you hover to open one, rather than fetching them for every message up front. [Commit](https://github.com/open-webui/open-webui/commit/4f69c33de0e9a8fde4f16d0b2f1ed8aac8741772)
|
||||
- ↩️ **Reset permissions to defaults.** The group and default permission dialogs now include a button to restore all permissions back to their built-in defaults in one step. [#25931](https://github.com/open-webui/open-webui/pull/25931)
|
||||
- 📥 **Chat import permission.** Administrators can now control whether users are allowed to import or clone chats, with a new "Allow Chat Import" permission. [Commit](https://github.com/open-webui/open-webui/commit/edf3ae920989b01383be543e7379d3eade03c0b6), [Commit](https://github.com/open-webui/open-webui/commit/9ccda6715c3b2dc2cbc1302d2396b2f0233bdea8), [Commit](https://github.com/open-webui/open-webui/commit/ed4cb358a06fc6962b378f20ef846fd2b0af90bc), [#25927](https://github.com/open-webui/open-webui/pull/25927)
|
||||
- 🔔 **Per-group user webhook permission.** Administrators can now control which users may set a personal notification webhook, with a new "User Webhooks" permission. [#25923](https://github.com/open-webui/open-webui/pull/25923)
|
||||
- ✍️ **Customizable autocomplete prompt.** Administrators can now set a custom prompt template for autocomplete generation from the admin interface settings. [Commit](https://github.com/open-webui/open-webui/commit/4dbb2f94a66d6e0035e2da857ddb6a841a68f862), [#25879](https://github.com/open-webui/open-webui/pull/25879)
|
||||
- 🔑 **Configurable secret key length.** The auto-generated secret key length can now be set with a new environment variable, instead of always using a fixed length. [Commit](https://github.com/open-webui/open-webui/commit/e473ab1231abedcb188c259b42ae7f2390739223), [#25906](https://github.com/open-webui/open-webui/pull/25906)
|
||||
- 🏟️ **Arena evaluation models configurable via environment.** Arena evaluation models can now be defined through an environment variable, which previously could not be set that way. [Commit](https://github.com/open-webui/open-webui/commit/fd56086e793a0eceb07a55ff972a5492d8f8a285)
|
||||
- ✏️ **Edit prompts from the menu.** The prompts list now has an Edit option in each prompt's menu, taking you straight to its editor. [#25789](https://github.com/open-webui/open-webui/pull/25789)
|
||||
- 📋 **Clone automations.** Automations now have a Clone option in their menu, so you can duplicate one as a starting point. [#25790](https://github.com/open-webui/open-webui/pull/25790)
|
||||
- 🔁 **Recurring calendar events.** The calendar event editor now includes a repeat option, so events can recur on a schedule. [#25865](https://github.com/open-webui/open-webui/pull/25865)
|
||||
- 🧷 **Separate skills import and export permissions.** Administrators can now control importing and exporting skills independently, with new skills import and export permissions. [#25921](https://github.com/open-webui/open-webui/pull/25921)
|
||||
- 🏷️ **Filter admin models by tag.** The admin Models settings page now has a tag filter for narrowing the model list by base-model tags. [Commit](https://github.com/open-webui/open-webui/commit/2bdd2ab94eefd3d75dd6511c445e302f82221b5d)
|
||||
- 📊 **Sortable analytics chat list.** The model chat list in analytics now has sortable column headers, so you can order it by title, last updated, or user. [Commit](https://github.com/open-webui/open-webui/commit/3730a9eaac68dff60b3ae5b4ed160b91480d66eb), [#26168](https://github.com/open-webui/open-webui/pull/26168)
|
||||
- 🔐 **Argon2 password hashing option.** Password hashing can now use Argon2 through a configurable algorithm setting, removing the 72-byte password length limit that came with the previous default. [Commit](https://github.com/open-webui/open-webui/commit/33cd199e6dffddd4ee8974af41ebb894871d74c1), [Commit](https://github.com/open-webui/open-webui/commit/a70a6589afad0b429cdd77afa62163391f406a87), [#25656](https://github.com/open-webui/open-webui/pull/25656)
|
||||
- 🔐 **Optional encryption of valve values at rest.** Tool and function valve values can now be encrypted at rest through a new opt-in setting, with existing stored values migrated automatically, so sensitive settings like API keys aren't kept in plaintext. [Commit](https://github.com/open-webui/open-webui/commit/b4073f6378392b23a3954e33031bf5e1d98e090a), [#23721](https://github.com/open-webui/open-webui/pull/23721)
|
||||
- 🗄️ **AWS RDS IAM database authentication.** The database connection can now authenticate using AWS RDS IAM tokens through a new opt-in setting, instead of only a static password. [Commit](https://github.com/open-webui/open-webui/commit/c0c6c2181a8dc57b62e8a5eabd550bf89db7ffed), [#23580](https://github.com/open-webui/open-webui/pull/23580)
|
||||
- 🔓 **Automatic auth for models with OAuth 2.1 tools.** When a model uses tools that require OAuth 2.1, Open WebUI now initiates the authorization flow automatically instead of failing the request. [Commit](https://github.com/open-webui/open-webui/commit/ae5d23f2267845922c2acb507a4a432908d03b41), [#23325](https://github.com/open-webui/open-webui/pull/23325), [#23272](https://github.com/open-webui/open-webui/issues/23272)
|
||||
- 🔤 **Custom tokenizer for token-based text splitting.** Token-based document splitting can now use a configurable Hugging Face tokenizer model, so chunking can match the tokenizer of the model you use. [Commit](https://github.com/open-webui/open-webui/commit/bb6b2db88b1e82395531f67db4f6accd49d8b9eb), [#24139](https://github.com/open-webui/open-webui/pull/24139)
|
||||
- 🔒 **Restrict OAuth scopes requested from MCP servers.** A new setting lets administrators limit which OAuth scopes Open WebUI requests when connecting to MCP servers. [Commit](https://github.com/open-webui/open-webui/commit/7be009649a0a94008484335c492ab0af18fde41f), [#25981](https://github.com/open-webui/open-webui/pull/25981), [#25978](https://github.com/open-webui/open-webui/issues/25978)
|
||||
- 🧩 **Filter Outlet Hook can now run on API requests and responses.** A filter function's outlet hook now runs for direct API callers, including streaming responses, so response post-processing isn't limited to the web interface; this is controlled by a new setting and on by default. [Commit](https://github.com/open-webui/open-webui/commit/390e200f76877b185002c88b4dda27b123d29e83), [#25650](https://github.com/open-webui/open-webui/pull/25650)
|
||||
- 🖥️ **Setting for terminal sidebar auto-open.** A new interface setting controls whether the files sidebar opens automatically when you select a terminal. [Commit](https://github.com/open-webui/open-webui/commit/958237473f8cbde97eb0df8c21f2a4de088c4459), [#25628](https://github.com/open-webui/open-webui/pull/25628)
|
||||
- 📌 **Reorder pinned notes by dragging.** Pinned notes in the sidebar can now be dragged to reorder them. [#25677](https://github.com/open-webui/open-webui/pull/25677)
|
||||
- 🔎 **Chat actions in search.** The search dialog now offers a context menu on each result, so you can act on a chat directly from search. [#25490](https://github.com/open-webui/open-webui/pull/25490)
|
||||
- 🔎 **Snippets in chat search results.** Searching your chats now shows a snippet of the matching content in each result, so you can tell results apart at a glance. [Commit](https://github.com/open-webui/open-webui/commit/0eba3df1199f56e8ac77213772a41313a3237296), [Commit](https://github.com/open-webui/open-webui/commit/67a7b23b85d2e3ce6b094b682ba9f07dc453d355), [Commit](https://github.com/open-webui/open-webui/commit/8927c9bb3d4b04f4fc8e443f42089f2a551276bc), [#25178](https://github.com/open-webui/open-webui/pull/25178)
|
||||
- 📝 **Formatted valve descriptions.** Valve descriptions for tools and functions now render Markdown, so they can include formatting and links. [Commit](https://github.com/open-webui/open-webui/commit/7c0b0e42f5afb0e9c39bd2d42e4822d64d0c7b3e)
|
||||
- 🔽 **Dropdown inputs for valve options.** Valve and confirmation inputs can now present a set of options as a dropdown instead of free text, making fixed-choice settings easier to configure. [Commit](https://github.com/open-webui/open-webui/commit/422a4768ea7428b5dd6d401ccfba00e1e86eb98a), [#26278](https://github.com/open-webui/open-webui/pull/26278)
|
||||
- 🔌 **Control the OAuth resource parameter for MCP connectors.** MCP connectors can now be set to always send, never send, or automatically decide whether to include the OAuth resource parameter, so they work with providers that reject it. [Commit](https://github.com/open-webui/open-webui/commit/5576e6ed8a80a4032b7c6cb3ee0cda0254019355)
|
||||
- 🔎 **SERPHouse web search.** SERPHouse can now be used as a web search provider. [Commit](https://github.com/open-webui/open-webui/commit/3a232f5e9a4d31a6b74cb34d007b581feeb2f005), [Commit](https://github.com/open-webui/open-webui/commit/dd4f43bfdb793c4276e6468b65fc32d9b881578a), [#26254](https://github.com/open-webui/open-webui/pull/26254)
|
||||
- 🔎 **Microsoft Web IQ web search.** Microsoft Web IQ can now be used as a web search provider, with a matching page-browse loader. [#26178](https://github.com/open-webui/open-webui/pull/26178)
|
||||
- ⚠️ **Optional web search confirmation.** Administrators can now require users to confirm before a web search runs, with a banner and message making it clear when search is about to be used. [Commit](https://github.com/open-webui/open-webui/commit/fa76764c3b7f99c5adacd34dabd51ead09542c13), [#24942](https://github.com/open-webui/open-webui/pull/24942)
|
||||
- 🪪 **Client User-Agent forwarded to model backends.** The browser's User-Agent is now passed through to all model backends, so upstream services can see the originating client. [#26333](https://github.com/open-webui/open-webui/pull/26333)
|
||||
- 🖐️ **Drag items from the sidebar into chat.** Folders, notes, and models — including pinned notes — can now be dragged from the sidebar into the chat input. [#25771](https://github.com/open-webui/open-webui/pull/25771), [Commit](https://github.com/open-webui/open-webui/commit/dc1bc41d2e), [#26384](https://github.com/open-webui/open-webui/pull/26384)
|
||||
- 🏷️ **Tag suggestions in the model editor.** The model editor now suggests existing tags as you type, making it easier to reuse a consistent set. [Commit](https://github.com/open-webui/open-webui/commit/b58b0ea7ca849b89d217e1077498a8a3fc92471f), [#25703](https://github.com/open-webui/open-webui/pull/25703)
|
||||
- 🗣️ **Voice suggestions in the model editor.** The model editor now offers a dropdown of available text-to-speech voices, making it easier to pick one. [Commit](https://github.com/open-webui/open-webui/commit/a5c945940134b957dbad47790b5baafeecdac6c4), [#25706](https://github.com/open-webui/open-webui/pull/25706)
|
||||
- 🎛️ **Unified model picker for workspace base model.** Choosing a base model in the model editor now uses the searchable model selector instead of a plain field, making it easier to find and pick the right model. [Commit](https://github.com/open-webui/open-webui/commit/c89fd237b822877bffbb33a37402622983c7189d), [#24576](https://github.com/open-webui/open-webui/issues/24576)
|
||||
- 🔍 **Searchable pickers in the model editor.** Attaching actions, filters, tools, knowledge, and skills to a model now uses type-to-search pickers instead of long checkbox lists, making large libraries easier to manage. [Commit](https://github.com/open-webui/open-webui/commit/61cee42ded4e84e31d7cc9b168994ef165efa8ca)
|
||||
- 🖼️ **iPhone images work with OpenAI image editing.** Uploaded images are now normalized before being sent to OpenAI image editing, fixing edits that failed for certain iPhone photo formats, with a new admin toggle to control the behavior. [Commit](https://github.com/open-webui/open-webui/commit/39837e0a3afd17b7ff617d97dddf4c5d6446e42d), [Commit](https://github.com/open-webui/open-webui/commit/2d3035a1122123df471ac9d0a591a9465a8212e2), [#26252](https://github.com/open-webui/open-webui/pull/26252), [#26249](https://github.com/open-webui/open-webui/issues/26249)
|
||||
- 🟢 **Loaded-model indicator for llama.cpp.** Models served through llama.cpp now report whether they're currently loaded in memory, including the sleeping state, so the loaded indicator works for them too. [Commit](https://github.com/open-webui/open-webui/commit/b696c5deff15d4c85c84c5bac244f062b1bc879a)
|
||||
- 🧱 **Structured model output rendered on the client.** Reasoning, tool calls, and server-side tool steps such as web and file search are now rendered in the browser from the model's structured output instead of being flattened into the message text on the server, giving more accurate and editable rendering of these items. [Commit](https://github.com/open-webui/open-webui/commit/0443ab3a61492799f1aaa449f89cbd8aa5912f57), [Commit](https://github.com/open-webui/open-webui/commit/c33fadc26671190c94d86485e6e2ef2f6fd486a3)
|
||||
- 📜 **Custom CA bundle for outbound connections.** A new environment variable lets you point Open WebUI at a custom CA certificate bundle, and the per-connection SSL settings now accept a bundle path, so deployments behind a corporate or internal CA can keep certificate verification on instead of disabling it. [Commit](https://github.com/open-webui/open-webui/commit/a54878b14f044d4aa1d8cf5be6f8ce9fc4285438), [Commit](https://github.com/open-webui/open-webui/commit/8b9e28b50354307a314111262f6737b6d8aa4685)
|
||||
- 🖥️ **More terminal server orchestrator controls.** Admins connecting an orchestrator terminal server can now configure session lifecycle policies and refresh or reset running terminal sessions, including targeting only idle ones, from the connection settings. [Commit](https://github.com/open-webui/open-webui/commit/7e8153e889a59afe4cf77261ea5e1ef5a66665f1)
|
||||
- 📁 **Terminal file browser can stay within a root folder.** The terminal file navigator now anchors to a defined root and home directory, so users can be kept within their workspace instead of browsing into system folders by accident. [Commit](https://github.com/open-webui/open-webui/commit/a0c2ec3d2cf8d696ede479211330eec2da360d39)
|
||||
- 🧠 **Memory toggle follows the server default.** When a user hasn't set their own memory preference, it now follows the admin's global memory setting instead of defaulting to off. [#25909](https://github.com/open-webui/open-webui/pull/25909)
|
||||
- 🧹 **Unshare all shared chats at once.** The Shared Chats dialog now has a button to stop sharing every shared chat in one action. [#25848](https://github.com/open-webui/open-webui/pull/25848)
|
||||
- 📈 **Richer analytics with a date picker.** The analytics dashboard now lets you choose a date range and shows additional columns. [#25922](https://github.com/open-webui/open-webui/pull/25922), [#25919](https://github.com/open-webui/open-webui/issues/25919)
|
||||
- 🔢 **Chat and file counts in their dialogs.** The Chats and Files dialogs now show the total number of chats and files in their titles. [#25872](https://github.com/open-webui/open-webui/pull/25872), [#25873](https://github.com/open-webui/open-webui/pull/25873)
|
||||
- ⚡ **Faster math rendering.** Rendered math is now cached and reused, so messages with repeated or unchanged math expressions render more efficiently. [#25847](https://github.com/open-webui/open-webui/pull/25847)
|
||||
- ⚡ **Lighter Markdown setup.** Markdown extension setup now runs once instead of on every render, avoiding repeated work and extension stacking. [#25837](https://github.com/open-webui/open-webui/pull/25837)
|
||||
- ⚡ **Snappier read-only code blocks.** Read-only code blocks now skip language auto-detection, so they render faster. [#25824](https://github.com/open-webui/open-webui/pull/25824)
|
||||
- ⚡ **Non-blocking audio model loading.** Loading speech models no longer blocks the server, keeping it responsive while they initialize. [#25806](https://github.com/open-webui/open-webui/pull/25806)
|
||||
- ⚡ **Faster URL safety checks.** The safety check on fetched URLs now resolves addresses off the main loop, so it no longer blocks other work. [#25825](https://github.com/open-webui/open-webui/pull/25825)
|
||||
- ⚡ **Fewer queries for channel reactions and replies.** Channel reactions and thread replies now load through batched queries, reducing database load on busy channels. [#25831](https://github.com/open-webui/open-webui/pull/25831)
|
||||
- ⚡ **Lighter streaming.** Streaming responses now skip re-processing message content that hasn't changed, reducing work on every update. [#26325](https://github.com/open-webui/open-webui/pull/26325), [#26326](https://github.com/open-webui/open-webui/pull/26326)
|
||||
- ⚡ **Smoother tool-call rendering.** Displaying tool calls now parses their content iteratively, avoiding slowdowns on deeply nested data. [#26146](https://github.com/open-webui/open-webui/pull/26146)
|
||||
- ⚡ **Hidden tool-call details cost nothing.** When tool-call arguments are collapsed, they are no longer rendered behind the scenes, noticeably speeding up chats with heavy tool use. [Commit](https://github.com/open-webui/open-webui/commit/b7934e918223ec0a9e972accd647a8654e503156), [#26147](https://github.com/open-webui/open-webui/pull/26147)
|
||||
- ⚡ **Leaner knowledge-file reading for agents.** The built-in tools that let a model read knowledge files now return output in bounded, paginated chunks with a default and a hard cap, instead of potentially returning an entire large file at once, sharply reducing token usage. [Commit](https://github.com/open-webui/open-webui/commit/a285a390c12e27e614d1ba9ffb92d1a0e49e7dfa), [#26139](https://github.com/open-webui/open-webui/issues/26139)
|
||||
- ⚡ **Lighter, faster file search on large knowledge bases.** Listing and searching files no longer returns each file's full extracted text by default, and content matching is now length-bounded, so these requests are far lighter and searching across very large knowledge bases is dramatically faster. [Commit](https://github.com/open-webui/open-webui/commit/36d08fa2a7), [Commit](https://github.com/open-webui/open-webui/commit/46c1d6591badb6ab567ba1b8fae23475d5da105a), [Commit](https://github.com/open-webui/open-webui/commit/ab84bbf08c5935f1a19044ef581986f83311da8b), [#25774](https://github.com/open-webui/open-webui/pull/25774), [#25741](https://github.com/open-webui/open-webui/issues/25741), [#26145](https://github.com/open-webui/open-webui/pull/26145), [#25867](https://github.com/open-webui/open-webui/issues/25867)
|
||||
- ⚡ **Faster password hashing and bulk user import.** Password hashing and verification no longer block the server, and importing users from a CSV is now processed in a single batch, keeping large imports and sign-ins responsive. [Commit](https://github.com/open-webui/open-webui/commit/6fdf9b4340), [#25804](https://github.com/open-webui/open-webui/pull/25804), [#25805](https://github.com/open-webui/open-webui/pull/25805)
|
||||
- ⚡ **Non-blocking model downloads.** Downloading large Ollama models no longer blocks the server on file reads and checksums, keeping it responsive during big downloads. [#25829](https://github.com/open-webui/open-webui/pull/25829)
|
||||
- ⚡ **Non-blocking uploads and link fetches.** Hashing uploaded files and fetching URLs now run off the main loop, so large uploads and link previews don't hold up other requests. [#25822](https://github.com/open-webui/open-webui/pull/25822)
|
||||
- ⚡ **More blocking work moved off the main loop.** Additional blocking operations in audio, pipelines, and plugin handling now run in worker threads, keeping the server responsive under load. [#26381](https://github.com/open-webui/open-webui/pull/26381)
|
||||
- ⚡ **Unreachable backends don't stall model loading.** Loading models and tool servers no longer blocks on backends that are down or slow to respond, so the model list stays responsive when one connection is unreachable. [#26289](https://github.com/open-webui/open-webui/pull/26289)
|
||||
- ⚡ **Batched streaming updates.** Streaming responses now group small updates of the same type before sending them, reducing overhead during fast token streams and tool-call output. [Commit](https://github.com/open-webui/open-webui/commit/7240517807a8b0097065f7cbbb384d34084f90fd), [#26202](https://github.com/open-webui/open-webui/pull/26202)
|
||||
- 🔄 **General improvements.** Various improvements were implemented across the application to enhance performance, stability, and security.
|
||||
- 🌐 **Updated translations.** Catalan, Brazilian Portuguese (pt-BR), Irish, German (de-DE), and Spanish (es-ES) translations were updated.
|
||||
|
||||
### Fixed
|
||||
|
||||
- 🛡️ **Security Advisory**: This release includes security and access-control fixes. We recommend updating production deployments at your earliest convenience. Not all security fixes in this version may be enumerated in the fixed section — some may be withheld for a short time to give administrators time to upgrade. [Advisories](https://github.com/open-webui/open-webui/security)
|
||||
- 🔐 **Knowledge base write access enforced on upload.** Attaching an uploaded file to a knowledge base now requires the same write access as the rest of the knowledge API, so users without write access can no longer add files to a collection by referencing its ID. [#26001](https://github.com/open-webui/open-webui/pull/26001)
|
||||
- 🗝️ **API key permission enforced on all key endpoints.** Viewing and deleting API keys now respects the API keys permission, matching the protection already applied to key creation. [#25992](https://github.com/open-webui/open-webui/pull/25992)
|
||||
- 🔊 **Text-to-speech permission enforced on the speech endpoint.** The OpenAI speech proxy now honors the text-to-speech permission, so it can no longer be used by people who are not allowed to use that feature. [#25993](https://github.com/open-webui/open-webui/pull/25993)
|
||||
- 🎲 **Model access enforced on arena fallback.** Reaching a model indirectly through an arena model on background and task requests now enforces that model's access rules, closing a path that could otherwise bypass them. [#26046](https://github.com/open-webui/open-webui/pull/26046)
|
||||
- ⏰ **Scheduled automations stop for deactivated accounts.** Scheduled automations now re-check the owner's account status and permissions before each run, so they stop when an account is deactivated or has automations access revoked. [#26047](https://github.com/open-webui/open-webui/pull/26047)
|
||||
- 🚧 **Heavily encoded paths rejected behind the proxy.** Request paths that remain encoded after repeated decoding are now rejected instead of forwarded, preventing a path traversal that could otherwise slip through. [#26050](https://github.com/open-webui/open-webui/pull/26050)
|
||||
- 🌐 **Image URL fetches hardened against DNS rebinding.** Fetching user-supplied image URLs now re-checks the destination address at connection time, closing a path that could be used to reach internal addresses behind a public hostname. [#25960](https://github.com/open-webui/open-webui/pull/25960)
|
||||
- 🛂 **Web fetch blocklist matches on hostname.** The web fetch filter now matches entries against the request's hostname on domain boundaries, so blocked hosts can no longer slip through with an added path and lookalike domains are no longer mistaken for allowed ones. [#25949](https://github.com/open-webui/open-webui/pull/25949)
|
||||
- 🪪 **MCP connectors request least-privilege scopes.** MCP connectors that register dynamically over OAuth now request only the scopes for the specific resource rather than the authorization server's full catalog. [#25958](https://github.com/open-webui/open-webui/pull/25958)
|
||||
- 🙈 **Channel member lists no longer expose private data.** Viewing a channel's members now returns only basic profile details, instead of also exposing other members' settings, linked-account data, and personal information. [Commit](https://github.com/open-webui/open-webui/commit/fbcdcf146b99b5002705060a8243eee769108f9e)
|
||||
- 🛟 **SCIM sync can't demote an admin.** A SCIM provisioning sync that marks a user inactive can no longer strip an existing administrator's role, preventing an instance from being locked out of its own administration. [#25948](https://github.com/open-webui/open-webui/pull/25948)
|
||||
- 👻 **Collaborative notes reject unauthenticated presence events.** The remaining real-time note-collaboration events now require an authenticated session, so presence and cursors can no longer be spoofed by someone who only knows a note's ID. [#25946](https://github.com/open-webui/open-webui/pull/25946)
|
||||
- ⏱️ **Login timing no longer reveals which accounts exist.** Sign-in now takes the same amount of time whether or not an account exists, removing a timing difference that could be used to discover valid accounts. [Commit](https://github.com/open-webui/open-webui/commit/993e74912199c66c522f08ec81abe31d76985e39), [Commit](https://github.com/open-webui/open-webui/commit/7b29834d4216e5db70b68f3598fa1ad654d3512b)
|
||||
- 🔌 **Terminal connections can't be redirected to another user.** Terminal session identifiers are now safely encoded before being passed upstream, closing a way to tamper with the connection's user identity. [#26042](https://github.com/open-webui/open-webui/pull/26042)
|
||||
- 📡 **Real-time events only reach your own session.** The server now verifies that a real-time event is delivered only to the requesting user's own active session, instead of trusting a client-supplied session identifier. [#25763](https://github.com/open-webui/open-webui/pull/25763)
|
||||
- 🔓 **Revoked sessions are rejected on real-time connections.** Real-time and terminal WebSocket connections now honor token revocation and expiry, so a signed-out or expired session can no longer keep a live connection open. [Commit](https://github.com/open-webui/open-webui/commit/33b91bd8ae8a100a5a306c91441a7d0b422c4cde), [#25764](https://github.com/open-webui/open-webui/pull/25764), [#25686](https://github.com/open-webui/open-webui/pull/25686)
|
||||
- 🕳️ **Another DNS-rebinding gap closed in URL fetching.** Fetching a URL's content now re-checks the destination address at connection time, closing another path that could reach internal addresses behind a public hostname. [#25775](https://github.com/open-webui/open-webui/pull/25775)
|
||||
- 🗣️ **Azure speech input is escaped.** Voice and language values are now escaped when building Azure text-to-speech requests, preventing malformed or injected markup. [#25776](https://github.com/open-webui/open-webui/pull/25776)
|
||||
- ⚙️ **Interface settings update respects its permission.** Saving interface settings now enforces the interface permission, so users without it can no longer change those settings through the API. [#25996](https://github.com/open-webui/open-webui/pull/25996)
|
||||
- 🗄️ **Unknown knowledge collections are denied by default.** Retrieval now rejects unknown or unscoped collection names by default, closing a legacy path that could be used to reach collections outside the normal access checks. [Commit](https://github.com/open-webui/open-webui/commit/d99ac7d3f83b25161ca775229150c8f7c74cceee)
|
||||
- 🙈 **Error responses no longer leak internals.** Server error responses now return sanitized messages instead of raw exception text, so internal details aren't exposed to signed-in users. [Commit](https://github.com/open-webui/open-webui/commit/ee5de69e374aabf5631da18a5bbc1c285ee6f7a1), [Commit](https://github.com/open-webui/open-webui/commit/0cc331d1c60341bb06b78ceeecfb2db86179c93e), [Commit](https://github.com/open-webui/open-webui/commit/396d9ac18193d43e40fb9d068075d4b780e971d7), [Commit](https://github.com/open-webui/open-webui/commit/0883638027a9b3cb7c9851f031c4f5fc1af1f25d), [#26375](https://github.com/open-webui/open-webui/pull/26375), [#26374](https://github.com/open-webui/open-webui/issues/26374)
|
||||
- 📏 **Upload size limit enforced on the server.** The maximum upload size is now enforced server-side, so it can't be bypassed by a client that ignores the limit. [Commit](https://github.com/open-webui/open-webui/commit/f8ec63203c4408c46bb06698ae624d17b01b9301), [Commit](https://github.com/open-webui/open-webui/commit/d3676b4f71bfdbaf4e4d76943c51117e18932ccf), [#25869](https://github.com/open-webui/open-webui/pull/25869)
|
||||
- 🖼️ **OAuth profile pictures are validated.** Profile picture URLs from OAuth providers are now validated and their type checked when stored, preventing unsafe image sources. [Commit](https://github.com/open-webui/open-webui/commit/eb53281c9acb3660e09554a8dbde0a0b42646b70), [#24548](https://github.com/open-webui/open-webui/pull/24548)
|
||||
- 📦 **Security updates to frontend dependencies.** Several frontend dependencies were updated to patch known security vulnerabilities. [#26281](https://github.com/open-webui/open-webui/pull/26281)
|
||||
- 🤝 **Chat sharing respects the user-sharing permission.** The share-chat dialog now hides the option to share with specific users from people who lack that permission, matching the access rules enforced elsewhere. [#25915](https://github.com/open-webui/open-webui/pull/25915)
|
||||
- 📤 **Chat export respects its permission everywhere.** Every chat export menu now checks the export permission, so users without it can no longer export chats through one of the dropdown menus. [#25914](https://github.com/open-webui/open-webui/pull/25914)
|
||||
- 📂 **File write access requires real ownership.** Editing or deleting a file through a knowledge base or workspace model now requires that the object's owner actually owns the file, so a read-only file can no longer gain write access by being referenced from an object you control. [#26032](https://github.com/open-webui/open-webui/pull/26032)
|
||||
- 🖌️ **Image edit endpoint enforces permission.** The image-edit endpoint now checks the image-edit switch and the image-generation permission, matching image generation, so it can't be called by users who lack access. [#26009](https://github.com/open-webui/open-webui/pull/26009)
|
||||
- 📁 **Folder permission enforced on all folder actions.** Every folder operation now checks the folders permission, so the setting is respected consistently instead of only when listing folders. [Commit](https://github.com/open-webui/open-webui/commit/19a176fd36bea15c49d7f2d1539b4832e57a8bc2)
|
||||
- 🧩 **Code Execution settings collapse when off.** The Code Execution settings section now collapses when the toggle is disabled, keeping the settings page tidy. [#25970](https://github.com/open-webui/open-webui/pull/25970)
|
||||
- 📅 **German date format in Notes.** Dates in the Notes view now display correctly for German, where they previously failed to render. [#25985](https://github.com/open-webui/open-webui/pull/25985)
|
||||
- 🎙️ **ElevenLabs speech keeps working when voices can't load.** Text-to-speech through ElevenLabs no longer fails when the available-voice list can't be fetched, instead of rejecting every voice. [Commit](https://github.com/open-webui/open-webui/commit/bb1419328b11b801b4c939dfc112700ba6f6fdab), [#26075](https://github.com/open-webui/open-webui/issues/26075)
|
||||
- 🪟 **Default Permissions modal resets on close.** Closing the Default Permissions dialog without saving now discards unsaved edits instead of keeping them around the next time you open it. [Commit](https://github.com/open-webui/open-webui/commit/78a5015846a9e55ff2bc9d6cc98f880437abe8ed)
|
||||
- 👯 **Side-by-side chat with the same model.** Running two panes with the same model no longer leaves one pane stuck waiting or showing the other pane's reply after a reload, since each pane's messages are now tracked separately. [Commit](https://github.com/open-webui/open-webui/commit/56ae99e96a845289b5787d2dd26a3d828f2295e7), [#25982](https://github.com/open-webui/open-webui/issues/25982)
|
||||
- 💾 **Model edits no longer lost when changing access.** Adjusting a model's access no longer auto-saves on its own and discards your other unsaved changes to that model. [#26004](https://github.com/open-webui/open-webui/pull/26004)
|
||||
- 🔧 **Parallel tool calls over the Anthropic-compatible API.** External Anthropic-compatible clients calling Open WebUI's messages endpoint now receive tool calls reliably when a model issues several at once or returns them in its final message. [Commit](https://github.com/open-webui/open-webui/commit/4210cae68e30173d7902582d32128dd699d5628a), [#25963](https://github.com/open-webui/open-webui/pull/25963), [#25964](https://github.com/open-webui/open-webui/discussions/25964)
|
||||
- 🗃️ **Prompt caching preserved over the Anthropic-compatible API.** Requests through the Anthropic-compatible API now keep their prompt-caching markers instead of having them stripped, so clients that rely on caching work as intended. [Commit](https://github.com/open-webui/open-webui/commit/caedcbae4988ef59ea7052b2a3198e2da4b5291a), [#25998](https://github.com/open-webui/open-webui/pull/25998), [#25964](https://github.com/open-webui/open-webui/discussions/25964)
|
||||
- 🔁 **Fewer redundant data loads.** Several views no longer fire duplicate background fetches at once, avoiding occasional glitches from overlapping requests. [#25943](https://github.com/open-webui/open-webui/pull/25943), [#25942](https://github.com/open-webui/open-webui/pull/25942), [#25934](https://github.com/open-webui/open-webui/pull/25934), [#25935](https://github.com/open-webui/open-webui/pull/25935), [#25838](https://github.com/open-webui/open-webui/pull/25838), [Commit](https://github.com/open-webui/open-webui/commit/e8d55c0a8beac9de0b2a0fe90f0bc0f9b64c1c1f)
|
||||
- 🔎 **Steadier search boxes across admin and workspace.** Search fields for users, knowledge, prompts, tools, and similar lists now run only as you type and reset to the first page correctly, instead of occasionally re-searching on their own. [Commit](https://github.com/open-webui/open-webui/commit/fc9c2ea1915accd1f6edca467e965283dff71cd7), [#25938](https://github.com/open-webui/open-webui/pull/25938)
|
||||
- 📊 **Admin feedback list loads again on PostgreSQL.** The admin feedback list no longer fails to load on PostgreSQL setups, where it previously returned a server error. [Commit](https://github.com/open-webui/open-webui/commit/7ee75a0c04a31528954903e88c9213d5fbb31aa7), [#25953](https://github.com/open-webui/open-webui/issues/25953)
|
||||
- 🗂️ **Deleting nested folders checks chats correctly.** Deleting a folder that contains subfolders now accounts for the chats inside those subfolders when applying the delete-permission check, instead of only the top-level folder's chats. [Commit](https://github.com/open-webui/open-webui/commit/232421f40b84590e6d6fdecab4e43274aac37add), [#25920](https://github.com/open-webui/open-webui/issues/25920)
|
||||
- 🖱️ **Dragging chats into folders is more reliable.** Dragging a chat into a folder no longer throws an error in cases where the chat couldn't be resolved. [#25928](https://github.com/open-webui/open-webui/pull/25928)
|
||||
- 🛠️ **Workspace menu shows for the skills permission.** Users who only have the skills permission now see the Workspace entry in their menu, which previously appeared only for other workspace permissions. [#25925](https://github.com/open-webui/open-webui/pull/25925)
|
||||
- 🧠 **Admins can always reach memories.** Administrators can now use the memories endpoints regardless of the memories permission toggle, matching how admin access works for other features. [#25924](https://github.com/open-webui/open-webui/pull/25924)
|
||||
- 🖼️ **Image settings page survives a config load failure.** The admin image settings page no longer crashes when its configuration fails to load, showing the page instead. [#25933](https://github.com/open-webui/open-webui/pull/25933)
|
||||
- 🧵 **Code blocks render in channel threads.** Code blocks now display correctly in a channel's thread view, where duplicated message identifiers previously broke their rendering. [Commit](https://github.com/open-webui/open-webui/commit/7d1f9415807a47e0da4f862327e9802a3b839753), [#25917](https://github.com/open-webui/open-webui/pull/25917)
|
||||
- 🔵 **No more false unread badges on chats.** Chats no longer show an unread indicator after automatic changes like title generation or pinning, archiving, and moving them between folders, and newly created chats are marked read correctly so they don't appear unread after a refresh. [#25912](https://github.com/open-webui/open-webui/pull/25912), [#25782](https://github.com/open-webui/open-webui/pull/25782), [#25108](https://github.com/open-webui/open-webui/issues/25108)
|
||||
- 📌 **Pinned notes stay in sync.** Pinning, unpinning, or deleting a note now updates the sidebar's pinned list consistently, instead of showing a stale pin state. [#25918](https://github.com/open-webui/open-webui/pull/25918), [#25640](https://github.com/open-webui/open-webui/pull/25640)
|
||||
- 📅 **All-day calendar events keep their date.** Saving an all-day calendar event no longer shifts it by a day for users in certain time zones. [#25864](https://github.com/open-webui/open-webui/pull/25864)
|
||||
- 🧷 **Damaged chat history recovers more reliably.** When a chat's current position is missing or points at a malformed message, Open WebUI now repairs it from the latest valid message — on both the client and the server — instead of risking a broken history view. [Commit](https://github.com/open-webui/open-webui/commit/2308b59f135e4c2da11eabdf2306e55a5dd4e9fb), [Commit](https://github.com/open-webui/open-webui/commit/a146e17bdcaf94fee3a98aa36b4f401e1f06c1d4), [#26298](https://github.com/open-webui/open-webui/pull/26298), [#26258](https://github.com/open-webui/open-webui/pull/26258), [#26257](https://github.com/open-webui/open-webui/issues/26257)
|
||||
- 💾 **Saving a chat no longer drops messages.** Chat updates are now merged with the existing history on the server, with explicit tracking of deleted messages, instead of overwriting it, preventing message loss from concurrent or partial saves. [Commit](https://github.com/open-webui/open-webui/commit/22a44e67a8ba781feb8f2a267fed0c40213d8432), [Commit](https://github.com/open-webui/open-webui/commit/3319b6410e1b600b7a885a5fb78573e9a2061c22), [Commit](https://github.com/open-webui/open-webui/commit/24b8619f64731788ac38813768abbb64405effa4), [#25657](https://github.com/open-webui/open-webui/pull/25657)
|
||||
- 📺 **Channel message updates stay in their channel.** Streaming updates to a channel message are now skipped if the message no longer exists or belongs to a different channel, preventing stray updates. [Commit](https://github.com/open-webui/open-webui/commit/ac3449cac91e62b08a7c28e54fcd044d14dea791)
|
||||
- 📌 **Pinned channel messages update for everyone.** Pinning or unpinning a channel message now updates live for all members and works from thread views, instead of only changing for the person who pinned it. [Commit](https://github.com/open-webui/open-webui/commit/7ea7680f563da30b121258e5a7d7123185c4da2a)
|
||||
- 📄 **Mistral OCR uploads work again.** Document OCR through Mistral has been repaired after an upstream library change broke its file uploads. [#25779](https://github.com/open-webui/open-webui/pull/25779)
|
||||
- 🗂️ **Chroma collection detection fixed.** Open WebUI now correctly detects existing Chroma collections, fixing a case where it always reported them as missing. [#25780](https://github.com/open-webui/open-webui/pull/25780)
|
||||
- 📊 **Vega-Lite charts render reliably.** Vega-Lite charts in chat are now detected by their code block language tag, so they render correctly. [#25843](https://github.com/open-webui/open-webui/pull/25843)
|
||||
- 🏷️ **Long chat tag lists scroll.** The tags section in the chat menu now scrolls instead of overflowing when a chat has many tags. [#26031](https://github.com/open-webui/open-webui/pull/26031)
|
||||
- ⌨️ **Enter key shows correctly on iOS.** The Enter key symbol in the keyboard shortcuts list no longer renders as an emoji on iOS. [#26173](https://github.com/open-webui/open-webui/pull/26173)
|
||||
- 🔗 **Whitespace in names no longer breaks MCP connections.** User name and info headers are now trimmed before being forwarded, fixing MCP connection failures when a display name contained leading or trailing whitespace. [#26182](https://github.com/open-webui/open-webui/pull/26182), [#26181](https://github.com/open-webui/open-webui/issues/26181)
|
||||
- 🈳 **Search no longer fires mid-composition.** Typing in search with an input method editor (such as Japanese, Chinese, or Korean) no longer triggers a search when you press Enter to confirm a composition. [#26238](https://github.com/open-webui/open-webui/pull/26238), [#26285](https://github.com/open-webui/open-webui/pull/26285), [#26172](https://github.com/open-webui/open-webui/issues/26172)
|
||||
- 🧰 **Valves icon stays visible.** The icon for configuring valves no longer disappears, so user-configurable tool and function settings remain reachable. [#26256](https://github.com/open-webui/open-webui/pull/26256)
|
||||
- 🎛️ **Chat controls persist across navigation.** Edits to chat controls are now kept when navigating between chats, and reverting a control to the chat's saved value persists correctly, instead of being lost. [#26336](https://github.com/open-webui/open-webui/pull/26336), [#25793](https://github.com/open-webui/open-webui/pull/25793)
|
||||
- 🔍 **Chat search tool handles empty queries.** The built-in chat search tool no longer crashes when called with an empty query. [Commit](https://github.com/open-webui/open-webui/commit/b854eb09b13216f914ce5fd07ab717b8f752882b), [#26310](https://github.com/open-webui/open-webui/issues/26310)
|
||||
- 📑 **More robust MinerU document processing.** Document processing through MinerU now handles its ZIP results more safely, including very large outputs. [Commit](https://github.com/open-webui/open-webui/commit/23d03d6aaebcced6c1e39e98dfff76ab73df8804), [#26263](https://github.com/open-webui/open-webui/pull/26263)
|
||||
- ⏰ **Scheduled automations with session-auth tools work.** Automations that use session-authenticated tools or terminals now authenticate correctly when running on a schedule, instead of failing. [Commit](https://github.com/open-webui/open-webui/commit/5b1c42e81a3ef3ad5ce5852dbf84020cb5e2498c), [#26247](https://github.com/open-webui/open-webui/pull/26247), [#26137](https://github.com/open-webui/open-webui/issues/26137)
|
||||
- 📝 **Model system prompt preserved with knowledge.** A model's system prompt is no longer dropped when knowledge retrieval runs with native tool calling. [Commit](https://github.com/open-webui/open-webui/commit/cfb49c4c181a96d5df07fbfacd819639baef0bab), [#26217](https://github.com/open-webui/open-webui/pull/26217)
|
||||
- 🔑 **Expired sessions return you to sign-in.** When a request fails because your session has expired, Open WebUI now redirects you to the sign-in page instead of leaving you on a broken view. [Commit](https://github.com/open-webui/open-webui/commit/5922727402593900758d84004f950071c701f6de), [#26237](https://github.com/open-webui/open-webui/pull/26237)
|
||||
- 🎯 **Ejecting a workspace model unloads the right model.** Unloading a workspace model now resolves to its underlying base model, so the correct model is freed from memory. [Commit](https://github.com/open-webui/open-webui/commit/464e703e4716812d015966582152ddd8a2c71572), [#26269](https://github.com/open-webui/open-webui/pull/26269)
|
||||
- 🔄 **Edited models refresh in the admin list.** After editing a model in the admin settings, the models list now updates right away instead of needing a manual reload. [Commit](https://github.com/open-webui/open-webui/commit/b34d6c836ee43d0e9721fa4fd6457d934e3e2a17)
|
||||
- 🗂️ **Workspace model bulk actions and search work across pages.** Bulk actions on workspace models now apply across all of them, and search results paginate correctly. [#26274](https://github.com/open-webui/open-webui/pull/26274)
|
||||
- 🧩 **MCP resource results come through.** Tool results that return resource content — including binary blobs and URI references — are no longer silently dropped, and image results are attached as files. [#25260](https://github.com/open-webui/open-webui/pull/25260), [#24038](https://github.com/open-webui/open-webui/issues/24038), [Commit](https://github.com/open-webui/open-webui/commit/783205a965c556815fae84b64d74f26a2e5e5729)
|
||||
- 🔗 **Broader MCP server compatibility for OAuth.** Open WebUI now discovers an MCP server's protected resource metadata even when the server doesn't advertise it, and recognizes more OAuth preflight variations, so more MCP servers connect. [#25980](https://github.com/open-webui/open-webui/pull/25980), [#25954](https://github.com/open-webui/open-webui/issues/25954), [Commit](https://github.com/open-webui/open-webui/commit/45fea34bd0c8ce54b0822499c40e3e6964220354), [#26068](https://github.com/open-webui/open-webui/pull/26068)
|
||||
- 📤 **Clearer upload error messages.** Failed uploads now show a readable explanation instead of an opaque error stub. [#25961](https://github.com/open-webui/open-webui/pull/25961)
|
||||
- 📋 **Cloned prompts get a proper title.** Cloning a prompt now adds the clone suffix to the correct field, so the duplicate is named as expected. [#25800](https://github.com/open-webui/open-webui/pull/25800)
|
||||
- 📐 **Long default group names don't overflow.** A long default group name no longer overflows its row in the admin authentication settings. [#25685](https://github.com/open-webui/open-webui/pull/25685)
|
||||
- 🖐️ **Sidebar drags don't trigger uploads.** Dragging a chat item in the sidebar no longer shows the file-upload overlay. [#25675](https://github.com/open-webui/open-webui/pull/25675)
|
||||
- 🔁 **Recovers from a stuck streaming response.** If the signal that a response finished is missed — for example after a mobile app is backgrounded mid-stream — Open WebUI now recovers the chat instead of leaving it stuck in a streaming state. [Commit](https://github.com/open-webui/open-webui/commit/aa851d93c63e7da6e94292d0b7586674339d47e5), [Commit](https://github.com/open-webui/open-webui/commit/edf2c6c8f76e7f6a5917e991f371a604adc34c5f), [Commit](https://github.com/open-webui/open-webui/commit/2856def6c05b2fb8c55b4e7170f05db0c4f956f1), [#26320](https://github.com/open-webui/open-webui/pull/26320), [#26315](https://github.com/open-webui/open-webui/issues/26315)
|
||||
- 🧠 **Model skills load on demand instead of filling the prompt.** A model's attached skills are now presented to the model as a manifest it can load when needed, rather than having their full content inserted into the system prompt; skills you mention inline still get their content included directly. [Commit](https://github.com/open-webui/open-webui/commit/e6d35fc4cca4f4b1e5cad97d7b7e3089ef832018), [Commit](https://github.com/open-webui/open-webui/commit/44b9463498085741669e6f5d92e21b5ecc5fd795), [#25592](https://github.com/open-webui/open-webui/issues/25592), [#25599](https://github.com/open-webui/open-webui/pull/25599)
|
||||
- 🗂️ **Empty metadata no longer breaks Chroma indexing.** Document metadata with empty values is now filtered out before indexing, fixing a case that could fail on Chroma. [Commit](https://github.com/open-webui/open-webui/commit/118549caf3), [#26342](https://github.com/open-webui/open-webui/pull/26342), [#26339](https://github.com/open-webui/open-webui/issues/26339)
|
||||
- 🔁 **Updating a knowledge file won't break the knowledge base.** When a file's content is updated, its new embeddings are now added before the old ones are removed, so a failed reindex leaves the knowledge base intact and usable instead of empty. [Commit](https://github.com/open-webui/open-webui/commit/248315de14d4537e0f2ec3f94dee8a7334cad248), [#23789](https://github.com/open-webui/open-webui/pull/23789), [#23787](https://github.com/open-webui/open-webui/issues/23787)
|
||||
- 🔤 **Documents with special tokens index correctly.** Measuring chunk sizes no longer fails when a document contains text that looks like a special token. [#26210](https://github.com/open-webui/open-webui/pull/26210)
|
||||
- 📝 **Note file attachments stay in sync.** Updating the files attached to a note now keeps the editor and saved note in sync. [Commit](https://github.com/open-webui/open-webui/commit/5055fb85aa8c8d5ef785daea7438498e36ddf33f)
|
||||
- 📱 **Better banner layout on mobile.** Notification banners now lay out correctly on small screens. [Commit](https://github.com/open-webui/open-webui/commit/4ed45ce84394c435405f93d03f07dabd797cdec3), [#24912](https://github.com/open-webui/open-webui/pull/24912)
|
||||
- 📂 **Knowledge file listing includes attached files.** Listing files through the knowledge tools now also shows files attached directly to a model, not only those inside a knowledge base, fixing cases where listing returned no results for a model with a single attached file. [Commit](https://github.com/open-webui/open-webui/commit/40b655e99e2c6dd802654ec0cdac38a4bcda08b3), [#26301](https://github.com/open-webui/open-webui/issues/26301)
|
||||
- 🏷️ **Chat titles generate after long first responses.** A new chat now gets its title even when the first response takes a long time, such as one with extensive reasoning or many tool calls, instead of staying "New Chat". [Commit](https://github.com/open-webui/open-webui/commit/754787f43dffad3dce2c90e4fd0417b1f9dbb3c0), [#26240](https://github.com/open-webui/open-webui/issues/26240)
|
||||
- 🔌 **Cancelling an MCP request no longer errors.** Stopping a response that was using MCP tools now shuts the connection down cleanly instead of surfacing a server error. [Commit](https://github.com/open-webui/open-webui/commit/ff5cec43bd360829cfdcc6a5253d1ba63f236b7f)
|
||||
- 🧠 **Reasoning details preserved across turns.** Models that return structured or encrypted reasoning data, such as Gemini, no longer have their assistant message split mid-stream, keeping reasoning continuity across turns. [Commit](https://github.com/open-webui/open-webui/commit/75db531c1238af113bb2b211882713e5e2f459cf), [#23852](https://github.com/open-webui/open-webui/pull/23852)
|
||||
- 📡 **Error messages show for non-standard streaming responses.** Providers that send errors over non-standard server-sent events now surface a readable error instead of nothing. [#23228](https://github.com/open-webui/open-webui/pull/23228)
|
||||
- 🔑 **Whitespace in terminal server keys no longer breaks auth.** Terminal server API keys are now trimmed before use, so a key with stray leading or trailing whitespace still authenticates. [Commit](https://github.com/open-webui/open-webui/commit/fe3300bd6581aa469c2cdf757700ecc65a200df4), [Commit](https://github.com/open-webui/open-webui/commit/d6cda4a04b2e3a48855fc91abb2376cfd3a0378d)
|
||||
- 🔥 **One bad URL no longer fails Firecrawl scraping.** When fetching multiple pages through Firecrawl, a single failing URL is now skipped instead of aborting the whole batch, and rate limits are respected between requests. [Commit](https://github.com/open-webui/open-webui/commit/6f8221df58b17334233ac6bfe069b8f837f677d6), [#24183](https://github.com/open-webui/open-webui/pull/24183)
|
||||
- 📱 **Usable chat input on mobile with many tools.** When skills, tools, terminal, web search, and image generation buttons fill the chat input, the row of buttons now scrolls horizontally while the menu, voice, and send controls stay reachable, instead of pushing them off-screen. [Commit](https://github.com/open-webui/open-webui/commit/6f8221df58b17334233ac6bfe069b8f837f677d6), [#26142](https://github.com/open-webui/open-webui/issues/26142)
|
||||
- 👤 **Owner avatars only show on shared folders.** Chat owner avatars in a folder's chat list now appear only when the folder is actually shared, instead of showing whenever owner information happened to be present. [Commit](https://github.com/open-webui/open-webui/commit/9802b0d13563b3535b86a350bab000d84686b1e9)
|
||||
- 📜 **No stray scrollbar on the About page.** Extra spacing that caused an unnecessary scrollbar on the About settings page has been removed. [#25802](https://github.com/open-webui/open-webui/pull/25802)
|
||||
- 🚪 **Sign out works from the Account Pending page.** Signing out while your account is pending now goes through the proper sign-out flow, so single sign-on sessions are ended and you are no longer left stuck on the pending screen. [#25681](https://github.com/open-webui/open-webui/pull/25681), [#25644](https://github.com/open-webui/open-webui/issues/25644)
|
||||
- 🔢 **Built-in tools accept numeric arguments.** Built-in tools no longer crash when a model passes a number or a string where a specific scalar type is expected; values are now coerced to the declared type. [Commit](https://github.com/open-webui/open-webui/commit/c4688b958d7f7929f5f4303493ca311c2c121683), [#25638](https://github.com/open-webui/open-webui/pull/25638), [#25731](https://github.com/open-webui/open-webui/pull/25731), [#25641](https://github.com/open-webui/open-webui/issues/25641)
|
||||
- ⏱️ **MinerU timeout saves.** The MinerU API timeout can now be saved from the admin settings, accepting a numeric value. [Commit](https://github.com/open-webui/open-webui/commit/3fd0384ffcd0eddd6f4c688475f8ad5d3b4de510), [#25604](https://github.com/open-webui/open-webui/pull/25604), [#25603](https://github.com/open-webui/open-webui/issues/25603)
|
||||
- 🔧 **Background completion no longer clears active tasks.** Finishing a chat in the background no longer wipes the set of active tasks, fixing a case where ongoing task indicators could be lost. [Commit](https://github.com/open-webui/open-webui/commit/388f62f8a002b789887d016892a1bf152c9d90af), [#25217](https://github.com/open-webui/open-webui/issues/25217)
|
||||
- 👁️ **Workspace base model selector respects visibility.** The base model selector in the workspace now hides models you don't have access to, matching their visibility settings. [#25668](https://github.com/open-webui/open-webui/pull/25668)
|
||||
- 🧵 **Channel threads bind to the right channel.** A channel thread's parent and replies are now tied to the channel in the URL, preventing mismatches when switching channels. [#25766](https://github.com/open-webui/open-webui/pull/25766)
|
||||
- 🗑️ **Unsharing cleans up orphaned rows.** Unsharing a chat now handles leftover shared-chat records, avoiding stale entries. [#25632](https://github.com/open-webui/open-webui/pull/25632)
|
||||
- 🔎 **Web search results reach the model with retrieval on.** Web search results are now passed to the model even when embedding and retrieval are enabled, instead of being left out. [#25600](https://github.com/open-webui/open-webui/pull/25600)
|
||||
- 🔢 **Group count follows search.** The groups count now reflects the filtered search results instead of the full list. [#25689](https://github.com/open-webui/open-webui/pull/25689)
|
||||
- ␣ **Space key works when renaming.** Pressing space while renaming a file or folder no longer opens it, so spaces can be typed in names. [#25627](https://github.com/open-webui/open-webui/pull/25627)
|
||||
- 🩹 **Missing local embedding model no longer blocks startup.** A missing local embedding model now surfaces as a deferred error instead of preventing the server from starting. [#25683](https://github.com/open-webui/open-webui/pull/25683)
|
||||
- 🔤 **Consistent settings label capitalization.** Toggle labels in settings now use consistent title casing. [#25765](https://github.com/open-webui/open-webui/pull/25765)
|
||||
- ♿ **Better screen-reader labels on toggles.** Integration and switch toggles now expose proper accessibility labels and pressed state for screen readers. [#25258](https://github.com/open-webui/open-webui/pull/25258), [#25230](https://github.com/open-webui/open-webui/pull/25230)
|
||||
- 📜 **Long dropdowns scroll.** Dropdown selects now scroll when their list is long, so all options stay reachable. [Commit](https://github.com/open-webui/open-webui/commit/4bc463072185d0d7c1c9218cd4487501090eccfe), [#25608](https://github.com/open-webui/open-webui/pull/25608)
|
||||
- 🔽 **Collapsible sections don't misfire on load.** Collapsible sections no longer trigger their change action when first rendered, avoiding unintended toggles on page load. [Commit](https://github.com/open-webui/open-webui/commit/c93d4f04aad1b0d4f8a8bda7ac403b2c8ee35f38), [#25229](https://github.com/open-webui/open-webui/pull/25229)
|
||||
- ➗ **Large math expressions no longer crash rendering.** Parsing math delimiters no longer overflows on very large or deeply nested input, so messages with heavy math render instead of failing. [#25845](https://github.com/open-webui/open-webui/pull/25845)
|
||||
- 🗄️ **Oversized chunks no longer break Milvus indexing.** Overly long text chunks are now trimmed before being sent to Milvus, so a single large chunk can no longer fail the whole batch and leave a file with no embeddings. [#25857](https://github.com/open-webui/open-webui/pull/25857), [#25858](https://github.com/open-webui/open-webui/pull/25858)
|
||||
- 📝 **Code editor stays open when empty.** The code editor drawer no longer collapses when its content is empty. [#25855](https://github.com/open-webui/open-webui/pull/25855)
|
||||
- 💽 **Settings no longer lost after a restart.** Admin configuration is now stored more reliably, fixing cases where external connections and model parameters could be lost after restarting the server. [Commit](https://github.com/open-webui/open-webui/commit/5cdcdbaeec9fc8156721c38c33ec37956962871c), [Commit](https://github.com/open-webui/open-webui/commit/21f9e5295bf484169d72f4538f7c926b5519723c), [Commit](https://github.com/open-webui/open-webui/commit/8958b64b5a7e96cd8c2260571b54324ca3bfe127), [#24743](https://github.com/open-webui/open-webui/issues/24743), [#25911](https://github.com/open-webui/open-webui/pull/25911), [#25959](https://github.com/open-webui/open-webui/pull/25959)
|
||||
- 📜 **Visible chat scrollbar.** The chat area now shows a scrollbar, making it easier to scroll through long responses. [Commit](https://github.com/open-webui/open-webui/commit/d56e1cb0b9), [#25833](https://github.com/open-webui/open-webui/issues/25833)
|
||||
- 🎚️ **Default model parameters apply to requests.** Default model parameters are now applied to outbound requests, so settings like temperature and the context window take effect as configured. [Commit](https://github.com/open-webui/open-webui/commit/cd6cc39c6d), [Commit](https://github.com/open-webui/open-webui/commit/19db873603215773f9a64e03785a1a076dc6c8a8), [#24930](https://github.com/open-webui/open-webui/issues/24930), [#26209](https://github.com/open-webui/open-webui/issues/26209)
|
||||
- 🟢 **Ollama loaded-model indicator restored.** The indicator showing which Ollama model is loaded in VRAM works again after recent changes. [#25586](https://github.com/open-webui/open-webui/issues/25586), [#25732](https://github.com/open-webui/open-webui/issues/25732)
|
||||
- 🪪 **Static MCP connectors recover missing OAuth details.** MCP connectors configured with static OAuth credentials now fill in a missing scope or resource from the server's published metadata, so they connect correctly instead of failing when those values were left out. [Commit](https://github.com/open-webui/open-webui/commit/88901bfa041ddcceab1cd4a97f08f0b43835eb05), [#25898](https://github.com/open-webui/open-webui/issues/25898)
|
||||
- 📊 **Token usage and cost stats no longer wiped by background tasks.** A response's token usage and cost are now preserved when background tasks like title, tag, and follow-up generation run on the same chat, instead of being overwritten. [Commit](https://github.com/open-webui/open-webui/commit/95391221dfabbfcd9090ab472b1c02a4c75c0387)
|
||||
- 🔗 **Model share link updated.** Sharing a model now opens the current community post page, fixing the link that pointed at the old endpoint. [#25801](https://github.com/open-webui/open-webui/pull/25801)
|
||||
|
||||
### Changed
|
||||
|
||||
- ⚠️ **Database Migrations**: This update contains database migrations. Please be sure to back up your database before updating, as downgrading after the migration is not supported.
|
||||
- 🔔 **System events now fire automatically.** With the new event system, Open WebUI emits events for activity like startup, sign-ins, and configuration changes, so any webhook you already have configured may begin receiving calls for these newly emitted events after upgrading. Review your event and webhook settings after updating so you only receive the events you want. [Commit](https://github.com/open-webui/open-webui/commit/b5c43968db0ea1556b228d143ae5946dc4e944ba)
|
||||
- 🔀 **Native tool calling is now the default.** Every chat and model that had not explicitly chosen a tool-calling mode now runs Native, which relies on a model's built-in tool support, while the old behavior has been renamed "Legacy" and made the explicit opt-out; if your models depend on the previous approach you must switch them back to "Legacy" per chat, per model, or globally in your default model parameters to preserve their behavior. [Commit](https://github.com/open-webui/open-webui/commit/b1d40f340921c27eb9a965b9feeb2563856e25e2)
|
||||
- 🗂️ **Authentication settings moved to their own page.** LDAP, OAuth, and related authentication settings have moved out of the General settings page into a dedicated Authentication page in the admin panel. [Commit](https://github.com/open-webui/open-webui/commit/5cdcdbaeec9fc8156721c38c33ec37956962871c)
|
||||
- 🎓 **Several features are no longer beta.** Memories, Notes, Channels, and High Contrast Mode have graduated out of beta and no longer carry a beta label. [Commit](https://github.com/open-webui/open-webui/commit/7b55a63fc7ee323e9114713ce1d2f3f688aa37e6)
|
||||
- 🔧 **Local web fetch setting renamed.** The "ENABLE_RAG_LOCAL_WEB_FETCH" environment variable is now "ENABLE_LOCAL_WEB_FETCH", reflecting that it applies beyond retrieval; the old name still works as a deprecated alias. [Commit](https://github.com/open-webui/open-webui/commit/e3ba6984534898695b47ee4fc3d6b746e2865abc)
|
||||
- 🔧 **You.com search key renamed.** You.com web search now prefers the "YDC_API_KEY" environment variable, with the previous "YOUCOM_API_KEY" still accepted as a fallback. [Commit](https://github.com/open-webui/open-webui/commit/df634bb64f5043b0292e43c69bd1d31676c89328), [#26316](https://github.com/open-webui/open-webui/pull/26316)
|
||||
- 🧪 **Client-side Python now runs sandboxed.** Client-side Python (Pyodide) now runs in a sandboxed, opaque-origin iframe by default, isolating executed code from your session, cookies, local storage, and the app's own endpoints, while full Python, JavaScript, and external network access keep working. Code that relied on reaching same-origin Open WebUI endpoints from Pyodide will no longer be able to, and Pyodide is now marked legacy in the admin Code Execution settings. [Commit](https://github.com/open-webui/open-webui/commit/516051304e1b1f250c34438746ade673a79bd40c), [Commit](https://github.com/open-webui/open-webui/commit/c7be66626fd10c75ec35f662a709129ba1b020ec), [Commit](https://github.com/open-webui/open-webui/commit/62ae2069183109d878d72b9444a0e7c4f6c66caa), [Commit](https://github.com/open-webui/open-webui/commit/518702caae5a6484e71aa79e8ab908ec398290a7), [Commit](https://github.com/open-webui/open-webui/commit/03a8363583b7e0e04760d49f1e8d28dbbfefee4d)
|
||||
|
||||
## [0.9.6] - 2026-06-01
|
||||
|
||||
### Added
|
||||
|
||||
+1
-1
@@ -126,7 +126,7 @@ RUN chown -R $UID:$GID /app $HOME
|
||||
# Install common system dependencies
|
||||
RUN apt-get update && \
|
||||
apt-get install -y --no-install-recommends \
|
||||
git build-essential pandoc gcc netcat-openbsd curl jq \
|
||||
git build-essential pandoc gcc netcat-openbsd curl jq ca-certificates \
|
||||
libmariadb-dev \
|
||||
python3-dev \
|
||||
ffmpeg libsm6 libxext6 zstd \
|
||||
|
||||
@@ -27,58 +27,84 @@ For more information, be sure to check out our [Open WebUI Documentation](https:
|
||||
|
||||
## Key Features of Open WebUI ⭐
|
||||
|
||||
- 🚀 **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.
|
||||
- 🚀 **Effortless Setup**: Install seamlessly via pip, uv, Docker, or Kubernetes (kubectl, kustomize, or helm), with `:ollama` and `:cuda` tagged images available for container deployments.
|
||||
|
||||
- 🤝 **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**.
|
||||
- 🤝 **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.
|
||||
|
||||
- 🛡️ **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.
|
||||
- 🔐 **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.
|
||||
|
||||
- 📱 **Responsive Design**: Enjoy a seamless experience across Desktop PC, Laptop, and Mobile devices.
|
||||
- 🧩 **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.
|
||||
|
||||
- 📱 **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.
|
||||
- 🤖 **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/).
|
||||
|
||||
- ✒️🔢 **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!
|
||||
|
||||
## 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 Terminal** ([open-webui/open-terminal](https://github.com/open-webui/open-terminal)): 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** · Enterprise ([open-webui/terminals](https://github.com/open-webui/terminals)): Per-user isolated containers with separate credentials, resource limits, and network rules. Automatic lifecycle management on Docker or Kubernetes.
|
||||
|
||||
- 💻 **cptr** ([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.
|
||||
|
||||
- 🔄 **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!
|
||||
@@ -222,6 +248,10 @@ 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
|
||||
[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
|
||||
|
||||
<a href="https://star-history.com/#open-webui/open-webui&Date">
|
||||
|
||||
@@ -11,6 +11,7 @@ import uvicorn
|
||||
app = typer.Typer()
|
||||
|
||||
KEY_FILE = Path.cwd() / '.webui_secret_key'
|
||||
DEFAULT_SECRET_KEY_LENGTH = 24
|
||||
|
||||
|
||||
def version_callback(value: bool) -> None:
|
||||
@@ -37,8 +38,11 @@ def serve(
|
||||
if os.getenv('WEBUI_SECRET_KEY') is None:
|
||||
typer.echo('Loading WEBUI_SECRET_KEY from file, not provided as an environment variable.')
|
||||
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}')
|
||||
KEY_FILE.write_bytes(base64.b64encode(random.randbytes(12)))
|
||||
KEY_FILE.write_bytes(base64.b64encode(random.randbytes(key_length)))
|
||||
typer.echo(f'Loading WEBUI_SECRET_KEY from {KEY_FILE}')
|
||||
os.environ['WEBUI_SECRET_KEY'] = KEY_FILE.read_text()
|
||||
|
||||
|
||||
+994
-1930
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,29 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import errno
|
||||
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):
|
||||
DEFAULT = lambda msg='': f'{msg if msg else ""}'
|
||||
MODEL_ADDED = lambda model='': f"The model '{model}' has been added successfully."
|
||||
@@ -18,7 +39,7 @@ class ERROR_MESSAGES(str, Enum):
|
||||
def __str__(self) -> str:
|
||||
return super().__str__()
|
||||
|
||||
DEFAULT = lambda err='': f'{"Something went wrong :/" if err == "" else "[ERROR: " + str(err) + "]"}'
|
||||
DEFAULT = _error_message
|
||||
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.'
|
||||
DELETE_USER_ERROR = 'Oops! Something went wrong. We encountered an issue while trying to delete the user. Please give it another shot.'
|
||||
|
||||
@@ -291,6 +291,7 @@ if 'postgres://' in DATABASE_URL:
|
||||
DATABASE_URL = DATABASE_URL.replace('postgres://', 'postgresql://')
|
||||
|
||||
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')
|
||||
try:
|
||||
@@ -498,6 +499,65 @@ else:
|
||||
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'
|
||||
|
||||
_aiohttp_timeout_raw = os.getenv('AIOHTTP_CLIENT_TIMEOUT', '')
|
||||
@@ -507,7 +567,10 @@ except (ValueError, TypeError):
|
||||
AIOHTTP_CLIENT_TIMEOUT = 300
|
||||
|
||||
|
||||
AIOHTTP_CLIENT_SESSION_SSL = os.getenv('AIOHTTP_CLIENT_SESSION_SSL', 'True').lower() == 'true'
|
||||
# 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'))
|
||||
|
||||
# When False (default), outbound HTTP requests do not follow 3xx redirects.
|
||||
AIOHTTP_CLIENT_ALLOW_REDIRECTS = os.getenv('AIOHTTP_CLIENT_ALLOW_REDIRECTS', 'False').lower() == 'true'
|
||||
@@ -533,7 +596,10 @@ except (ValueError, TypeError):
|
||||
AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER_DATA = 10
|
||||
|
||||
|
||||
AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL = os.getenv('AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL', 'True').lower() == 'true'
|
||||
# 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', '')
|
||||
|
||||
@@ -616,6 +682,8 @@ WEBUI_SECRET_KEY = os.getenv(
|
||||
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_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)
|
||||
@@ -662,6 +730,7 @@ 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')
|
||||
|
||||
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',
|
||||
r'^(?=.*[a-z])(?=.*[A-Z])(?=.*\d)(?=.*[^\w\s]).{8,}$',
|
||||
@@ -686,6 +755,9 @@ BYPASS_RETRIEVAL_ACCESS_CONTROL = os.getenv('BYPASS_RETRIEVAL_ACCESS_CONTROL', '
|
||||
# for non-admin users. When False (default), unknown collection names are
|
||||
# denied — closing the legacy unscoped namespace.
|
||||
ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS = os.getenv('ENABLE_RETRIEVAL_UNSCOPED_COLLECTIONS', 'False').lower() == 'true'
|
||||
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,
|
||||
# and chunked splitting) before sending files to processing engines. Useful when
|
||||
@@ -862,6 +934,7 @@ else:
|
||||
ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION = (
|
||||
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'
|
||||
|
||||
# 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
|
||||
@@ -985,6 +1058,12 @@ if OFFLINE_MODE:
|
||||
os.environ['HF_HUB_OFFLINE'] = '1'
|
||||
ENABLE_VERSION_UPDATE_CHECK = False
|
||||
|
||||
####################################
|
||||
# Pyodide file persistence
|
||||
####################################
|
||||
|
||||
ENABLE_PYODIDE_FILE_PERSISTENCE = os.getenv('ENABLE_PYODIDE_FILE_PERSISTENCE', 'false').lower() == 'true'
|
||||
|
||||
####################################
|
||||
# Audit logging
|
||||
####################################
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,265 +0,0 @@
|
||||
"""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)
|
||||
@@ -5,10 +5,12 @@ import logging
|
||||
import os
|
||||
import sys
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
|
||||
|
||||
from open_webui.env import (
|
||||
DATABASE_ENABLE_IAM_TOKEN_AUTH,
|
||||
DATABASE_ENABLE_SESSION_SHARING,
|
||||
DATABASE_ENABLE_SQLITE_WAL,
|
||||
DATABASE_POOL_MAX_OVERFLOW,
|
||||
@@ -27,6 +29,7 @@ from open_webui.env import (
|
||||
OPEN_WEBUI_DIR,
|
||||
)
|
||||
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.declarative import declarative_base
|
||||
from sqlalchemy.orm import Session, scoped_session, sessionmaker
|
||||
@@ -146,6 +149,63 @@ _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
|
||||
|
||||
|
||||
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)
|
||||
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:
|
||||
"""Convert a sync database URL to its async driver equivalent.
|
||||
|
||||
@@ -268,6 +328,8 @@ else:
|
||||
else:
|
||||
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)
|
||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
|
||||
@@ -344,6 +406,8 @@ else:
|
||||
pool_pre_ping=True,
|
||||
)
|
||||
|
||||
enable_iam_token_auth(async_engine)
|
||||
|
||||
|
||||
AsyncSessionLocal = async_sessionmaker(
|
||||
bind=async_engine,
|
||||
|
||||
+566
-981
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,7 @@ import logging.config
|
||||
import logging
|
||||
import alembic.context
|
||||
from open_webui.env import DATABASE_PASSWORD, DATABASE_URL, LOG_FORMAT
|
||||
from open_webui.internal.db import extract_ssl_params_from_url, reattach_ssl_params_to_url
|
||||
from open_webui.internal.db import enable_iam_token_auth, extract_ssl_params_from_url, reattach_ssl_params_to_url
|
||||
from open_webui.models.auths import Auth
|
||||
from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401
|
||||
from sqlalchemy import create_engine, engine_from_config, pool
|
||||
@@ -68,6 +68,7 @@ def _get_engine_connectable():
|
||||
def run_migrations_online() -> None:
|
||||
"""Execute migrations against a live database connection."""
|
||||
live_connectable = _get_engine_connectable()
|
||||
enable_iam_token_auth(live_connectable)
|
||||
with live_connectable.connect() as live_connection:
|
||||
alembic.context.configure(
|
||||
connection=live_connection,
|
||||
|
||||
@@ -49,7 +49,7 @@ def upgrade():
|
||||
|
||||
# Step 3: Migrate data from 'old_chat' to 'chat' (only if old_chat exists)
|
||||
# Re-check columns after potential rename above
|
||||
current_cols = {c['name'] for c in inspector.get_columns('chat')}
|
||||
current_cols = {c['name'] for c in sa.inspect(conn).get_columns('chat')}
|
||||
if 'old_chat' in current_cols:
|
||||
chat_table = table(
|
||||
'chat',
|
||||
@@ -76,8 +76,12 @@ def upgrade():
|
||||
|
||||
|
||||
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
|
||||
op.add_column('chat', sa.Column('old_chat', sa.Text(), nullable=True))
|
||||
if 'old_chat' not in columns:
|
||||
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'
|
||||
chat_table = table(
|
||||
@@ -87,14 +91,14 @@ def downgrade():
|
||||
sa.Column('old_chat', sa.Text()),
|
||||
)
|
||||
|
||||
connection = op.get_bind()
|
||||
results = connection.execute(select(chat_table.c.id, chat_table.c.chat))
|
||||
for row in results:
|
||||
text_data = json.dumps(row.chat) if row.chat is not None else None
|
||||
connection.execute(sa.update(chat_table).where(chat_table.c.id == row.id).values(old_chat=text_data))
|
||||
if 'chat' in columns:
|
||||
results = conn.execute(select(chat_table.c.id, chat_table.c.chat))
|
||||
for row in results:
|
||||
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))
|
||||
|
||||
# Step 3: Remove the new 'chat' JSON column
|
||||
op.drop_column('chat', 'chat')
|
||||
# Step 3: Remove the new 'chat' JSON column
|
||||
op.drop_column('chat', 'chat')
|
||||
|
||||
# Step 4: Rename 'old_chat' back to 'chat'
|
||||
op.alter_column('chat', 'old_chat', new_column_name='chat', existing_type=sa.Text())
|
||||
|
||||
+584
@@ -0,0 +1,584 @@
|
||||
"""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}])
|
||||
@@ -0,0 +1,40 @@
|
||||
"""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')
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
"""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')
|
||||
@@ -0,0 +1,44 @@
|
||||
"""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')
|
||||
@@ -6,6 +6,7 @@ import logging
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
import bcrypt
|
||||
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.utils.validate import validate_profile_image_url
|
||||
@@ -15,6 +16,11 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
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
|
||||
"""Maps a user ID to an email/password pair with an active flag."""
|
||||
@@ -142,13 +148,15 @@ class AuthsTable:
|
||||
log.info('authenticate_user: %s', email)
|
||||
resolved = await Users.get_user_by_email(email, db=db)
|
||||
if not resolved:
|
||||
await verify_password(PLACEHOLDER_HASH)
|
||||
return
|
||||
# load the credential row and verify the password hash
|
||||
async with get_async_db_context(db) as session:
|
||||
credential = await session.get(Auth, resolved.id)
|
||||
if not credential or not credential.active:
|
||||
await verify_password(PLACEHOLDER_HASH)
|
||||
return
|
||||
if not verify_password(credential.password):
|
||||
if not await verify_password(credential.password):
|
||||
return
|
||||
return resolved
|
||||
|
||||
|
||||
@@ -3,8 +3,10 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
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.utils.response import normalize_usage
|
||||
from open_webui.utils.response import merge_usage, normalize_usage
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import (
|
||||
JSON,
|
||||
@@ -107,6 +109,9 @@ class ChatMessage(Base):
|
||||
# Usage (tokens, timing, etc.)
|
||||
usage = Column(JSON, nullable=True)
|
||||
|
||||
# Context compaction checkpoint
|
||||
context_summary = Column(Text, nullable=True)
|
||||
|
||||
# Timestamps
|
||||
created_at = Column(BigInteger, index=True)
|
||||
updated_at = Column(BigInteger)
|
||||
@@ -141,6 +146,7 @@ class ChatMessageModel(BaseModel):
|
||||
status_history: Optional[list] = None
|
||||
error: Optional[dict | str] = None
|
||||
usage: Optional[dict] = None
|
||||
context_summary: Optional[str] = None
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
@@ -192,13 +198,13 @@ class ChatMessageTable:
|
||||
existing.status_history = data.get('status_history') or data.get('statusHistory')
|
||||
if 'error' in data:
|
||||
existing.error = data.get('error')
|
||||
if 'context_summary' in data or 'contextSummary' in data:
|
||||
existing.context_summary = data.get('context_summary') or data.get('contextSummary')
|
||||
# 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_usage = normalize_usage(existing.usage or {}) if existing.usage else {}
|
||||
existing.usage = existing_usage if usage == existing_usage else merge_usage(existing_usage, usage)
|
||||
existing.updated_at = now
|
||||
await db.commit()
|
||||
await db.refresh(existing)
|
||||
@@ -223,6 +229,7 @@ class ChatMessageTable:
|
||||
status_history=data.get('status_history') or data.get('statusHistory'),
|
||||
error=data.get('error'),
|
||||
usage=usage,
|
||||
context_summary=data.get('context_summary') or data.get('contextSummary'),
|
||||
created_at=timestamp,
|
||||
updated_at=now,
|
||||
)
|
||||
@@ -249,6 +256,7 @@ class ChatMessageTable:
|
||||
'parent_id': 'parentId',
|
||||
'model_id': 'model',
|
||||
'status_history': 'statusHistory',
|
||||
'context_summary': 'contextSummary',
|
||||
'created_at': 'timestamp',
|
||||
}
|
||||
# DB-internal columns excluded from the reconstructed message dict.
|
||||
@@ -440,6 +448,44 @@ class ChatMessageTable:
|
||||
result = await db.execute(stmt)
|
||||
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(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
|
||||
@@ -34,6 +34,7 @@ from sqlalchemy import (
|
||||
update,
|
||||
)
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
from sqlalchemy.sql import exists
|
||||
from sqlalchemy.sql.expression import bindparam
|
||||
|
||||
@@ -179,6 +180,7 @@ class ChatTitleIdResponse(BaseModel):
|
||||
updated_at: int
|
||||
created_at: int
|
||||
last_read_at: int | None = None
|
||||
snippet: str | None = None
|
||||
|
||||
|
||||
class SharedChatResponse(BaseModel):
|
||||
@@ -294,6 +296,58 @@ class ChatTable:
|
||||
|
||||
return changed
|
||||
|
||||
def _repair_chat_current_id(self, chat: dict) -> bool:
|
||||
history = chat.get('history')
|
||||
if not isinstance(history, dict):
|
||||
return False
|
||||
|
||||
messages = history.get('messages')
|
||||
if not isinstance(messages, dict):
|
||||
return False
|
||||
|
||||
current_id = history.get('currentId')
|
||||
current_message = messages.get(current_id)
|
||||
output = []
|
||||
if isinstance(current_message, dict):
|
||||
output = current_message.get('output') or []
|
||||
|
||||
output_role = next(
|
||||
(item.get('role') for item in output if isinstance(item, dict) and item.get('role')),
|
||||
None,
|
||||
)
|
||||
current_is_bad_leaf = (
|
||||
isinstance(current_message, dict)
|
||||
and output_role == 'assistant'
|
||||
and current_message.get('parentId') is None
|
||||
and not current_message.get('timestamp')
|
||||
and len(messages) > 1
|
||||
)
|
||||
if (
|
||||
isinstance(current_message, dict)
|
||||
and current_message.get('id')
|
||||
and current_message.get('role')
|
||||
and not current_is_bad_leaf
|
||||
):
|
||||
return False
|
||||
|
||||
latest_leaf_id = None
|
||||
latest_timestamp = -1
|
||||
for message_id, message in messages.items():
|
||||
if not isinstance(message, dict) or not message.get('role'):
|
||||
continue
|
||||
|
||||
children_ids = message.get('childrenIds') if isinstance(message.get('childrenIds'), list) else []
|
||||
timestamp = message.get('timestamp') or 0
|
||||
if len(children_ids) == 0 and timestamp > latest_timestamp:
|
||||
latest_leaf_id = message_id
|
||||
latest_timestamp = timestamp
|
||||
|
||||
if not latest_leaf_id or latest_leaf_id == current_id:
|
||||
return False
|
||||
|
||||
history['currentId'] = latest_leaf_id
|
||||
return True
|
||||
|
||||
async def insert_new_chat(
|
||||
self, id: str, user_id: str, form_data: ChatForm, db: AsyncSession | None = None
|
||||
) -> ChatModel | None:
|
||||
@@ -309,6 +363,7 @@ class ChatTable:
|
||||
'folder_id': form_data.folder_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
'last_read_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -445,7 +500,6 @@ class ChatTable:
|
||||
clean_title = self._clean_null_bytes(title)
|
||||
chat_item.title = clean_title
|
||||
chat_item.chat = {**(chat_item.chat or {}), 'title': clean_title}
|
||||
chat_item.updated_at = int(time.time())
|
||||
await session.commit()
|
||||
await session.refresh(chat_item)
|
||||
return ChatModel.model_validate(chat_item)
|
||||
@@ -497,6 +551,68 @@ class ChatTable:
|
||||
if msg.get('parentId') and msg['parentId'] not in messages_map
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def merge_history(existing_history: dict | None, incoming_history: dict | None) -> dict:
|
||||
existing = (existing_history or {}).get('messages') or {}
|
||||
incoming = (incoming_history or {}).get('messages') or {}
|
||||
merged = {**existing, **incoming}
|
||||
merged = {message_id: message for message_id, message in merged.items() if isinstance(message, dict)}
|
||||
|
||||
for message in merged.values():
|
||||
message['childrenIds'] = []
|
||||
for message_id, message in merged.items():
|
||||
parent_id = message.get('parentId')
|
||||
if parent_id in merged:
|
||||
merged[parent_id]['childrenIds'].append(message_id)
|
||||
|
||||
current_id = (incoming_history or {}).get('currentId')
|
||||
if current_id not in merged:
|
||||
current_id = (existing_history or {}).get('currentId')
|
||||
if current_id not in merged:
|
||||
current_id = None
|
||||
|
||||
return {**(existing_history or {}), **(incoming_history or {}), 'messages': merged, 'currentId': current_id}
|
||||
|
||||
@staticmethod
|
||||
def delete_message_from_history(history: dict, message_id: str) -> set[str]:
|
||||
messages = history.get('messages') or {}
|
||||
message = messages.get(message_id)
|
||||
if not isinstance(message, dict):
|
||||
return set()
|
||||
|
||||
parent_id = message.get('parentId')
|
||||
child_ids = [child_id for child_id in (message.get('childrenIds') or []) if child_id in messages]
|
||||
grandchild_ids = [
|
||||
grandchild_id
|
||||
for child_id in child_ids
|
||||
for grandchild_id in (messages.get(child_id, {}).get('childrenIds') or [])
|
||||
if grandchild_id in messages
|
||||
]
|
||||
|
||||
if parent_id in messages:
|
||||
messages[parent_id]['childrenIds'] = [
|
||||
child_id for child_id in (messages[parent_id].get('childrenIds') or []) if child_id != message_id
|
||||
] + grandchild_ids
|
||||
|
||||
for grandchild_id in grandchild_ids:
|
||||
messages[grandchild_id]['parentId'] = parent_id
|
||||
|
||||
deleted_ids = {message_id, *child_ids}
|
||||
for deleted_id in deleted_ids:
|
||||
messages.pop(deleted_id, None)
|
||||
|
||||
current_id = parent_id
|
||||
child_ids = (
|
||||
[child_id for child_id, child in messages.items() if child.get('parentId') is None]
|
||||
if current_id is None
|
||||
else messages.get(current_id, {}).get('childrenIds', [])
|
||||
)
|
||||
while child_ids:
|
||||
current_id = child_ids[-1]
|
||||
child_ids = messages.get(current_id, {}).get('childrenIds', [])
|
||||
history['currentId'] = current_id if current_id in messages else None
|
||||
return deleted_ids
|
||||
|
||||
async def backfill_messages_by_chat_id(self, chat_id: str, user_id: str, messages: dict[str, dict]) -> None:
|
||||
"""Write messages to the ``chat_message`` table so future lookups
|
||||
use the fast path. Errors are logged but never raised.
|
||||
@@ -517,18 +633,11 @@ class ChatTable:
|
||||
async def reconcile_messages_by_chat_id(self, chat_id: str, user_id: str, messages: dict[str, dict]) -> None:
|
||||
"""Sync ``chat_message`` rows with the committed JSON blob.
|
||||
|
||||
Upserts current messages via ``backfill_messages_by_chat_id``
|
||||
and deletes orphaned rows whose message_id no longer appears
|
||||
in the blob. Best-effort: errors are logged but never raised.
|
||||
Upserts current messages via ``backfill_messages_by_chat_id``.
|
||||
Best-effort: errors are logged but never raised.
|
||||
"""
|
||||
try:
|
||||
await self.backfill_messages_by_chat_id(chat_id, user_id, messages)
|
||||
|
||||
existing_map = await ChatMessages.get_messages_map_by_chat_id(chat_id)
|
||||
if existing_map is not None:
|
||||
orphaned_ids = set(existing_map.keys()) - set(messages.keys())
|
||||
if orphaned_ids:
|
||||
await ChatMessages.delete_message_ids_by_chat_id(chat_id, orphaned_ids)
|
||||
except Exception as e:
|
||||
log.warning('Failed to reconcile chat_message rows for chat %s: %s', chat_id, e)
|
||||
|
||||
@@ -606,14 +715,46 @@ class ChatTable:
|
||||
user_id = chat.user_id
|
||||
chat = chat.chat
|
||||
history = chat.get('history', {})
|
||||
messages = history.setdefault('messages', {})
|
||||
|
||||
if message_id in history.get('messages', {}):
|
||||
history['messages'][message_id] = {
|
||||
**history['messages'][message_id],
|
||||
if message_id in messages:
|
||||
messages[message_id] = {
|
||||
**messages[message_id],
|
||||
**message,
|
||||
}
|
||||
else:
|
||||
history['messages'][message_id] = message
|
||||
message_parent_id = message.get('parentId')
|
||||
parent_id = message_parent_id
|
||||
if parent_id is None:
|
||||
for existing_id, existing_message in messages.items():
|
||||
if message_id in existing_message.get('childrenIds', []):
|
||||
parent_id = existing_id
|
||||
break
|
||||
|
||||
parent = messages.get(parent_id) if parent_id else None
|
||||
output = message.get('output') or []
|
||||
output_role = next(
|
||||
(item.get('role') for item in output if isinstance(item, dict) and item.get('role')),
|
||||
None,
|
||||
)
|
||||
role = message.get('role') or output_role
|
||||
if not role:
|
||||
parent_role = parent.get('role') if parent else None
|
||||
if parent_role == 'user':
|
||||
role = 'assistant'
|
||||
elif parent_role == 'assistant':
|
||||
role = 'user'
|
||||
else:
|
||||
role = 'assistant'
|
||||
|
||||
messages[message_id] = {
|
||||
**message,
|
||||
'id': message.get('id') or message_id,
|
||||
'parentId': message_parent_id if message_parent_id is not None else parent_id,
|
||||
'childrenIds': message.get('childrenIds') if isinstance(message.get('childrenIds'), list) else [],
|
||||
'role': role,
|
||||
'timestamp': message.get('timestamp') or int(time.time()),
|
||||
}
|
||||
|
||||
history['currentId'] = message_id
|
||||
|
||||
@@ -625,13 +766,33 @@ class ChatTable:
|
||||
message_id=message_id,
|
||||
chat_id=id,
|
||||
user_id=user_id,
|
||||
data=history['messages'][message_id],
|
||||
data=messages[message_id],
|
||||
)
|
||||
except Exception as e:
|
||||
log.warning(f'Failed to write to chat_message table: {e}')
|
||||
|
||||
return await self.update_chat_by_id(id, chat)
|
||||
|
||||
async def delete_message_from_chat_by_id_and_message_id(self, id: str, message_id: str) -> ChatModel | None:
|
||||
chat_model = await self.get_chat_by_id(id)
|
||||
if chat_model is None:
|
||||
return None
|
||||
|
||||
chat = chat_model.chat
|
||||
history = chat.get('history', {})
|
||||
deleted_ids = self.delete_message_from_history(history, message_id)
|
||||
if not deleted_ids:
|
||||
return chat_model
|
||||
|
||||
messages = history.get('messages') or {}
|
||||
chat['history'] = history
|
||||
updated_chat = await self.update_chat_by_id(id, chat)
|
||||
|
||||
await self.backfill_messages_by_chat_id(id, chat_model.user_id, messages)
|
||||
await ChatMessages.delete_message_ids_by_chat_id(id, deleted_ids)
|
||||
|
||||
return updated_chat
|
||||
|
||||
async def add_message_status_to_chat_by_id_and_message_id(
|
||||
self, id: str, message_id: str, status: dict
|
||||
) -> ChatModel | None:
|
||||
@@ -748,6 +909,7 @@ class ChatTable:
|
||||
chat = await session.get(Chat, id)
|
||||
chat.pinned = not chat.pinned
|
||||
chat.updated_at = int(time.time())
|
||||
chat.last_read_at = int(time.time())
|
||||
await session.commit()
|
||||
await session.refresh(chat)
|
||||
return ChatModel.model_validate(chat)
|
||||
@@ -761,6 +923,7 @@ class ChatTable:
|
||||
chat.archived = not chat.archived
|
||||
chat.folder_id = None
|
||||
chat.updated_at = int(time.time())
|
||||
chat.last_read_at = int(time.time())
|
||||
await session.commit()
|
||||
await session.refresh(chat)
|
||||
return ChatModel.model_validate(chat)
|
||||
@@ -829,6 +992,15 @@ class ChatTable:
|
||||
for chat in all_chats
|
||||
]
|
||||
|
||||
async def count_archived_chats_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
db: AsyncSession | None = None,
|
||||
) -> int:
|
||||
async with get_async_db_context(db) as session:
|
||||
result = await session.execute(select(func.count(Chat.id)).filter_by(user_id=user_id, archived=True))
|
||||
return result.scalar() or 0
|
||||
|
||||
async def get_shared_chat_list_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
@@ -957,6 +1129,85 @@ class ChatTable:
|
||||
all_chats = result.scalars().all()
|
||||
return [ChatModel.model_validate(chat) for chat in all_chats]
|
||||
|
||||
async def get_chat_metas_by_chat_ids(
|
||||
self,
|
||||
chat_ids: list[str],
|
||||
include_archived: bool = False,
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[dict]:
|
||||
async with get_async_db_context(db) as session:
|
||||
stmt = select(Chat.meta).filter(Chat.id.in_(chat_ids))
|
||||
if not include_archived:
|
||||
stmt = stmt.filter_by(archived=False)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return [meta for meta in result.scalars().all() if isinstance(meta, dict)]
|
||||
|
||||
async def get_chats_by_model_id(
|
||||
self,
|
||||
model_id: str,
|
||||
filter: dict | None = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: AsyncSession | None = None,
|
||||
) -> dict:
|
||||
from open_webui.models.users import User
|
||||
|
||||
async with get_async_db_context(db) as session:
|
||||
chat_ids = (
|
||||
select(ChatMessage.chat_id).filter(ChatMessage.model_id == model_id).group_by(ChatMessage.chat_id)
|
||||
)
|
||||
|
||||
if filter:
|
||||
if filter.get('start_date'):
|
||||
chat_ids = chat_ids.filter(ChatMessage.created_at >= filter.get('start_date'))
|
||||
if filter.get('end_date'):
|
||||
chat_ids = chat_ids.filter(ChatMessage.created_at <= filter.get('end_date'))
|
||||
|
||||
chat_ids = chat_ids.subquery()
|
||||
|
||||
stmt = (
|
||||
select(Chat.id, Chat.user_id, Chat.title, Chat.updated_at, User.name.label('user_name'))
|
||||
.join(chat_ids, chat_ids.c.chat_id == Chat.id)
|
||||
.outerjoin(User, User.id == Chat.user_id)
|
||||
)
|
||||
|
||||
order_by = filter.get('order_by') if filter else None
|
||||
direction = filter.get('direction') if filter else None
|
||||
is_asc = direction == 'asc'
|
||||
|
||||
if order_by == 'title':
|
||||
primary_sort = Chat.title.asc() if is_asc else Chat.title.desc()
|
||||
elif order_by == 'user_name':
|
||||
primary_sort = User.name.asc() if is_asc else User.name.desc()
|
||||
else:
|
||||
primary_sort = Chat.updated_at.asc() if is_asc else Chat.updated_at.desc()
|
||||
|
||||
stmt = stmt.order_by(primary_sort, Chat.id.asc())
|
||||
|
||||
count_result = await session.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
return {
|
||||
'items': [
|
||||
{
|
||||
'chat_id': chat.id,
|
||||
'user_id': chat.user_id,
|
||||
'user_name': chat.user_name,
|
||||
'first_message': chat.title,
|
||||
'updated_at': chat.updated_at,
|
||||
}
|
||||
for chat in result.all()
|
||||
],
|
||||
'total': total,
|
||||
}
|
||||
|
||||
# retrieve conversation
|
||||
async def get_chat_by_id(
|
||||
self,
|
||||
@@ -970,7 +1221,10 @@ class ChatTable:
|
||||
if chat_item is None:
|
||||
return None
|
||||
|
||||
if self._sanitize_chat_row(chat_item):
|
||||
repaired_history = self._repair_chat_current_id(chat_item.chat or {})
|
||||
if repaired_history:
|
||||
flag_modified(chat_item, 'chat')
|
||||
if self._sanitize_chat_row(chat_item) or repaired_history:
|
||||
await session.commit()
|
||||
await session.refresh(chat_item)
|
||||
|
||||
@@ -1006,7 +1260,17 @@ class ChatTable:
|
||||
async with get_async_db_context(db) as session:
|
||||
result = await session.execute(select(Chat).filter_by(id=id, user_id=user_id))
|
||||
chat = result.scalars().first()
|
||||
return ChatModel.model_validate(chat) if chat else None
|
||||
if not chat:
|
||||
return None
|
||||
|
||||
repaired_history = self._repair_chat_current_id(chat.chat or {})
|
||||
if repaired_history:
|
||||
flag_modified(chat, 'chat')
|
||||
if self._sanitize_chat_row(chat) or repaired_history:
|
||||
await session.commit()
|
||||
await session.refresh(chat)
|
||||
|
||||
return ChatModel.model_validate(chat)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@@ -1353,6 +1617,42 @@ class ChatTable:
|
||||
for chat in all_chats
|
||||
]
|
||||
|
||||
async def get_all_chats_by_folder_id(
|
||||
self,
|
||||
folder_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 60,
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[dict]:
|
||||
"""Get chats in a folder across ALL users. Returns dicts with user_id."""
|
||||
async with get_async_db_context(db) as session:
|
||||
stmt = (
|
||||
select(Chat.id, Chat.title, Chat.user_id, Chat.updated_at, Chat.created_at, Chat.last_read_at)
|
||||
.filter_by(folder_id=folder_id)
|
||||
.filter(or_(Chat.pinned == False, Chat.pinned == None))
|
||||
.filter_by(archived=False)
|
||||
.order_by(Chat.updated_at.desc(), Chat.id)
|
||||
)
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await session.execute(stmt)
|
||||
all_chats = result.all()
|
||||
return [
|
||||
{
|
||||
'id': chat[0],
|
||||
'title': chat[1],
|
||||
'user_id': chat[2],
|
||||
'updated_at': chat[3],
|
||||
'created_at': chat[4],
|
||||
'last_read_at': chat[5],
|
||||
}
|
||||
for chat in all_chats
|
||||
]
|
||||
|
||||
async def get_chats_by_folder_ids_and_user_id(
|
||||
self, folder_ids: list[str], user_id: str, db: AsyncSession | None = None
|
||||
) -> list[ChatModel]:
|
||||
@@ -1377,6 +1677,7 @@ class ChatTable:
|
||||
chat = await session.get(Chat, id)
|
||||
chat.folder_id = folder_id
|
||||
chat.updated_at = int(time.time())
|
||||
chat.last_read_at = int(time.time())
|
||||
chat.pinned = False
|
||||
await session.commit()
|
||||
await session.refresh(chat)
|
||||
@@ -1521,6 +1822,21 @@ class ChatTable:
|
||||
log.info(f"Count of chats for folder '{folder_id}': {count}")
|
||||
return count
|
||||
|
||||
async def count_chats_by_folder_ids_and_user_id(
|
||||
self, folder_ids: list[str], user_id: str, db: AsyncSession | None = None
|
||||
) -> int:
|
||||
if not folder_ids:
|
||||
return 0
|
||||
|
||||
async with get_async_db_context(db) as session:
|
||||
result = await session.execute(
|
||||
select(func.count(Chat.id)).filter(Chat.user_id == user_id, Chat.folder_id.in_(folder_ids))
|
||||
)
|
||||
count = result.scalar()
|
||||
|
||||
log.info(f"Count of chats for folders '{folder_ids}': {count}")
|
||||
return count
|
||||
|
||||
async def delete_tag_by_id_and_user_id_and_tag_name(
|
||||
self, id: str, user_id: str, tag_name: str, db: AsyncSession | None = None
|
||||
) -> bool:
|
||||
|
||||
@@ -0,0 +1,343 @@
|
||||
"""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 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',),
|
||||
'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_version',
|
||||
'extra_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
|
||||
|
||||
|
||||
# ── 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 = 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."""
|
||||
async with get_async_db() as db:
|
||||
now = int(time.time())
|
||||
for key, value in 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():
|
||||
if key not in existing_keys:
|
||||
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_flattened_dict_configs() -> None:
|
||||
"""Reassemble dict config values flattened by the per-key migration."""
|
||||
if not Config.PERSISTENT_ENABLED:
|
||||
return
|
||||
|
||||
async with get_async_db() as db:
|
||||
repaired_keys: list[str] = []
|
||||
orphan_keys: list[str] = []
|
||||
|
||||
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 = int(time.time())
|
||||
else:
|
||||
db.add(Config(key=config_key, value=repaired, updated_at=int(time.time())))
|
||||
repaired_keys.append(config_key)
|
||||
|
||||
if orphan_keys:
|
||||
await db.execute(delete(Config).where(Config.key.in_(orphan_keys)))
|
||||
|
||||
if repaired_keys or orphan_keys:
|
||||
await db.commit()
|
||||
log.info('Repaired flattened dict config rows for %s', ', '.join(repaired_keys))
|
||||
@@ -133,6 +133,12 @@ class ModelHistoryEntry(BaseModel):
|
||||
lost: int
|
||||
|
||||
|
||||
class ModelHistoryCounts(BaseModel):
|
||||
date: str
|
||||
won: int = 0
|
||||
lost: int = 0
|
||||
|
||||
|
||||
class ModelHistoryResponse(BaseModel):
|
||||
model_id: str
|
||||
history: list[ModelHistoryEntry]
|
||||
@@ -216,12 +222,15 @@ class FeedbackTable:
|
||||
) -> FeedbackListResponse:
|
||||
async with get_async_db_context(db) as db:
|
||||
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:
|
||||
# Apply model_id filter (exact match)
|
||||
model_id = filter.get('model_id')
|
||||
if model_id:
|
||||
stmt = stmt.filter(Feedback.data['model_id'].as_string() == model_id)
|
||||
model_id_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')
|
||||
direction = filter.get('direction')
|
||||
@@ -250,9 +259,9 @@ class FeedbackTable:
|
||||
else:
|
||||
stmt = stmt.order_by(Feedback.created_at.desc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
# Count before pagination without wrapping the ordered item query.
|
||||
count_result = await db.execute(count_stmt)
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
@@ -375,6 +384,45 @@ class FeedbackTable:
|
||||
|
||||
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 with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()))
|
||||
|
||||
@@ -201,6 +201,18 @@ class FilesTable:
|
||||
result = await db.execute(select(File))
|
||||
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:
|
||||
file = await self.get_file_by_id(id, db=db)
|
||||
if not file:
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Optional
|
||||
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import JSON, BigInteger, Boolean, Column, Text, delete, func, select
|
||||
from sqlalchemy import JSON, BigInteger, Boolean, Column, Text, delete, func, select, or_, and_
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -62,6 +62,20 @@ class FolderNameIdResponse(BaseModel):
|
||||
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
|
||||
updated_at: int
|
||||
|
||||
|
||||
####################
|
||||
# Forms
|
||||
####################
|
||||
@@ -130,6 +144,52 @@ class FolderTable:
|
||||
except Exception:
|
||||
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_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(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[list[FolderModel]]:
|
||||
@@ -188,6 +248,25 @@ class FolderTable:
|
||||
result = await db.execute(select(Folder).filter_by(parent_id=parent_id, user_id=user_id))
|
||||
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)
|
||||
folder_ids.extend(child.id for child in children)
|
||||
folders.extend(children)
|
||||
|
||||
return folder_ids
|
||||
|
||||
async def update_folder_parent_id_by_id_and_user_id(
|
||||
self,
|
||||
id: str,
|
||||
|
||||
@@ -7,7 +7,8 @@ import time
|
||||
|
||||
# local imports
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.users import UserModel, 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
|
||||
from sqlalchemy import BigInteger, Boolean, Column, Index, String, Text, delete, select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -143,7 +144,8 @@ class FunctionsTable:
|
||||
functions: list[FunctionWithValvesModel],
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[FunctionWithValvesModel]:
|
||||
# Synchronize functions for a user by updating existing ones, inserting new ones, and removing those that are no longer present.
|
||||
# Synchronize functions by updating existing ones, inserting new ones,
|
||||
# and removing those that are no longer present.
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Get existing functions
|
||||
@@ -156,24 +158,15 @@ class FunctionsTable:
|
||||
|
||||
# Update or insert 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:
|
||||
await db.execute(
|
||||
update(Function)
|
||||
.filter_by(id=func.id)
|
||||
.values(
|
||||
**func.model_dump(),
|
||||
user_id=user_id,
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
await db.execute(update(Function).filter_by(id=func.id).values(**func_data))
|
||||
else:
|
||||
new_func = Function(
|
||||
**{
|
||||
**func.model_dump(),
|
||||
'user_id': user_id,
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
new_func = Function(**func_data)
|
||||
db.add(new_func)
|
||||
|
||||
# Remove functions that are no longer present
|
||||
@@ -227,7 +220,15 @@ class FunctionsTable:
|
||||
functions = result.scalars().all()
|
||||
|
||||
if include_valves:
|
||||
return [FunctionWithValvesModel.model_validate(function) for function in functions]
|
||||
return [
|
||||
FunctionWithValvesModel.model_validate(
|
||||
{
|
||||
**FunctionModel.model_validate(function).model_dump(),
|
||||
'valves': decrypt_valves(function.valves),
|
||||
}
|
||||
)
|
||||
for function in functions
|
||||
]
|
||||
else:
|
||||
return [FunctionModel.model_validate(function) for function in functions]
|
||||
|
||||
@@ -283,7 +284,7 @@ class FunctionsTable:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
function = await db.get(Function, id)
|
||||
return function.valves if function.valves else {}
|
||||
return decrypt_valves(function.valves if function else None)
|
||||
except Exception as e:
|
||||
log.exception(f'Error getting function valves by id {id}: {e}')
|
||||
return None
|
||||
@@ -300,7 +301,7 @@ class FunctionsTable:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids)))
|
||||
functions = result.all()
|
||||
return {f.id: (f.valves if f.valves else {}) for f in functions}
|
||||
return {f.id: decrypt_valves(f.valves) for f in functions}
|
||||
except Exception as e:
|
||||
log.exception(f'Error batch-fetching function valves: {e}')
|
||||
return {}
|
||||
@@ -311,7 +312,7 @@ class FunctionsTable:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
function = await db.get(Function, id)
|
||||
function.valves = valves
|
||||
function.valves = encrypt_valves(valves)
|
||||
function.updated_at = int(time.time())
|
||||
await db.commit()
|
||||
await db.refresh(function)
|
||||
@@ -355,8 +356,8 @@ class FunctionsTable:
|
||||
if 'valves' not in user_settings['functions']:
|
||||
user_settings['functions']['valves'] = {}
|
||||
|
||||
return user_settings['functions']['valves'].get(id, {})
|
||||
except Exception as e:
|
||||
return decrypt_valves(user_settings['functions']['valves'].get(id))
|
||||
except Exception:
|
||||
log.exception(f'Error getting user values by id {id} and user id {user_id}')
|
||||
return None
|
||||
|
||||
@@ -373,12 +374,12 @@ class FunctionsTable:
|
||||
if 'valves' not in user_settings['functions']:
|
||||
user_settings['functions']['valves'] = {}
|
||||
|
||||
user_settings['functions']['valves'][id] = valves
|
||||
user_settings['functions']['valves'][id] = encrypt_valves(valves)
|
||||
|
||||
# Update the user settings in the database
|
||||
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
|
||||
|
||||
return user_settings['functions']['valves'][id]
|
||||
return valves
|
||||
except Exception as e:
|
||||
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
|
||||
return None
|
||||
|
||||
@@ -4,6 +4,7 @@ import time
|
||||
import uuid
|
||||
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.models.access_grants import AccessGrantModel, AccessGrants
|
||||
from open_webui.models.files import (
|
||||
@@ -31,6 +32,7 @@ from sqlalchemy import (
|
||||
update,
|
||||
)
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import defer
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -286,6 +288,17 @@ class KnowledgeTable:
|
||||
elif view_option == 'shared':
|
||||
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(
|
||||
db=db,
|
||||
query=stmt,
|
||||
@@ -369,6 +382,7 @@ class KnowledgeTable:
|
||||
# to avoid PostgreSQL "invalid memory alloc request
|
||||
# size" on large extracted-content rows (#24670).
|
||||
content_text = File.data['content'].as_string()
|
||||
content_text = func.substr(content_text, 1, RAG_FILE_CONTENT_SEARCH_MAX_CHARS)
|
||||
search_filter = or_(
|
||||
File.filename.ilike(f'%{q}%'),
|
||||
content_text.ilike(f'%{q}%'),
|
||||
@@ -405,6 +419,7 @@ class KnowledgeTable:
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
stmt = stmt.options(defer(File.data))
|
||||
result = await db.execute(stmt)
|
||||
rows = result.all()
|
||||
|
||||
@@ -412,7 +427,13 @@ class KnowledgeTable:
|
||||
for file, user, knowledge in rows:
|
||||
items.append(
|
||||
FileUserResponse(
|
||||
**FileModel.model_validate(file).model_dump(),
|
||||
id=file.id,
|
||||
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),
|
||||
collection=(await self._to_knowledge_model(knowledge, db=db)).model_dump(),
|
||||
)
|
||||
@@ -554,6 +575,7 @@ class KnowledgeTable:
|
||||
# to avoid PostgreSQL memory allocation failures on
|
||||
# large content (#24670).
|
||||
content_text = File.data['content'].as_string()
|
||||
content_text = func.substr(content_text, 1, RAG_FILE_CONTENT_SEARCH_MAX_CHARS)
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
File.filename.ilike(f'%{query_key}%'),
|
||||
@@ -592,17 +614,23 @@ class KnowledgeTable:
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
stmt = stmt.options(defer(File.data))
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
files = []
|
||||
for file, user in items:
|
||||
files.append(
|
||||
FileUserResponse(
|
||||
**FileModel.model_validate(file).model_dump(),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
files = [
|
||||
FileUserResponse(
|
||||
id=file.id,
|
||||
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),
|
||||
)
|
||||
for file, user in items
|
||||
]
|
||||
|
||||
return KnowledgeFileListResponse(
|
||||
items=files,
|
||||
@@ -765,6 +793,25 @@ class KnowledgeTable:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
async def update_knowledge_meta_by_id(
|
||||
self, id: str, meta: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[KnowledgeModel]:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(
|
||||
update(Knowledge)
|
||||
.filter_by(id=id)
|
||||
.values(
|
||||
meta=meta,
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
return await self.get_knowledge_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
async def delete_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
|
||||
@@ -4,11 +4,11 @@ from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
from typing import Literal
|
||||
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, String, Text, delete, select
|
||||
from sqlalchemy import JSON, BigInteger, Column, String, Text, delete, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
|
||||
@@ -19,7 +19,10 @@ class Memory(Base): # user memory store
|
||||
|
||||
id = Column(String, primary_key=True, unique=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
|
||||
meta = Column(JSON, nullable=True)
|
||||
updated_at = Column(BigInteger) # epoch seconds
|
||||
created_at = Column(BigInteger) # epoch seconds
|
||||
|
||||
@@ -29,17 +32,27 @@ class MemoryModel(BaseModel):
|
||||
|
||||
id: str
|
||||
user_id: str
|
||||
type: Literal['user', 'context'] = 'context'
|
||||
path: str | None = None
|
||||
content: str
|
||||
meta: dict | None = None
|
||||
updated_at: int # timestamp in epoch
|
||||
created_at: int # timestamp in epoch
|
||||
model_config = ConfigDict(from_attributes=True) # allows ORM mapping
|
||||
|
||||
|
||||
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(
|
||||
self,
|
||||
user_id: str,
|
||||
content: str,
|
||||
memory_type: str | None = None,
|
||||
path: str | None = None,
|
||||
meta: dict | None = None,
|
||||
db: AsyncSession | None = None,
|
||||
) -> MemoryModel | None:
|
||||
"""Persist a new memory entry and return the created model."""
|
||||
@@ -48,7 +61,10 @@ class MemoriesTable:
|
||||
record = Memory(
|
||||
id=str(uuid.uuid4()),
|
||||
user_id=user_id,
|
||||
type=self.normalize_memory_type(memory_type),
|
||||
path=path,
|
||||
content=content,
|
||||
meta=meta,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
@@ -61,7 +77,11 @@ class MemoriesTable:
|
||||
self,
|
||||
id: str,
|
||||
user_id: str,
|
||||
content: str,
|
||||
content: str | None,
|
||||
memory_type: str | None = None,
|
||||
path: str | None = None,
|
||||
update_path: bool = False,
|
||||
meta: dict | None = None,
|
||||
db: AsyncSession | None = None,
|
||||
) -> MemoryModel | None:
|
||||
async with get_async_db_context(db) as db:
|
||||
@@ -70,7 +90,14 @@ class MemoriesTable:
|
||||
if not memory or memory.user_id != user_id:
|
||||
return None
|
||||
|
||||
memory.content = content
|
||||
if content is not None:
|
||||
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())
|
||||
|
||||
await db.commit()
|
||||
@@ -139,5 +166,104 @@ class MemoriesTable:
|
||||
except Exception:
|
||||
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
|
||||
|
||||
@@ -328,7 +328,8 @@ class MessageTable:
|
||||
async with get_async_db_context(db) as db:
|
||||
message = await db.get(Message, parent_id)
|
||||
|
||||
if not message:
|
||||
# Thread parent must belong to the requested channel; never disclose a foreign-channel message.
|
||||
if not message or message.channel_id != channel_id:
|
||||
return []
|
||||
|
||||
result = await db.execute(
|
||||
@@ -500,6 +501,71 @@ class MessageTable:
|
||||
|
||||
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(
|
||||
self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
|
||||
) -> bool:
|
||||
|
||||
@@ -229,10 +229,25 @@ class ModelsTable:
|
||||
)
|
||||
return models
|
||||
|
||||
async def get_base_models(self, db: AsyncSession | None = None) -> list[ModelModel]:
|
||||
@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 == None))
|
||||
result = await db.execute(select(Model).filter(Model.base_model_id.is_(None)))
|
||||
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]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
return [
|
||||
@@ -395,11 +410,14 @@ class ModelsTable:
|
||||
self,
|
||||
user_id: str,
|
||||
is_admin: bool = False,
|
||||
is_base_model: bool = False,
|
||||
db: AsyncSession | None = None,
|
||||
) -> set[str]:
|
||||
"""Extract unique tag names from model meta, querying only the meta column."""
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Model.meta).filter(Model.base_model_id != None)
|
||||
stmt = select(Model.meta).filter(
|
||||
Model.base_model_id.is_(None) if is_base_model else Model.base_model_id.is_not(None)
|
||||
)
|
||||
|
||||
if not is_admin:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
|
||||
@@ -201,5 +201,15 @@ class SharedChatsTable:
|
||||
except Exception:
|
||||
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()
|
||||
|
||||
@@ -10,6 +10,7 @@ from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import UserResponse, Users
|
||||
from open_webui.utils.valves import decrypt_valves, encrypt_valves
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import BigInteger, Column, String, Text, delete, select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -35,6 +36,7 @@ class Tool(Base): # database table definition
|
||||
class ToolMeta(BaseModel):
|
||||
description: str | None = None
|
||||
manifest: dict | None = {}
|
||||
has_user_valves: bool = False
|
||||
|
||||
|
||||
class ToolModel(BaseModel):
|
||||
@@ -231,8 +233,8 @@ class ToolsTable:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
tool = await db.get(Tool, id)
|
||||
return tool.valves if tool.valves else {}
|
||||
except Exception as e:
|
||||
return decrypt_valves(tool.valves if tool else None)
|
||||
except Exception:
|
||||
log.exception(f'Error getting tool valves by id {id}')
|
||||
return None
|
||||
|
||||
@@ -241,7 +243,9 @@ class ToolsTable:
|
||||
) -> ToolValves | None:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(update(Tool).filter_by(id=id).values(valves=valves, updated_at=int(time.time())))
|
||||
await db.execute(
|
||||
update(Tool).filter_by(id=id).values(valves=encrypt_valves(valves), updated_at=int(time.time()))
|
||||
)
|
||||
await db.commit()
|
||||
return await self.get_tool_by_id(id, db=db)
|
||||
except Exception:
|
||||
@@ -260,7 +264,7 @@ class ToolsTable:
|
||||
if 'valves' not in user_settings['tools']:
|
||||
user_settings['tools']['valves'] = {}
|
||||
|
||||
return user_settings['tools']['valves'].get(id, {})
|
||||
return decrypt_valves(user_settings['tools']['valves'].get(id))
|
||||
except Exception as e:
|
||||
log.exception(f'Error getting user values by id {id} and user_id {user_id}: {e}')
|
||||
return None
|
||||
@@ -278,12 +282,12 @@ class ToolsTable:
|
||||
if 'valves' not in user_settings['tools']:
|
||||
user_settings['tools']['valves'] = {}
|
||||
|
||||
user_settings['tools']['valves'][id] = valves
|
||||
user_settings['tools']['valves'][id] = encrypt_valves(valves)
|
||||
|
||||
# Update the user settings in the database
|
||||
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
|
||||
|
||||
return user_settings['tools']['valves'][id]
|
||||
return valves
|
||||
except Exception as e:
|
||||
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
|
||||
return None
|
||||
|
||||
@@ -279,6 +279,11 @@ class UsersTable:
|
||||
oauth: dict | None = None,
|
||||
db: AsyncSession | None = None,
|
||||
) -> UserModel | None:
|
||||
try:
|
||||
profile_image_url = validate_profile_image_url(profile_image_url)
|
||||
except ValueError:
|
||||
profile_image_url = '/user.png'
|
||||
|
||||
async with get_async_db_context(db) as session:
|
||||
user = UserModel(
|
||||
**{
|
||||
@@ -606,6 +611,11 @@ class UsersTable:
|
||||
profile_image_url: str,
|
||||
db: AsyncSession | None = None,
|
||||
) -> UserModel | None:
|
||||
try:
|
||||
profile_image_url = validate_profile_image_url(profile_image_url)
|
||||
except ValueError:
|
||||
profile_image_url = '/user.png'
|
||||
|
||||
async with get_async_db_context(db) as session:
|
||||
user = await session.get(User, id)
|
||||
if user is None:
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Optional
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
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]],
|
||||
}
|
||||
@@ -6,7 +6,7 @@ from urllib.parse import quote
|
||||
import requests
|
||||
from langchain_core.document_loaders import BaseLoader
|
||||
from langchain_core.documents import Document
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
from open_webui.utils.headers import get_custom_headers, include_user_info_headers
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -19,6 +19,8 @@ class ExternalDocumentLoader(BaseLoader):
|
||||
api_key: str,
|
||||
mime_type=None,
|
||||
user=None,
|
||||
headers=None,
|
||||
metadata=None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
self.url = url
|
||||
@@ -28,6 +30,8 @@ class ExternalDocumentLoader(BaseLoader):
|
||||
self.mime_type = mime_type
|
||||
|
||||
self.user = user
|
||||
self.headers = headers
|
||||
self.metadata = metadata
|
||||
|
||||
def load(self) -> List[Document]:
|
||||
with open(self.file_path, 'rb') as f:
|
||||
@@ -45,6 +49,8 @@ class ExternalDocumentLoader(BaseLoader):
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
headers.update(get_custom_headers(self.headers, self.user, self.metadata))
|
||||
|
||||
if self.user is not None:
|
||||
headers = include_user_info_headers(headers, self.user)
|
||||
|
||||
|
||||
@@ -17,7 +17,12 @@ from langchain_community.document_loaders import (
|
||||
YoutubeLoader,
|
||||
)
|
||||
from langchain_core.documents import Document
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, GLOBAL_LOG_LEVEL, REQUESTS_VERIFY
|
||||
from open_webui.env import (
|
||||
AIOHTTP_CLIENT_SESSION_SSL,
|
||||
GLOBAL_LOG_LEVEL,
|
||||
MINERU_MAX_MARKDOWN_BYTES,
|
||||
REQUESTS_VERIFY,
|
||||
)
|
||||
from open_webui.retrieval.loaders.datalab_marker import DatalabMarkerLoader
|
||||
from open_webui.retrieval.loaders.external_document import ExternalDocumentLoader
|
||||
from open_webui.retrieval.loaders.mineru import MinerULoader
|
||||
@@ -183,6 +188,7 @@ class DoclingLoader:
|
||||
self.params = params or {}
|
||||
|
||||
def load(self) -> list[Document]:
|
||||
page_break_marker = '\f'
|
||||
with open(self.file_path, 'rb') as f:
|
||||
headers = {}
|
||||
if self.api_key:
|
||||
@@ -199,6 +205,7 @@ class DoclingLoader:
|
||||
},
|
||||
data={
|
||||
'image_export_mode': 'placeholder',
|
||||
'md_page_break_placeholder': page_break_marker,
|
||||
**self.params,
|
||||
},
|
||||
headers=headers,
|
||||
@@ -207,9 +214,19 @@ class DoclingLoader:
|
||||
if r.ok:
|
||||
result = r.json()
|
||||
document_data = result.get('document', {})
|
||||
text = document_data.get('md_content', '<No text content found>')
|
||||
md_content = document_data.get('md_content', '')
|
||||
text = md_content or '<No text content found>'
|
||||
|
||||
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)
|
||||
return [Document(page_content=text, metadata=metadata)]
|
||||
@@ -229,6 +246,7 @@ class Loader:
|
||||
def __init__(self, engine: str = '', **kwargs):
|
||||
self.engine = engine
|
||||
self.user = kwargs.get('user', None)
|
||||
self.metadata = kwargs.get('metadata', {})
|
||||
self.kwargs = kwargs
|
||||
|
||||
def load(self, filename: str, file_content_type: str, file_path: str) -> list[Document]:
|
||||
@@ -404,6 +422,12 @@ class Loader:
|
||||
api_key=self.kwargs.get('EXTERNAL_DOCUMENT_LOADER_API_KEY'),
|
||||
mime_type=file_content_type,
|
||||
user=self.user,
|
||||
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'):
|
||||
if self._is_text_file(file_ext, file_content_type):
|
||||
@@ -511,7 +535,6 @@ class Loader:
|
||||
mineru_timeout = int(mineru_timeout)
|
||||
except ValueError:
|
||||
mineru_timeout = 300
|
||||
|
||||
loader = MinerULoader(
|
||||
file_path=file_path,
|
||||
api_mode=self.kwargs.get('MINERU_API_MODE', 'local'),
|
||||
@@ -519,6 +542,7 @@ class Loader:
|
||||
api_key=self.kwargs.get('MINERU_API_KEY', ''),
|
||||
params=self.kwargs.get('MINERU_PARAMS', {}),
|
||||
timeout=mineru_timeout,
|
||||
max_markdown_bytes=MINERU_MAX_MARKDOWN_BYTES,
|
||||
)
|
||||
elif (
|
||||
self.engine == 'mistral_ocr'
|
||||
@@ -529,6 +553,7 @@ class Loader:
|
||||
base_url=self.kwargs.get('MISTRAL_OCR_API_BASE_URL'),
|
||||
api_key=self.kwargs.get('MISTRAL_OCR_API_KEY'),
|
||||
file_path=file_path,
|
||||
use_base64=self.kwargs.get('MISTRAL_OCR_USE_BASE64', False),
|
||||
)
|
||||
elif self.engine == 'paddleocr_vl' and self.kwargs.get('PADDLEOCR_VL_TOKEN') != '':
|
||||
loader = PaddleOCRVLLoader(
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
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__)
|
||||
|
||||
DEFAULT_MICROSOFT_WEB_IQ_API_BASE_URL = 'https://api.microsoft.ai/v3'
|
||||
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 or DEFAULT_MICROSOFT_WEB_IQ_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,20 +28,22 @@ class MinerULoader:
|
||||
api_key: str = '',
|
||||
params: dict = None,
|
||||
timeout: Optional[int] = 300,
|
||||
max_markdown_bytes: Optional[int] = None,
|
||||
):
|
||||
self.file_path = file_path
|
||||
self.api_mode = api_mode.lower()
|
||||
self.api_url = api_url.rstrip('/')
|
||||
self.api_key = api_key
|
||||
self.timeout = timeout
|
||||
self.max_markdown_bytes = max_markdown_bytes
|
||||
|
||||
# Parse params dict with defaults
|
||||
self.params = params or {}
|
||||
self.enable_ocr = params.get('enable_ocr', False)
|
||||
self.enable_formula = params.get('enable_formula', True)
|
||||
self.enable_table = params.get('enable_table', True)
|
||||
self.language = params.get('language', 'en')
|
||||
self.model_version = params.get('model_version', 'pipeline')
|
||||
self.enable_ocr = self.params.get('enable_ocr', False)
|
||||
self.enable_formula = self.params.get('enable_formula', True)
|
||||
self.enable_table = self.params.get('enable_table', True)
|
||||
self.language = self.params.get('language', 'en')
|
||||
self.model_version = self.params.get('model_version', 'pipeline')
|
||||
|
||||
self.page_ranges = self.params.pop('page_ranges', '')
|
||||
|
||||
@@ -435,67 +437,77 @@ class MinerULoader:
|
||||
detail=f'Error downloading results: {str(e)}',
|
||||
)
|
||||
|
||||
# Save ZIP to temporary file and extract
|
||||
# Save ZIP to temporary file before reading.
|
||||
tmp_zip_path = None
|
||||
markdown_content = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix='.zip') as tmp_zip:
|
||||
tmp_zip.write(response.content)
|
||||
tmp_zip_path = tmp_zip.name
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
# Extract ZIP
|
||||
with zipfile.ZipFile(tmp_zip_path, 'r') as zip_ref:
|
||||
zip_ref.extractall(tmp_dir)
|
||||
with zipfile.ZipFile(tmp_zip_path, 'r') as zip_ref:
|
||||
members = zip_ref.infolist()
|
||||
all_files = [member.filename for member in members]
|
||||
md_members = [member for member in members if member.filename.endswith('.md')]
|
||||
read_errors = []
|
||||
|
||||
# Find markdown file - search recursively for any .md file
|
||||
markdown_content = None
|
||||
found_md_path = None
|
||||
|
||||
# First, list all files in the ZIP for debugging
|
||||
all_files = []
|
||||
for root, dirs, files in os.walk(tmp_dir):
|
||||
for file in files:
|
||||
full_path = os.path.join(root, file)
|
||||
all_files.append(full_path)
|
||||
# Look for any .md file
|
||||
if file.endswith('.md'):
|
||||
found_md_path = full_path
|
||||
log.info(f'Found markdown file at: {full_path}')
|
||||
try:
|
||||
with open(full_path, 'r', encoding='utf-8') as f:
|
||||
markdown_content = f.read()
|
||||
if markdown_content: # Use the first non-empty markdown file
|
||||
break
|
||||
except Exception as e:
|
||||
log.warning(f'Failed to read {full_path}: {e}')
|
||||
for member in md_members:
|
||||
log.info(f'Found markdown file in ZIP: {member.filename}')
|
||||
try:
|
||||
with zip_ref.open(member, 'r') as f:
|
||||
if self.max_markdown_bytes is None:
|
||||
content = f.read()
|
||||
else:
|
||||
content = f.read(self.max_markdown_bytes + 1)
|
||||
if len(content) > self.max_markdown_bytes:
|
||||
raise HTTPException(
|
||||
status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f'Markdown file in results ZIP is too large: {member.filename}',
|
||||
)
|
||||
markdown_content = content.decode('utf-8')
|
||||
except UnicodeDecodeError as e:
|
||||
read_errors.append(f'{member.filename}: {e}')
|
||||
log.warning(f'Failed to decode {member.filename}: {e}')
|
||||
continue
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
read_errors.append(f'{member.filename}: {e}')
|
||||
log.warning(f'Failed to read {member.filename}: {e}')
|
||||
continue
|
||||
if markdown_content:
|
||||
break
|
||||
|
||||
if markdown_content is None:
|
||||
log.error(f'Available files in ZIP: {all_files}')
|
||||
# Try to provide more helpful error message
|
||||
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}"
|
||||
if read_errors:
|
||||
error_msg = f"Found .md files but couldn't read them: {read_errors}"
|
||||
else:
|
||||
error_msg = f'No .md files found in ZIP. Available files: {all_files}'
|
||||
raise HTTPException(
|
||||
status.HTTP_502_BAD_GATEWAY,
|
||||
detail=error_msg,
|
||||
)
|
||||
|
||||
# Clean up temporary ZIP file
|
||||
os.unlink(tmp_zip_path)
|
||||
|
||||
except zipfile.BadZipFile as e:
|
||||
raise HTTPException(
|
||||
status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f'Invalid ZIP file received: {e}',
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
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:
|
||||
raise HTTPException(
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import base64
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
@@ -37,6 +38,7 @@ class MistralLoader:
|
||||
timeout: int = 300, # 5 minutes default
|
||||
max_retries: int = 3,
|
||||
enable_debug_logging: bool = False,
|
||||
use_base64: bool = False,
|
||||
):
|
||||
"""
|
||||
Initializes the loader with enhanced features.
|
||||
@@ -47,6 +49,7 @@ class MistralLoader:
|
||||
timeout: Request timeout in seconds.
|
||||
max_retries: Maximum number of retry attempts.
|
||||
enable_debug_logging: Enable detailed debug logs.
|
||||
use_base64: Send the document as a data URL instead of uploading it first.
|
||||
"""
|
||||
if not api_key:
|
||||
raise ValueError('API key cannot be empty.')
|
||||
@@ -59,6 +62,7 @@ class MistralLoader:
|
||||
self.timeout = timeout
|
||||
self.max_retries = max_retries
|
||||
self.debug = enable_debug_logging
|
||||
self.use_base64 = use_base64
|
||||
|
||||
# PERFORMANCE OPTIMIZATION: Differentiated timeouts for different operations
|
||||
# This prevents long-running OCR operations from affecting quick operations
|
||||
@@ -261,33 +265,32 @@ class MistralLoader:
|
||||
url = f'{self.base_url}/files'
|
||||
|
||||
async def upload_request():
|
||||
# Create multipart writer for streaming upload
|
||||
writer = aiohttp.MultipartWriter('form-data')
|
||||
# Open inside the request so the handle stays valid for the whole
|
||||
# streamed POST and is closed right after.
|
||||
with open(self.file_path, 'rb') as f:
|
||||
writer = aiohttp.MultipartWriter('form-data')
|
||||
|
||||
# Add purpose field
|
||||
purpose_part = writer.append('ocr')
|
||||
purpose_part.set_content_disposition('form-data', name='purpose')
|
||||
# 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)
|
||||
# Stream the file. aiohttp builds a payload from the file object;
|
||||
# the previous aiohttp.streams.FilePayload was removed upstream
|
||||
# (payloads live in aiohttp.payload and there is no FilePayload),
|
||||
# so this path raised AttributeError on every async OCR upload.
|
||||
file_part = writer.append(f, {'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)')
|
||||
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)
|
||||
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)
|
||||
|
||||
@@ -417,6 +420,11 @@ class MistralLoader:
|
||||
|
||||
return await self._retry_request_async(ocr_request)
|
||||
|
||||
def _get_file_data_url(self) -> str:
|
||||
with open(self.file_path, 'rb') as f:
|
||||
encoded_file = base64.b64encode(f.read()).decode('utf-8')
|
||||
return f'data:application/pdf;base64,{encoded_file}'
|
||||
|
||||
def _delete_file(self, file_id: str) -> None:
|
||||
"""Deletes the file from Mistral storage (sync version)."""
|
||||
log.info(f'Deleting uploaded file ID: {file_id}')
|
||||
@@ -566,6 +574,12 @@ class MistralLoader:
|
||||
start_time = time.time()
|
||||
|
||||
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(f'Sync OCR workflow completed in {total_time:.2f}s, produced {len(documents)} documents')
|
||||
return documents
|
||||
|
||||
# 1. Upload file
|
||||
file_id = self._upload_file()
|
||||
|
||||
@@ -617,6 +631,13 @@ class MistralLoader:
|
||||
|
||||
try:
|
||||
async with self._get_session() as session:
|
||||
if self.use_base64:
|
||||
ocr_response = await self._process_ocr_async(session, self._get_file_data_url())
|
||||
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
|
||||
|
||||
# 1. Upload file with streaming
|
||||
file_id = await self._upload_file_async(session)
|
||||
|
||||
|
||||
@@ -39,11 +39,13 @@ from open_webui.models.chats import Chats
|
||||
from open_webui.models.files import Files
|
||||
from open_webui.models.knowledge import Knowledges
|
||||
from open_webui.models.notes import Notes
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.users import UserModel
|
||||
from open_webui.retrieval.loaders.youtube import YoutubeLoader
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.external import retrieve_external_knowledge
|
||||
from open_webui.retrieval.vector.factory import VECTOR_DB_CLIENT
|
||||
from open_webui.retrieval.vector.main import GetResult
|
||||
from open_webui.retrieval.vector.main import GetResult, SearchResult
|
||||
from open_webui.retrieval.web.utils import get_web_loader
|
||||
from open_webui.utils.access_control.files import has_access_to_file
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
@@ -63,65 +65,84 @@ def is_youtube_url(url: str) -> bool:
|
||||
return re.match(youtube_regex, url) is not None
|
||||
|
||||
|
||||
def get_loader(request, url: str):
|
||||
LOADER_CONFIG_KEYS = {
|
||||
'youtube_language': 'rag.youtube_loader_language',
|
||||
'youtube_proxy_url': 'rag.youtube_loader_proxy_url',
|
||||
'web_loader_ssl_verification': 'web.loader.ssl_verification',
|
||||
'web_loader_concurrent_requests': 'web.loader.concurrent_requests',
|
||||
'web_search_trust_env': 'web.search.trust_env',
|
||||
'CONTENT_EXTRACTION_ENGINE': 'rag.content_extraction_engine',
|
||||
'DATALAB_MARKER_API_KEY': 'rag.datalab_marker_api_key',
|
||||
'DATALAB_MARKER_API_BASE_URL': 'rag.datalab_marker_api_base_url',
|
||||
'DATALAB_MARKER_ADDITIONAL_CONFIG': 'rag.datalab_marker_additional_config',
|
||||
'DATALAB_MARKER_SKIP_CACHE': 'rag.datalab_marker_skip_cache',
|
||||
'DATALAB_MARKER_FORCE_OCR': 'rag.datalab_marker_force_ocr',
|
||||
'DATALAB_MARKER_PAGINATE': 'rag.datalab_marker_paginate',
|
||||
'DATALAB_MARKER_STRIP_EXISTING_OCR': 'rag.datalab_marker_strip_existing_ocr',
|
||||
'DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION': 'rag.datalab_marker_disable_image_extraction',
|
||||
'DATALAB_MARKER_FORMAT_LINES': 'rag.datalab_marker_format_lines',
|
||||
'DATALAB_MARKER_USE_LLM': 'rag.datalab_marker_use_llm',
|
||||
'DATALAB_MARKER_OUTPUT_FORMAT': 'rag.datalab_marker_output_format',
|
||||
'EXTERNAL_DOCUMENT_LOADER_URL': 'rag.external_document_loader_url',
|
||||
'EXTERNAL_DOCUMENT_LOADER_API_KEY': 'rag.external_document_loader_api_key',
|
||||
'EXTERNAL_DOCUMENT_LOADER_HEADERS': 'rag.external_document_loader_headers',
|
||||
'TIKA_SERVER_URL': 'rag.tika_server_url',
|
||||
'DOCLING_SERVER_URL': 'rag.docling_server_url',
|
||||
'DOCLING_API_KEY': 'rag.docling_api_key',
|
||||
'DOCLING_PARAMS': 'rag.docling_params',
|
||||
'PDF_EXTRACT_IMAGES': 'rag.pdf_extract_images',
|
||||
'PDF_LOADER_MODE': 'rag.pdf_loader_mode',
|
||||
'DOCUMENT_INTELLIGENCE_ENDPOINT': 'rag.document_intelligence_endpoint',
|
||||
'DOCUMENT_INTELLIGENCE_KEY': 'rag.document_intelligence_key',
|
||||
'DOCUMENT_INTELLIGENCE_MODEL': 'rag.document_intelligence_model',
|
||||
'MISTRAL_OCR_API_BASE_URL': 'rag.mistral_ocr_api_base_url',
|
||||
'MISTRAL_OCR_API_KEY': 'rag.mistral_ocr_api_key',
|
||||
'MISTRAL_OCR_USE_BASE64': 'rag.mistral_ocr_use_base64',
|
||||
'PADDLEOCR_VL_BASE_URL': 'rag.paddleocr_vl_base_url',
|
||||
'PADDLEOCR_VL_TOKEN': 'rag.paddleocr_vl_token',
|
||||
'MINERU_API_MODE': 'rag.mineru_api_mode',
|
||||
'MINERU_API_URL': 'rag.mineru_api_url',
|
||||
'MINERU_API_KEY': 'rag.mineru_api_key',
|
||||
'MINERU_API_TIMEOUT': 'rag.mineru_api_timeout',
|
||||
'MINERU_PARAMS': 'rag.mineru_params',
|
||||
'MINERU_FILE_EXTENSIONS': 'rag.mineru_file_extensions',
|
||||
}
|
||||
|
||||
|
||||
async def get_loader_config():
|
||||
values = await Config.get_many(*LOADER_CONFIG_KEYS.values())
|
||||
return {name: values.get(key) for name, key in LOADER_CONFIG_KEYS.items()}
|
||||
|
||||
|
||||
def get_loader(request, url: str, config: dict):
|
||||
if is_youtube_url(url):
|
||||
return YoutubeLoader(
|
||||
url,
|
||||
language=request.app.state.config.YOUTUBE_LOADER_LANGUAGE,
|
||||
proxy_url=request.app.state.config.YOUTUBE_LOADER_PROXY_URL,
|
||||
language=config.get('youtube_language'),
|
||||
proxy_url=config.get('youtube_proxy_url'),
|
||||
)
|
||||
else:
|
||||
return get_web_loader(
|
||||
url,
|
||||
verify_ssl=request.app.state.config.ENABLE_WEB_LOADER_SSL_VERIFICATION,
|
||||
requests_per_second=request.app.state.config.WEB_LOADER_CONCURRENT_REQUESTS,
|
||||
trust_env=request.app.state.config.WEB_SEARCH_TRUST_ENV,
|
||||
)
|
||||
|
||||
|
||||
def build_loader_from_config(request):
|
||||
"""Build a Loader instance with the admin's configured extraction engine settings."""
|
||||
from open_webui.retrieval.loaders.main import Loader
|
||||
|
||||
config = request.app.state.config
|
||||
return Loader(
|
||||
engine=config.CONTENT_EXTRACTION_ENGINE,
|
||||
DATALAB_MARKER_API_KEY=config.DATALAB_MARKER_API_KEY,
|
||||
DATALAB_MARKER_API_BASE_URL=config.DATALAB_MARKER_API_BASE_URL,
|
||||
DATALAB_MARKER_ADDITIONAL_CONFIG=config.DATALAB_MARKER_ADDITIONAL_CONFIG,
|
||||
DATALAB_MARKER_SKIP_CACHE=config.DATALAB_MARKER_SKIP_CACHE,
|
||||
DATALAB_MARKER_FORCE_OCR=config.DATALAB_MARKER_FORCE_OCR,
|
||||
DATALAB_MARKER_PAGINATE=config.DATALAB_MARKER_PAGINATE,
|
||||
DATALAB_MARKER_STRIP_EXISTING_OCR=config.DATALAB_MARKER_STRIP_EXISTING_OCR,
|
||||
DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION=config.DATALAB_MARKER_DISABLE_IMAGE_EXTRACTION,
|
||||
DATALAB_MARKER_FORMAT_LINES=config.DATALAB_MARKER_FORMAT_LINES,
|
||||
DATALAB_MARKER_USE_LLM=config.DATALAB_MARKER_USE_LLM,
|
||||
DATALAB_MARKER_OUTPUT_FORMAT=config.DATALAB_MARKER_OUTPUT_FORMAT,
|
||||
EXTERNAL_DOCUMENT_LOADER_URL=config.EXTERNAL_DOCUMENT_LOADER_URL,
|
||||
EXTERNAL_DOCUMENT_LOADER_API_KEY=config.EXTERNAL_DOCUMENT_LOADER_API_KEY,
|
||||
TIKA_SERVER_URL=config.TIKA_SERVER_URL,
|
||||
DOCLING_SERVER_URL=config.DOCLING_SERVER_URL,
|
||||
DOCLING_API_KEY=config.DOCLING_API_KEY,
|
||||
DOCLING_PARAMS=config.DOCLING_PARAMS,
|
||||
PDF_EXTRACT_IMAGES=config.PDF_EXTRACT_IMAGES,
|
||||
PDF_LOADER_MODE=config.PDF_LOADER_MODE,
|
||||
DOCUMENT_INTELLIGENCE_ENDPOINT=config.DOCUMENT_INTELLIGENCE_ENDPOINT,
|
||||
DOCUMENT_INTELLIGENCE_KEY=config.DOCUMENT_INTELLIGENCE_KEY,
|
||||
DOCUMENT_INTELLIGENCE_MODEL=config.DOCUMENT_INTELLIGENCE_MODEL,
|
||||
MISTRAL_OCR_API_BASE_URL=config.MISTRAL_OCR_API_BASE_URL,
|
||||
MISTRAL_OCR_API_KEY=config.MISTRAL_OCR_API_KEY,
|
||||
PADDLEOCR_VL_BASE_URL=config.PADDLEOCR_VL_BASE_URL,
|
||||
PADDLEOCR_VL_TOKEN=config.PADDLEOCR_VL_TOKEN,
|
||||
MINERU_API_MODE=config.MINERU_API_MODE,
|
||||
MINERU_API_URL=config.MINERU_API_URL,
|
||||
MINERU_API_KEY=config.MINERU_API_KEY,
|
||||
MINERU_API_TIMEOUT=config.MINERU_API_TIMEOUT,
|
||||
MINERU_PARAMS=config.MINERU_PARAMS,
|
||||
MINERU_FILE_EXTENSIONS=config.MINERU_FILE_EXTENSIONS,
|
||||
return get_web_loader(
|
||||
url,
|
||||
verify_ssl=config.get('web_loader_ssl_verification'),
|
||||
requests_per_second=config.get('web_loader_concurrent_requests'),
|
||||
trust_env=config.get('web_search_trust_env'),
|
||||
)
|
||||
|
||||
|
||||
def _extract_text_from_binary_response(request, response: requests.Response, url: str) -> tuple[str, list]:
|
||||
def build_loader_from_config(request, config: dict):
|
||||
"""Build a Loader instance with the admin's configured extraction engine settings."""
|
||||
from open_webui.retrieval.loaders.main import Loader
|
||||
|
||||
loader_config = {key: config.get(key) for key in LOADER_CONFIG_KEYS if key.isupper()}
|
||||
return Loader(
|
||||
engine=loader_config['CONTENT_EXTRACTION_ENGINE'],
|
||||
**{key: value for key, value in loader_config.items() if key != 'CONTENT_EXTRACTION_ENGINE'},
|
||||
)
|
||||
|
||||
|
||||
def _extract_text_from_binary_response(
|
||||
request, response: requests.Response, url: str, loader_config: dict
|
||||
) -> tuple[str, list]:
|
||||
"""Download response body to a temp file and extract text using the Loader pipeline."""
|
||||
import mimetypes
|
||||
import tempfile
|
||||
@@ -150,7 +171,7 @@ def _extract_text_from_binary_response(request, response: requests.Response, url
|
||||
tmp_path = tmp.name
|
||||
|
||||
try:
|
||||
loader = build_loader_from_config(request)
|
||||
loader = build_loader_from_config(request, loader_config)
|
||||
docs = loader.load(filename, content_type, tmp_path)
|
||||
for doc in docs:
|
||||
doc.metadata['source'] = url
|
||||
@@ -170,8 +191,17 @@ def _is_text_content_type(content_type: str) -> bool:
|
||||
return not ct # empty / missing → assume HTML
|
||||
|
||||
|
||||
def get_content_from_url(request, url: str) -> str:
|
||||
from open_webui.retrieval.web.utils import validate_url
|
||||
async def get_content_from_url(request, url: str) -> str:
|
||||
loader_config = await get_loader_config()
|
||||
|
||||
# The rest of this function performs synchronous, blocking work: an SSRF-guarded
|
||||
# `requests` probe and a synchronous document loader (`loader.load()`). Run it in a
|
||||
# worker thread so the event loop stays free while waiting on network/parsing.
|
||||
return await asyncio.to_thread(_get_content_from_url_sync, request, url, loader_config)
|
||||
|
||||
|
||||
def _get_content_from_url_sync(request, url: str, loader_config):
|
||||
from open_webui.retrieval.web.utils import validate_url, _SSRFSafeAdapter
|
||||
|
||||
# Validate URL before making any request (blocks private IPs, non-HTTP, filter list)
|
||||
validate_url(url)
|
||||
@@ -183,7 +213,7 @@ def get_content_from_url(request, url: str) -> str:
|
||||
# when allow_redirects=False, causing the binary-content path to run
|
||||
# and produce empty docs → HTTP 400.
|
||||
if is_youtube_url(url):
|
||||
loader = get_loader(request, url)
|
||||
loader = get_loader(request, url, loader_config)
|
||||
docs = loader.load()
|
||||
content = ' '.join([doc.page_content for doc in docs])
|
||||
return content, docs
|
||||
@@ -194,7 +224,11 @@ def get_content_from_url(request, url: str) -> str:
|
||||
# re-validation would let an attacker reach private IPs (RFC1918, loopback,
|
||||
# cloud-metadata 169.254.169.254) via a public host that redirects internally.
|
||||
try:
|
||||
response = requests.get(url, stream=True, timeout=30, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS)
|
||||
# Probe through the connect-time SSRF guard; bare requests.get re-resolves (DNS-rebinding gap).
|
||||
session = requests.Session()
|
||||
session.mount('http://', _SSRFSafeAdapter())
|
||||
session.mount('https://', _SSRFSafeAdapter())
|
||||
response = session.get(url, stream=True, timeout=30, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS)
|
||||
response.raise_for_status()
|
||||
content_type = response.headers.get('Content-Type', '')
|
||||
except Exception:
|
||||
@@ -205,14 +239,14 @@ def get_content_from_url(request, url: str) -> str:
|
||||
if response is None or _is_text_content_type(content_type):
|
||||
if response is not None:
|
||||
response.close()
|
||||
loader = get_loader(request, url)
|
||||
loader = get_loader(request, url, loader_config)
|
||||
docs = loader.load()
|
||||
content = ' '.join([doc.page_content for doc in docs])
|
||||
return content, docs
|
||||
|
||||
# Binary content (PDF, DOCX, XLSX, PPTX, etc.) — download and extract
|
||||
try:
|
||||
return _extract_text_from_binary_response(request, response, url)
|
||||
return _extract_text_from_binary_response(request, response, url, loader_config)
|
||||
finally:
|
||||
response.close()
|
||||
|
||||
@@ -255,21 +289,7 @@ class VectorSearchRetriever(BaseRetriever):
|
||||
limit=self.top_k,
|
||||
)
|
||||
|
||||
ids = result.ids[0]
|
||||
metadatas = result.metadatas[0]
|
||||
documents = result.documents[0]
|
||||
|
||||
results = []
|
||||
for idx in range(len(ids)):
|
||||
metadata = metadatas[idx]
|
||||
metadata[CHUNK_HASH_KEY] = _content_hash(documents[idx])
|
||||
results.append(
|
||||
Document(
|
||||
metadata=metadata,
|
||||
page_content=documents[idx],
|
||||
)
|
||||
)
|
||||
return results
|
||||
return _search_result_to_documents(result)
|
||||
|
||||
|
||||
def query_doc(collection_name: str, query_embedding: list[float], k: int, user: UserModel = None):
|
||||
@@ -338,9 +358,96 @@ def get_enriched_texts(collection_result: GetResult) -> list[str]:
|
||||
return enriched_texts
|
||||
|
||||
|
||||
def _search_result_to_documents(result: SearchResult | None) -> list[Document]:
|
||||
ids = result.ids[0] if result and result.ids else []
|
||||
metadatas = result.metadatas[0] if result and result.metadatas else []
|
||||
documents = result.documents[0] if result and result.documents else []
|
||||
distances = result.distances[0] if result and result.distances else []
|
||||
|
||||
docs = []
|
||||
for idx in range(len(ids)):
|
||||
document = documents[idx]
|
||||
metadata = dict(metadatas[idx] or {})
|
||||
metadata[CHUNK_HASH_KEY] = _content_hash(document)
|
||||
if idx < len(distances):
|
||||
metadata.setdefault('score', distances[idx])
|
||||
docs.append(Document(metadata=metadata, page_content=document))
|
||||
return docs
|
||||
|
||||
|
||||
def _supports_native_hybrid_search() -> bool:
|
||||
supports_hybrid_search = getattr(ASYNC_VECTOR_DB_CLIENT, 'supports_hybrid_search', None)
|
||||
if supports_hybrid_search is not None:
|
||||
return bool(supports_hybrid_search)
|
||||
return callable(getattr(ASYNC_VECTOR_DB_CLIENT, 'hybrid_search', None))
|
||||
|
||||
|
||||
async def query_doc_with_native_hybrid_search(
|
||||
collection_name: str,
|
||||
query: str,
|
||||
embedding_function,
|
||||
k: int,
|
||||
reranking_function,
|
||||
k_reranker: int,
|
||||
r: float,
|
||||
hybrid_bm25_weight: float,
|
||||
) -> Optional[dict]:
|
||||
try:
|
||||
if not _supports_native_hybrid_search():
|
||||
return None
|
||||
|
||||
query_vectors = []
|
||||
if hybrid_bm25_weight < 1:
|
||||
query_vectors = [await embedding_function(query, RAG_EMBEDDING_QUERY_PREFIX)]
|
||||
|
||||
result = await ASYNC_VECTOR_DB_CLIENT.hybrid_search(
|
||||
collection_name=collection_name,
|
||||
query=query,
|
||||
vectors=query_vectors,
|
||||
limit=k,
|
||||
hybrid_bm25_weight=hybrid_bm25_weight,
|
||||
)
|
||||
if result is None:
|
||||
return None
|
||||
|
||||
documents = _search_result_to_documents(result)
|
||||
if not documents:
|
||||
return {'distances': [[]], 'documents': [[]], 'metadatas': [[]]}
|
||||
|
||||
compressor = RerankCompressor(
|
||||
embedding_function=embedding_function,
|
||||
top_n=k_reranker,
|
||||
reranking_function=reranking_function,
|
||||
r_score=r,
|
||||
)
|
||||
compressed = await compressor.acompress_documents(documents, query)
|
||||
|
||||
distances = [d.metadata.get('score') for d in compressed]
|
||||
documents = [d.page_content for d in compressed]
|
||||
metadatas = [d.metadata for d in compressed]
|
||||
|
||||
if k < k_reranker:
|
||||
sorted_items = sorted(zip(distances, documents, metadatas), key=lambda x: x[0], reverse=True)
|
||||
sorted_items = sorted_items[:k]
|
||||
|
||||
if sorted_items:
|
||||
distances, documents, metadatas = map(list, zip(*sorted_items))
|
||||
else:
|
||||
distances, documents, metadatas = [], [], []
|
||||
|
||||
return {
|
||||
'distances': [distances],
|
||||
'documents': [documents],
|
||||
'metadatas': [metadatas],
|
||||
}
|
||||
except Exception as e:
|
||||
log.debug(f'Native hybrid search failed for {collection_name}, falling back to legacy hybrid search: {e}')
|
||||
return None
|
||||
|
||||
|
||||
async def query_doc_with_hybrid_search(
|
||||
collection_name: str,
|
||||
collection_result: GetResult,
|
||||
collection_result: Optional[GetResult],
|
||||
query: str,
|
||||
embedding_function,
|
||||
k: int,
|
||||
@@ -349,8 +456,26 @@ async def query_doc_with_hybrid_search(
|
||||
r: float,
|
||||
hybrid_bm25_weight: float,
|
||||
enable_enriched_texts: bool = False,
|
||||
native_hybrid_search: bool = True,
|
||||
) -> dict:
|
||||
try:
|
||||
if native_hybrid_search and not enable_enriched_texts:
|
||||
native_result = await query_doc_with_native_hybrid_search(
|
||||
collection_name=collection_name,
|
||||
query=query,
|
||||
embedding_function=embedding_function,
|
||||
k=k,
|
||||
reranking_function=reranking_function,
|
||||
k_reranker=k_reranker,
|
||||
r=r,
|
||||
hybrid_bm25_weight=hybrid_bm25_weight,
|
||||
)
|
||||
if native_result is not None:
|
||||
return native_result
|
||||
|
||||
if collection_result is None:
|
||||
collection_result = await ASYNC_VECTOR_DB_CLIENT.get(collection_name=collection_name)
|
||||
|
||||
# First check if collection_result has the required attributes
|
||||
if (
|
||||
not collection_result
|
||||
@@ -539,8 +664,15 @@ async def query_collection(
|
||||
embedding_function,
|
||||
k: int,
|
||||
) -> dict:
|
||||
config = await Config.get_many(
|
||||
'rag.enable_hybrid_search',
|
||||
'rag.top_k_reranker',
|
||||
'rag.relevance_threshold',
|
||||
'rag.hybrid_bm25_weight',
|
||||
'rag.enable_hybrid_search_enriched_texts',
|
||||
)
|
||||
# When request is provided, try hybrid search + reranking if enabled
|
||||
if request and request.app.state.config.ENABLE_RAG_HYBRID_SEARCH:
|
||||
if request and config.get('rag.enable_hybrid_search'):
|
||||
try:
|
||||
reranking_function = (
|
||||
(lambda query, documents: request.app.state.RERANKING_FUNCTION(query, documents))
|
||||
@@ -553,10 +685,10 @@ async def query_collection(
|
||||
embedding_function=embedding_function,
|
||||
k=k,
|
||||
reranking_function=reranking_function,
|
||||
k_reranker=request.app.state.config.TOP_K_RERANKER,
|
||||
r=request.app.state.config.RELEVANCE_THRESHOLD,
|
||||
hybrid_bm25_weight=request.app.state.config.HYBRID_BM25_WEIGHT,
|
||||
enable_enriched_texts=request.app.state.config.ENABLE_RAG_HYBRID_SEARCH_ENRICHED_TEXTS,
|
||||
k_reranker=config.get('rag.top_k_reranker'),
|
||||
r=config.get('rag.relevance_threshold'),
|
||||
hybrid_bm25_weight=config.get('rag.hybrid_bm25_weight'),
|
||||
enable_enriched_texts=config.get('rag.enable_hybrid_search_enriched_texts'),
|
||||
)
|
||||
except Exception as e:
|
||||
log.debug(f'Hybrid search failed, falling back to vector search: {e}')
|
||||
@@ -623,6 +755,28 @@ async def query_collection_with_hybrid_search(
|
||||
) -> dict:
|
||||
results = []
|
||||
error = False
|
||||
|
||||
if not enable_enriched_texts:
|
||||
|
||||
async def process_native_query(collection_name, query):
|
||||
result = await query_doc_with_native_hybrid_search(
|
||||
collection_name=collection_name,
|
||||
query=query,
|
||||
embedding_function=embedding_function,
|
||||
k=k,
|
||||
reranking_function=reranking_function,
|
||||
k_reranker=k_reranker,
|
||||
r=r,
|
||||
hybrid_bm25_weight=hybrid_bm25_weight,
|
||||
)
|
||||
return result
|
||||
|
||||
native_task_results = await asyncio.gather(
|
||||
*[process_native_query(collection_name, query) for collection_name in collection_names for query in queries]
|
||||
)
|
||||
if native_task_results and all(result is not None for result in native_task_results):
|
||||
return merge_and_sort_query_results(native_task_results, k=k)
|
||||
|
||||
# Fetch every collection's contents once up front so the
|
||||
# per-query/per-document loop below can reuse them. Each fetch
|
||||
# offloads to a worker thread, so run them concurrently with
|
||||
@@ -657,6 +811,7 @@ async def query_collection_with_hybrid_search(
|
||||
r=r,
|
||||
hybrid_bm25_weight=hybrid_bm25_weight,
|
||||
enable_enriched_texts=enable_enriched_texts,
|
||||
native_hybrid_search=False,
|
||||
)
|
||||
return result, None
|
||||
except Exception as e:
|
||||
@@ -927,15 +1082,15 @@ def get_embedding_function(
|
||||
concurrent_requests=0,
|
||||
) -> Awaitable:
|
||||
if embedding_engine == '':
|
||||
if embedding_function is None:
|
||||
raise ValueError(
|
||||
'No embedding model is loaded. Set RAG_EMBEDDING_MODEL to a valid '
|
||||
'SentenceTransformer model name, or configure an external '
|
||||
'RAG_EMBEDDING_ENGINE (ollama, openai, azure_openai).'
|
||||
)
|
||||
|
||||
# Sentence transformers: CPU-bound sync operation
|
||||
async def async_embedding_function(query, prefix=None, user=None):
|
||||
# Deferred so a missing local model degrades RAG instead of crashing boot.
|
||||
if embedding_function is None:
|
||||
raise ValueError(
|
||||
'No embedding model is loaded. Set RAG_EMBEDDING_MODEL to a valid '
|
||||
'SentenceTransformer model name, or configure an external '
|
||||
'RAG_EMBEDDING_ENGINE (ollama, openai, azure_openai).'
|
||||
)
|
||||
return await asyncio.to_thread(
|
||||
(
|
||||
lambda query, prefix=None: embedding_function.encode(
|
||||
@@ -1165,6 +1320,7 @@ async def get_sources_from_items(
|
||||
):
|
||||
log.debug(f'items: {items} {queries} {embedding_function} {reranking_function} {full_context}')
|
||||
|
||||
bypass_embedding_and_retrieval = await Config.get('rag.bypass_embedding_and_retrieval')
|
||||
extracted_collections = []
|
||||
query_results = []
|
||||
|
||||
@@ -1244,14 +1400,14 @@ async def get_sources_from_items(
|
||||
}
|
||||
|
||||
elif item.get('type') == 'url':
|
||||
content, docs = get_content_from_url(request, item.get('url'))
|
||||
content, docs = await get_content_from_url(request, item.get('url'))
|
||||
if docs:
|
||||
query_result = {
|
||||
'documents': [[content]],
|
||||
'metadatas': [[{'url': item.get('url'), 'name': item.get('url')}]],
|
||||
}
|
||||
elif item.get('type') == 'file':
|
||||
if item.get('context') == 'full' or request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL:
|
||||
if item.get('context') == 'full' or bypass_embedding_and_retrieval:
|
||||
if item.get('file', {}).get('data', {}).get('content', ''):
|
||||
# Manual Full Mode Toggle
|
||||
# Used from chat file modal, we can assume that the file content will be available from item.get("file").get("data", {}).get("content")
|
||||
@@ -1323,50 +1479,61 @@ async def get_sources_from_items(
|
||||
permission='read',
|
||||
)
|
||||
):
|
||||
if item.get('context') == 'full' or request.app.state.config.BYPASS_EMBEDDING_AND_RETRIEVAL:
|
||||
if knowledge_base and (
|
||||
user.role == 'admin'
|
||||
or knowledge_base.user_id == user.id
|
||||
or await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge_base.id,
|
||||
permission='read',
|
||||
)
|
||||
):
|
||||
files = await Knowledges.get_files_by_id(knowledge_base.id)
|
||||
if (knowledge_base.meta or {}).get('source') == 'external':
|
||||
query_result = await retrieve_external_knowledge(
|
||||
request,
|
||||
knowledge_base,
|
||||
queries=queries,
|
||||
count=k,
|
||||
user=user,
|
||||
)
|
||||
extracted_collections.append(knowledge_base.id)
|
||||
|
||||
documents = []
|
||||
metadatas = []
|
||||
for file in files:
|
||||
documents.append(file.data.get('content', ''))
|
||||
metadatas.append(
|
||||
{
|
||||
'file_id': file.id,
|
||||
'name': file.filename,
|
||||
'source': file.filename,
|
||||
}
|
||||
)
|
||||
|
||||
query_result = {
|
||||
'documents': [documents],
|
||||
'metadatas': [metadatas],
|
||||
}
|
||||
else:
|
||||
if item.get('legacy'):
|
||||
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
|
||||
collection_names = item.get('collection_names', [])
|
||||
else:
|
||||
# Legacy KB: item.collection_names is client-supplied.
|
||||
# Validate against the KB's actual files to prevent
|
||||
# cross-tenant collection name substitution.
|
||||
if item.get('context') == 'full' or bypass_embedding_and_retrieval:
|
||||
if knowledge_base and (
|
||||
user.role == 'admin'
|
||||
or knowledge_base.user_id == user.id
|
||||
or await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge_base.id,
|
||||
permission='read',
|
||||
)
|
||||
):
|
||||
files = await Knowledges.get_files_by_id(knowledge_base.id)
|
||||
owned_names = {f'file-{f.id}' for f in files}
|
||||
owned_names.add(knowledge_base.id)
|
||||
valid_names = [n for n in (item.get('collection_names') or []) if n in owned_names]
|
||||
collection_names = valid_names if valid_names else [knowledge_base.id]
|
||||
|
||||
documents = []
|
||||
metadatas = []
|
||||
for file in files:
|
||||
documents.append(file.data.get('content', ''))
|
||||
metadatas.append(
|
||||
{
|
||||
'file_id': file.id,
|
||||
'name': file.filename,
|
||||
'source': file.filename,
|
||||
}
|
||||
)
|
||||
|
||||
query_result = {
|
||||
'documents': [documents],
|
||||
'metadatas': [metadatas],
|
||||
}
|
||||
else:
|
||||
collection_names.append(item['id'])
|
||||
if item.get('legacy'):
|
||||
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
|
||||
collection_names = item.get('collection_names', [])
|
||||
else:
|
||||
# Legacy KB: item.collection_names is client-supplied.
|
||||
# Validate against the KB's actual files to prevent
|
||||
# cross-tenant collection name substitution.
|
||||
files = await Knowledges.get_files_by_id(knowledge_base.id)
|
||||
owned_names = {f'file-{f.id}' for f in files}
|
||||
owned_names.add(knowledge_base.id)
|
||||
valid_names = [n for n in (item.get('collection_names') or []) if n in owned_names]
|
||||
collection_names = valid_names if valid_names else [knowledge_base.id]
|
||||
else:
|
||||
collection_names.append(item['id'])
|
||||
|
||||
elif item.get('docs'):
|
||||
# BYPASS_WEB_SEARCH_EMBEDDING_AND_RETRIEVAL
|
||||
@@ -1374,6 +1541,10 @@ async def get_sources_from_items(
|
||||
'documents': [[doc.get('content') for doc in item.get('docs')]],
|
||||
'metadatas': [[doc.get('metadata') for doc in item.get('docs')]],
|
||||
}
|
||||
elif item.get('type') == 'web_search' and item.get('collection_name'):
|
||||
# Trusted server-generated collection; authorized by
|
||||
# filter_accessible_collections below (allowlists web-search-*).
|
||||
collection_names.append(item['collection_name'])
|
||||
elif item.get('collection_name'):
|
||||
if BYPASS_RETRIEVAL_ACCESS_CONTROL:
|
||||
collection_names.append(item['collection_name'])
|
||||
|
||||
@@ -82,6 +82,10 @@ class AsyncVectorDBClient:
|
||||
(e.g. already inside a worker thread)."""
|
||||
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:
|
||||
return await asyncio.to_thread(self._sync.has_collection, collection_name)
|
||||
|
||||
@@ -103,6 +107,25 @@ class AsyncVectorDBClient:
|
||||
) -> Optional[SearchResult]:
|
||||
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(
|
||||
self._sync.hybrid_search,
|
||||
collection_name,
|
||||
query,
|
||||
vectors,
|
||||
filter,
|
||||
limit,
|
||||
hybrid_bm25_weight,
|
||||
)
|
||||
|
||||
async def query(
|
||||
self,
|
||||
collection_name: str,
|
||||
|
||||
@@ -57,7 +57,11 @@ class ChromaClient(VectorDBBase):
|
||||
|
||||
def has_collection(self, collection_name: str) -> bool:
|
||||
# Check if the collection exists based on the collection name.
|
||||
collection_names = self.client.list_collections()
|
||||
# chromadb's list_collections() returns Collection objects (1.x), so a
|
||||
# bare `name in collections` membership test is always False — compare
|
||||
# against the names. (hasattr guard tolerates versions that yield names.)
|
||||
collections = self.client.list_collections()
|
||||
collection_names = [c.name if hasattr(c, 'name') else c for c in collections]
|
||||
return collection_name in collection_names
|
||||
|
||||
def delete_collection(self, collection_name: str):
|
||||
|
||||
@@ -27,9 +27,15 @@ from open_webui.retrieval.vector.main import (
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from pymilvus import Collection, DataType, FieldSchema, connections
|
||||
from pymilvus import MilvusClient as Client
|
||||
from pymilvus.exceptions import MilvusException
|
||||
|
||||
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
|
||||
|
||||
|
||||
class MilvusClient(VectorDBBase):
|
||||
def __init__(self):
|
||||
@@ -270,18 +276,28 @@ class MilvusClient(VectorDBBase):
|
||||
self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector']))
|
||||
|
||||
log.info(f'Inserting {len(items)} items into collection {self.collection_prefix}_{collection_name}.')
|
||||
return self.client.insert(
|
||||
collection_name=f'{self.collection_prefix}_{collection_name}',
|
||||
data=[
|
||||
data = []
|
||||
for item in items:
|
||||
text = item['text'] or ''
|
||||
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'],
|
||||
'vector': item['vector'],
|
||||
'data': {'text': item['text']},
|
||||
'data': {'text': text},
|
||||
'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]):
|
||||
# Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
|
||||
@@ -298,18 +314,28 @@ class MilvusClient(VectorDBBase):
|
||||
self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector']))
|
||||
|
||||
log.info(f'Upserting {len(items)} items into collection {self.collection_prefix}_{collection_name}.')
|
||||
return self.client.upsert(
|
||||
collection_name=f'{self.collection_prefix}_{collection_name}',
|
||||
data=[
|
||||
data = []
|
||||
for item in items:
|
||||
text = item['text'] or ''
|
||||
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'],
|
||||
'vector': item['vector'],
|
||||
'data': {'text': item['text']},
|
||||
'data': {'text': text},
|
||||
'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(
|
||||
self,
|
||||
|
||||
@@ -31,10 +31,15 @@ from pymilvus import (
|
||||
connections,
|
||||
utility,
|
||||
)
|
||||
from pymilvus.exceptions import MilvusException
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
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;
|
||||
# values get interpolated into single-quoted literals. Reject anything that
|
||||
@@ -169,17 +174,34 @@ class MilvusClient(VectorDBBase):
|
||||
self._ensure_collection(mt_collection, dimension)
|
||||
collection = Collection(mt_collection)
|
||||
|
||||
entities = [
|
||||
{
|
||||
'id': item['id'],
|
||||
'vector': item['vector'],
|
||||
'text': item['text'],
|
||||
'metadata': item['metadata'],
|
||||
RESOURCE_ID_FIELD: resource_id,
|
||||
}
|
||||
for item in items
|
||||
]
|
||||
collection.insert(entities)
|
||||
entities = []
|
||||
for item in items:
|
||||
text = item['text'] or ''
|
||||
if len(text) > MILVUS_TEXT_MAX_LENGTH:
|
||||
log.warning(
|
||||
f'Milvus: truncating text id={item["id"]} '
|
||||
f'{len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars '
|
||||
f'(collection={mt_collection}, resource_id={resource_id})'
|
||||
)
|
||||
text = text[:MILVUS_TEXT_MAX_LENGTH]
|
||||
entities.append(
|
||||
{
|
||||
'id': item['id'],
|
||||
'vector': item['vector'],
|
||||
'text': text,
|
||||
'metadata': item['metadata'],
|
||||
RESOURCE_ID_FIELD: resource_id,
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
collection.insert(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(
|
||||
self,
|
||||
|
||||
@@ -24,7 +24,7 @@ from open_webui.retrieval.vector.main import (
|
||||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import merge_hybrid_search_results, process_metadata
|
||||
from open_webui.utils.misc import sanitize_text_for_db
|
||||
from pgvector.sqlalchemy import HALFVEC, Vector
|
||||
from sqlalchemy import (
|
||||
@@ -153,6 +153,7 @@ class PgvectorClient(VectorDBBase):
|
||||
|
||||
index_method, index_options = self._vector_index_configuration()
|
||||
self._ensure_vector_index(index_method, index_options)
|
||||
self._ensure_text_search_index()
|
||||
|
||||
self.session.execute(
|
||||
text(
|
||||
@@ -236,6 +237,19 @@ class PgvectorClient(VectorDBBase):
|
||||
f' {index_options}' if index_options else '',
|
||||
)
|
||||
|
||||
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:
|
||||
"""
|
||||
Check if the VECTOR_LENGTH matches the existing vector column dimension in the database.
|
||||
@@ -521,6 +535,71 @@ class PgvectorClient(VectorDBBase):
|
||||
log.exception(f'Error during search: {e}')
|
||||
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]:
|
||||
try:
|
||||
if PGVECTOR_PGCRYPTO:
|
||||
|
||||
@@ -63,6 +63,18 @@ class VectorDBBase(ABC):
|
||||
"""Search for similar vectors in a collection."""
|
||||
pass
|
||||
|
||||
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]:
|
||||
"""Search using a backend-native hybrid keyword/vector implementation when available."""
|
||||
return None
|
||||
|
||||
@abstractmethod
|
||||
def query(self, collection_name: str, filter: Dict, limit: Optional[int] = None) -> Optional[GetResult]:
|
||||
"""Query vectors from a collection using metadata filter."""
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from datetime import datetime
|
||||
import datetime as dt
|
||||
from typing import Any
|
||||
|
||||
from open_webui.retrieval.vector.main import SearchResult
|
||||
from open_webui.utils.misc import sanitize_text_for_db
|
||||
|
||||
KEYS_TO_EXCLUDE = ['content', 'pages', 'tables', 'paragraphs', 'sections', 'figures']
|
||||
@@ -21,9 +23,71 @@ def process_metadata(
|
||||
# Skip large fields
|
||||
if key in KEYS_TO_EXCLUDE:
|
||||
continue
|
||||
if value is None:
|
||||
continue
|
||||
# Convert non-serializable fields to strings
|
||||
if isinstance(value, (datetime, list, dict)):
|
||||
if isinstance(value, (dt.datetime, list, dict)):
|
||||
result[key] = sanitize_text_for_db(str(value))
|
||||
else:
|
||||
result[key] = sanitize_text_for_db(value)
|
||||
return result
|
||||
|
||||
|
||||
def merge_hybrid_search_results(
|
||||
vector_result: SearchResult | None,
|
||||
fts_results: list[dict[str, Any]],
|
||||
num_queries: int,
|
||||
limit: int,
|
||||
hybrid_bm25_weight: float,
|
||||
) -> SearchResult:
|
||||
rank_constant = 60.0
|
||||
bm25_weight = min(max(hybrid_bm25_weight, 0.0), 1.0)
|
||||
vector_weight = 1.0 - bm25_weight
|
||||
|
||||
ids = [[] for _ in range(num_queries)]
|
||||
distances = [[] for _ in range(num_queries)]
|
||||
documents = [[] for _ in range(num_queries)]
|
||||
metadatas = [[] for _ in range(num_queries)]
|
||||
|
||||
for qid in range(num_queries):
|
||||
candidates: dict[str, dict[str, Any]] = {}
|
||||
|
||||
if vector_result and vector_result.ids and qid < len(vector_result.ids):
|
||||
for rank, item_id in enumerate(vector_result.ids[qid] or [], start=1):
|
||||
score = vector_weight / (rank_constant + rank) if vector_weight > 0 else 0
|
||||
if score <= 0:
|
||||
continue
|
||||
|
||||
candidate = candidates.setdefault(
|
||||
item_id,
|
||||
{
|
||||
'score': 0.0,
|
||||
'document': vector_result.documents[qid][rank - 1],
|
||||
'metadata': vector_result.metadatas[qid][rank - 1],
|
||||
},
|
||||
)
|
||||
candidate['score'] += score
|
||||
|
||||
for rank, row in enumerate(fts_results, start=1):
|
||||
score = bm25_weight / (rank_constant + rank) if bm25_weight > 0 else 0
|
||||
if score <= 0:
|
||||
continue
|
||||
|
||||
item_id = row['id']
|
||||
candidate = candidates.setdefault(
|
||||
item_id,
|
||||
{
|
||||
'score': 0.0,
|
||||
'document': row['text'],
|
||||
'metadata': row['vmetadata'],
|
||||
},
|
||||
)
|
||||
candidate['score'] += score
|
||||
|
||||
ranked = sorted(candidates.items(), key=lambda item: item[1]['score'], reverse=True)[:limit]
|
||||
ids[qid] = [item_id for item_id, _ in ranked]
|
||||
distances[qid] = [candidate['score'] for _, candidate in ranked]
|
||||
documents[qid] = [candidate['document'] for _, candidate in ranked]
|
||||
metadatas[qid] = [candidate['metadata'] for _, candidate in ranked]
|
||||
|
||||
return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas)
|
||||
|
||||
@@ -4,7 +4,7 @@ from urllib.parse import urlparse
|
||||
|
||||
import validators
|
||||
from open_webui.retrieval.web.utils import resolve_hostname
|
||||
from open_webui.utils.misc import is_string_allowed
|
||||
from open_webui.utils.misc import is_host_allowed
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ def get_filtered_results(results, filter_list):
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if is_string_allowed(hostnames, filter_list):
|
||||
if is_host_allowed(hostnames, filter_list):
|
||||
filtered_results.append(result)
|
||||
continue
|
||||
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import requests
|
||||
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MICROSOFT_WEB_IQ_API_BASE_URL = 'https://api.microsoft.ai/v3'
|
||||
|
||||
|
||||
def search_microsoft_web_iq(
|
||||
api_base_url: str,
|
||||
api_key: str,
|
||||
query: str,
|
||||
count: int,
|
||||
filter_list: list[str | None] | None = None,
|
||||
language: str = 'en',
|
||||
user=None,
|
||||
) -> list[SearchResult]:
|
||||
try:
|
||||
api_base_url = (api_base_url or DEFAULT_MICROSOFT_WEB_IQ_API_BASE_URL).rstrip('/')
|
||||
headers = {
|
||||
'host': urlparse(api_base_url).netloc or 'api.microsoft.ai',
|
||||
'x-apikey': api_key,
|
||||
'content-type': 'application/json',
|
||||
}
|
||||
if user is not None:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
response = requests.post(
|
||||
f'{api_base_url}/search/web',
|
||||
json={
|
||||
'query': query,
|
||||
'maxResults': count,
|
||||
'language': language,
|
||||
'contentFormat': 'passage',
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
results = response.json().get('webResults', [])
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result['url'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('content'),
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
except Exception as e:
|
||||
log.error(f'Error searching with Microsoft Web IQ API: {e}')
|
||||
return []
|
||||
@@ -38,9 +38,7 @@ def search_perplexity(
|
||||
|
||||
"""
|
||||
|
||||
# Handle ConfigVar object
|
||||
if hasattr(api_key, '__str__'):
|
||||
api_key = str(api_key)
|
||||
api_key = str(api_key)
|
||||
|
||||
try:
|
||||
url = 'https://api.perplexity.ai/chat/completions'
|
||||
|
||||
@@ -29,12 +29,8 @@ def search_perplexity_search(
|
||||
|
||||
"""
|
||||
|
||||
# Handle ConfigVar object
|
||||
if hasattr(api_key, '__str__'):
|
||||
api_key = str(api_key)
|
||||
|
||||
if hasattr(api_url, '__str__'):
|
||||
api_url = str(api_url)
|
||||
api_key = str(api_key)
|
||||
api_url = str(api_url)
|
||||
|
||||
try:
|
||||
url = api_url
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
from open_webui.utils.session_pool import get_session
|
||||
|
||||
|
||||
async def search_serphouse(
|
||||
api_key: str,
|
||||
domain: str,
|
||||
query: str,
|
||||
count: int,
|
||||
filter_list: list[str | None] | None = None,
|
||||
) -> list[SearchResult]:
|
||||
"""Query SERPHouse and return normalised organic results."""
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
'https://api.serphouse.com/serp/live',
|
||||
params={
|
||||
'q': query,
|
||||
'domain': (domain or 'google.com').strip() or 'google.com',
|
||||
'device': 'desktop',
|
||||
'serp_type': 'web',
|
||||
'page': 1,
|
||||
'num_result': count,
|
||||
},
|
||||
headers={'Authorization': f'Bearer {api_key}', 'Accept': 'application/json'},
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
payload = await response.json()
|
||||
|
||||
organic = payload.get('results', {}).get('results', {}).get('organic', [])
|
||||
organic = sorted(organic, key=lambda item: item.get('position', 0))
|
||||
if filter_list:
|
||||
organic = get_filtered_results(organic, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=item.get('link', ''),
|
||||
title=item.get('title'),
|
||||
snippet=item.get('snippet'),
|
||||
)
|
||||
for item in organic[:count]
|
||||
if item.get('link')
|
||||
]
|
||||
@@ -21,7 +21,6 @@ from typing import (
|
||||
import aiohttp
|
||||
import aiohttp.resolver
|
||||
import certifi
|
||||
import requests
|
||||
import urllib3.connection
|
||||
import urllib3.connectionpool
|
||||
import validators
|
||||
@@ -31,12 +30,15 @@ from langchain_community.document_loaders import PlaywrightURLLoader, WebBaseLoa
|
||||
from langchain_community.document_loaders.base import BaseLoader
|
||||
from langchain_core.documents import Document
|
||||
from open_webui.config import (
|
||||
ENABLE_RAG_LOCAL_WEB_FETCH,
|
||||
ENABLE_LOCAL_WEB_FETCH,
|
||||
EXTERNAL_WEB_LOADER_API_KEY,
|
||||
EXTERNAL_WEB_LOADER_URL,
|
||||
FIRECRAWL_API_BASE_URL,
|
||||
FIRECRAWL_API_KEY,
|
||||
FIRECRAWL_TIMEOUT,
|
||||
MICROSOFT_WEB_IQ_API_BASE_URL,
|
||||
MICROSOFT_WEB_IQ_API_KEY,
|
||||
MICROSOFT_WEB_IQ_LANGUAGE,
|
||||
PLAYWRIGHT_TIMEOUT,
|
||||
PLAYWRIGHT_WS_URL,
|
||||
TAVILY_API_KEY,
|
||||
@@ -46,11 +48,17 @@ from open_webui.config import (
|
||||
WEB_LOADER_TIMEOUT,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_ALLOW_REDIRECTS, AIOHTTP_CLIENT_SESSION_SSL, USER_AGENT
|
||||
from open_webui.env import (
|
||||
AIOHTTP_CLIENT_ALLOW_REDIRECTS,
|
||||
AIOHTTP_CLIENT_SESSION_SSL,
|
||||
AIOHTTP_CLIENT_TIMEOUT,
|
||||
USER_AGENT,
|
||||
)
|
||||
from open_webui.retrieval.loaders.external_web import ExternalWebLoader
|
||||
from open_webui.retrieval.loaders.microsoft_web_iq import MicrosoftWebIQLoader
|
||||
from open_webui.retrieval.loaders.tavily import TavilyLoader
|
||||
from open_webui.retrieval.web.firecrawl import scrape_firecrawl_url
|
||||
from open_webui.utils.misc import is_string_allowed
|
||||
from open_webui.utils.misc import is_host_allowed
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -88,12 +96,14 @@ def validate_url(url: Union[str, Sequence[str]]):
|
||||
|
||||
# Blocklist check using unified filtering logic
|
||||
if WEB_FETCH_FILTER_LIST:
|
||||
if not is_string_allowed(url, WEB_FETCH_FILTER_LIST):
|
||||
# Match on the parsed hostname, not the full URL: a path component would
|
||||
# otherwise let any URL slip past a hostname-based block/allow entry.
|
||||
if not is_host_allowed(parsed_url.hostname, WEB_FETCH_FILTER_LIST):
|
||||
log.warning(f'URL blocked by filter list: {url}')
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
|
||||
if not ENABLE_RAG_LOCAL_WEB_FETCH:
|
||||
# Local web fetch is disabled, filter out any URLs that resolve to private IP addresses
|
||||
if not ENABLE_LOCAL_WEB_FETCH:
|
||||
# Local web fetch is disabled, filter out URLs that resolve to non-global IP addresses.
|
||||
parsed_url = urllib.parse.urlparse(url)
|
||||
# Get IPv4 and IPv6 addresses
|
||||
ipv4_addresses, ipv6_addresses = resolve_hostname(parsed_url.hostname)
|
||||
@@ -134,7 +144,7 @@ def _ssrf_safe_new_conn(self):
|
||||
infos = socket.getaddrinfo(host, port, 0, socket.SOCK_STREAM)
|
||||
if not infos:
|
||||
raise OSError(f'getaddrinfo for {host!r} returned empty list')
|
||||
if not ENABLE_RAG_LOCAL_WEB_FETCH:
|
||||
if not ENABLE_LOCAL_WEB_FETCH:
|
||||
for _, _, _, _, sa in infos:
|
||||
if not ipaddress.ip_address(sa[0]).is_global:
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
@@ -190,13 +200,26 @@ class _SSRFSafeResolver(aiohttp.resolver.DefaultResolver):
|
||||
|
||||
async def resolve(self, host, port=0, family=socket.AF_INET):
|
||||
results = await super().resolve(host, port, family)
|
||||
if not ENABLE_RAG_LOCAL_WEB_FETCH:
|
||||
if not ENABLE_LOCAL_WEB_FETCH:
|
||||
for entry in results:
|
||||
if not ipaddress.ip_address(entry['host']).is_global:
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
return results
|
||||
|
||||
|
||||
def get_ssrf_safe_session() -> aiohttp.ClientSession:
|
||||
"""A one-off aiohttp session that re-validates the connect-time IP via _SSRFSafeResolver,
|
||||
defeating DNS rebinding. Use for validate_url-gated fetches of user-supplied URLs that must
|
||||
not use the shared (rebinding-vulnerable) pool. Use as a context manager so it is closed:
|
||||
``async with get_ssrf_safe_session() as session: ...``.
|
||||
"""
|
||||
return aiohttp.ClientSession(
|
||||
connector=aiohttp.TCPConnector(resolver=_SSRFSafeResolver()),
|
||||
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
|
||||
trust_env=True,
|
||||
)
|
||||
|
||||
|
||||
def extract_metadata(soup, url):
|
||||
metadata = {'source': url}
|
||||
if title := soup.find('title'):
|
||||
@@ -303,8 +326,9 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
self.params = params or {}
|
||||
|
||||
def lazy_load(self) -> Iterator[Document]:
|
||||
try:
|
||||
for url in self.web_paths:
|
||||
for url in self.web_paths:
|
||||
try:
|
||||
self._sync_wait_for_rate_limit()
|
||||
doc = scrape_firecrawl_url(
|
||||
self.api_url,
|
||||
self.api_key,
|
||||
@@ -315,28 +339,39 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
)
|
||||
if doc is not None:
|
||||
yield doc
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.warning(f'Error extracting content from URLs with Firecrawl: {e}')
|
||||
else:
|
||||
raise e
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.warning(f'Error extracting content from {url} with Firecrawl: {e}')
|
||||
continue
|
||||
raise
|
||||
|
||||
async def alazy_load(self):
|
||||
try:
|
||||
docs = await run_in_threadpool(lambda: list(self.lazy_load()))
|
||||
for doc in docs:
|
||||
yield doc
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.warning(f'Error extracting content from URLs with Firecrawl: {e}')
|
||||
else:
|
||||
raise e
|
||||
for url in self.web_paths:
|
||||
try:
|
||||
await self._wait_for_rate_limit()
|
||||
doc = await run_in_threadpool(
|
||||
scrape_firecrawl_url,
|
||||
self.api_url,
|
||||
self.api_key,
|
||||
url,
|
||||
verify_ssl=self.verify_ssl,
|
||||
timeout=self.timeout,
|
||||
params=self.params,
|
||||
)
|
||||
if doc is not None:
|
||||
yield doc
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.warning(f'Error extracting content from {url} with Firecrawl: {e}')
|
||||
continue
|
||||
raise
|
||||
|
||||
|
||||
class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
def __init__(
|
||||
self,
|
||||
web_paths: Union[str, List[str]],
|
||||
api_base_url: str,
|
||||
api_key: str,
|
||||
extract_depth: Literal['basic', 'advanced'] = 'basic',
|
||||
continue_on_failure: bool = True,
|
||||
@@ -370,6 +405,7 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
|
||||
# Store parameters for creating TavilyLoader instances
|
||||
self.web_paths = web_paths if isinstance(web_paths, list) else [web_paths]
|
||||
self.api_base_url = api_base_url
|
||||
self.api_key = api_key
|
||||
self.extract_depth = extract_depth
|
||||
self.continue_on_failure = continue_on_failure
|
||||
@@ -445,6 +481,67 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
raise e
|
||||
|
||||
|
||||
class SafeMicrosoftWebIQLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
def __init__(
|
||||
self,
|
||||
web_paths: Union[str, List[str]],
|
||||
api_key: str,
|
||||
language: str = 'en',
|
||||
verify_ssl: bool = True,
|
||||
trust_env: bool = False,
|
||||
requests_per_second: Optional[float] = None,
|
||||
continue_on_failure: bool = True,
|
||||
timeout: Optional[int] = None,
|
||||
):
|
||||
self.web_paths = web_paths if isinstance(web_paths, list) else [web_paths]
|
||||
self.api_key = api_key
|
||||
self.language = language
|
||||
self.verify_ssl = verify_ssl
|
||||
self.trust_env = trust_env
|
||||
self.requests_per_second = requests_per_second
|
||||
self.last_request_time = None
|
||||
self.continue_on_failure = continue_on_failure
|
||||
self.timeout = timeout
|
||||
|
||||
def lazy_load(self) -> Iterator[Document]:
|
||||
valid_urls = []
|
||||
for url in self.web_paths:
|
||||
try:
|
||||
self._safe_process_url_sync(url)
|
||||
valid_urls.append(url)
|
||||
except Exception as e:
|
||||
log.warning(f'SSL verification failed for {url}: {str(e)}')
|
||||
if not self.continue_on_failure:
|
||||
raise e
|
||||
if not valid_urls:
|
||||
if self.continue_on_failure:
|
||||
log.warning('No valid URLs to process after SSL verification')
|
||||
return
|
||||
raise ValueError('No valid URLs to process after SSL verification')
|
||||
|
||||
loader = MicrosoftWebIQLoader(
|
||||
urls=valid_urls,
|
||||
api_base_url=self.api_base_url,
|
||||
api_key=self.api_key,
|
||||
language=self.language,
|
||||
verify_ssl=self.verify_ssl,
|
||||
timeout=self.timeout,
|
||||
continue_on_failure=self.continue_on_failure,
|
||||
)
|
||||
yield from loader.lazy_load()
|
||||
|
||||
async def alazy_load(self) -> AsyncIterator[Document]:
|
||||
try:
|
||||
docs = await run_in_threadpool(lambda: list(self.lazy_load()))
|
||||
for doc in docs:
|
||||
yield doc
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.warning(f'Error browsing URLs with Microsoft Web IQ: {e}')
|
||||
else:
|
||||
raise e
|
||||
|
||||
|
||||
class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessingMixin):
|
||||
"""Load HTML pages safely with Playwright, supporting SSL verification, rate limiting, and remote browser connection.
|
||||
|
||||
@@ -757,13 +854,13 @@ def get_web_loader(
|
||||
'trust_env': trust_env,
|
||||
}
|
||||
|
||||
if WEB_LOADER_ENGINE.value == '' or WEB_LOADER_ENGINE.value == 'safe_web':
|
||||
if WEB_LOADER_ENGINE == '' or WEB_LOADER_ENGINE == 'safe_web':
|
||||
WebLoaderClass = SafeWebBaseLoader
|
||||
|
||||
request_kwargs = {}
|
||||
if WEB_LOADER_TIMEOUT.value:
|
||||
if WEB_LOADER_TIMEOUT:
|
||||
try:
|
||||
timeout_value = float(WEB_LOADER_TIMEOUT.value)
|
||||
timeout_value = float(WEB_LOADER_TIMEOUT)
|
||||
except ValueError:
|
||||
timeout_value = None
|
||||
|
||||
@@ -773,31 +870,42 @@ def get_web_loader(
|
||||
if request_kwargs:
|
||||
web_loader_args['requests_kwargs'] = request_kwargs
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'playwright':
|
||||
if WEB_LOADER_ENGINE == 'playwright':
|
||||
WebLoaderClass = SafePlaywrightURLLoader
|
||||
web_loader_args['playwright_timeout'] = PLAYWRIGHT_TIMEOUT.value
|
||||
if PLAYWRIGHT_WS_URL.value:
|
||||
web_loader_args['playwright_ws_url'] = PLAYWRIGHT_WS_URL.value
|
||||
web_loader_args['playwright_timeout'] = PLAYWRIGHT_TIMEOUT
|
||||
if PLAYWRIGHT_WS_URL:
|
||||
web_loader_args['playwright_ws_url'] = PLAYWRIGHT_WS_URL
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'firecrawl':
|
||||
if WEB_LOADER_ENGINE == 'firecrawl':
|
||||
WebLoaderClass = SafeFireCrawlLoader
|
||||
web_loader_args['api_key'] = FIRECRAWL_API_KEY.value
|
||||
web_loader_args['api_url'] = FIRECRAWL_API_BASE_URL.value
|
||||
if FIRECRAWL_TIMEOUT.value:
|
||||
web_loader_args['api_key'] = FIRECRAWL_API_KEY
|
||||
web_loader_args['api_url'] = FIRECRAWL_API_BASE_URL
|
||||
if FIRECRAWL_TIMEOUT:
|
||||
try:
|
||||
web_loader_args['timeout'] = int(FIRECRAWL_TIMEOUT.value)
|
||||
web_loader_args['timeout'] = int(FIRECRAWL_TIMEOUT)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'tavily':
|
||||
if WEB_LOADER_ENGINE == 'tavily':
|
||||
WebLoaderClass = SafeTavilyLoader
|
||||
web_loader_args['api_key'] = TAVILY_API_KEY.value
|
||||
web_loader_args['extract_depth'] = TAVILY_EXTRACT_DEPTH.value
|
||||
web_loader_args['api_key'] = TAVILY_API_KEY
|
||||
web_loader_args['extract_depth'] = TAVILY_EXTRACT_DEPTH
|
||||
|
||||
if WEB_LOADER_ENGINE.value == 'external':
|
||||
if WEB_LOADER_ENGINE == 'microsoft_web_iq':
|
||||
WebLoaderClass = SafeMicrosoftWebIQLoader
|
||||
web_loader_args['api_base_url'] = MICROSOFT_WEB_IQ_API_BASE_URL
|
||||
web_loader_args['api_key'] = MICROSOFT_WEB_IQ_API_KEY
|
||||
web_loader_args['language'] = MICROSOFT_WEB_IQ_LANGUAGE
|
||||
if WEB_LOADER_TIMEOUT:
|
||||
try:
|
||||
web_loader_args['timeout'] = int(WEB_LOADER_TIMEOUT)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if WEB_LOADER_ENGINE == 'external':
|
||||
WebLoaderClass = ExternalWebLoader
|
||||
web_loader_args['external_url'] = EXTERNAL_WEB_LOADER_URL.value
|
||||
web_loader_args['external_api_key'] = EXTERNAL_WEB_LOADER_API_KEY.value
|
||||
web_loader_args['external_url'] = EXTERNAL_WEB_LOADER_URL
|
||||
web_loader_args['external_api_key'] = EXTERNAL_WEB_LOADER_API_KEY
|
||||
|
||||
if WebLoaderClass:
|
||||
web_loader = WebLoaderClass(**web_loader_args)
|
||||
@@ -811,6 +919,6 @@ def get_web_loader(
|
||||
return web_loader
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Invalid WEB_LOADER_ENGINE: {WEB_LOADER_ENGINE.value}. '
|
||||
"Please set it to 'safe_web', 'playwright', 'firecrawl', or 'tavily'."
|
||||
f'Invalid WEB_LOADER_ENGINE: {WEB_LOADER_ENGINE}. '
|
||||
"Please set it to 'safe_web', 'playwright', 'firecrawl', 'tavily', 'external', or 'microsoft_web_iq'."
|
||||
)
|
||||
|
||||
@@ -28,6 +28,8 @@ router = APIRouter()
|
||||
class ModelAnalyticsEntry(BaseModel):
|
||||
model_id: str
|
||||
count: int
|
||||
unique_users: int = 0
|
||||
unique_chats: int = 0
|
||||
|
||||
|
||||
class ModelAnalyticsResponse(BaseModel):
|
||||
@@ -65,8 +67,16 @@ async def get_model_analytics(
|
||||
counts = await ChatMessages.get_message_count_by_model(
|
||||
start_date=start_date, end_date=end_date, group_id=group_id, db=db
|
||||
)
|
||||
unique_counts = await ChatMessages.get_unique_counts_by_model(
|
||||
start_date=start_date, end_date=end_date, group_id=group_id, db=db
|
||||
)
|
||||
models = [
|
||||
ModelAnalyticsEntry(model_id=model_id, count=count)
|
||||
ModelAnalyticsEntry(
|
||||
model_id=model_id,
|
||||
count=count,
|
||||
unique_users=unique_counts.get(model_id, {}).get('unique_users', 0),
|
||||
unique_chats=unique_counts.get(model_id, {}).get('unique_chats', 0),
|
||||
)
|
||||
for model_id, count in sorted(counts.items(), key=lambda x: -x[1])
|
||||
]
|
||||
return ModelAnalyticsResponse(models=models)
|
||||
@@ -269,6 +279,9 @@ class ModelChatsResponse(BaseModel):
|
||||
total: int
|
||||
|
||||
|
||||
MODEL_CHAT_ORDER_FIELDS = {'title', 'updated_at', 'user_name'}
|
||||
|
||||
|
||||
@router.get('/models/{model_id:path}/chats', response_model=ModelChatsResponse)
|
||||
async def get_model_chats(
|
||||
model_id: str,
|
||||
@@ -276,65 +289,34 @@ async def get_model_chats(
|
||||
end_date: Optional[int] = Query(None),
|
||||
skip: int = Query(0),
|
||||
limit: int = Query(50, le=100),
|
||||
order_by: str = Query('updated_at'),
|
||||
direction: str = Query('desc'),
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
"""Get chats that used a specific model, with preview and feedback info."""
|
||||
filter = {}
|
||||
if start_date:
|
||||
filter['start_date'] = start_date
|
||||
if end_date:
|
||||
filter['end_date'] = end_date
|
||||
if order_by in MODEL_CHAT_ORDER_FIELDS:
|
||||
filter['order_by'] = order_by
|
||||
if direction in {'asc', 'desc'}:
|
||||
filter['direction'] = direction
|
||||
|
||||
# Get chat IDs that used this model
|
||||
chat_ids = await ChatMessages.get_chat_ids_by_model_id(
|
||||
result = await Chats.get_chats_by_model_id(
|
||||
model_id=model_id,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
filter=filter,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
db=db,
|
||||
)
|
||||
|
||||
if not chat_ids:
|
||||
return ModelChatsResponse(chats=[], total=0)
|
||||
|
||||
# Get chat details from messages only
|
||||
chats_data = []
|
||||
for chat_id in chat_ids:
|
||||
messages = await ChatMessages.get_messages_by_chat_id(chat_id, db=db)
|
||||
if not messages:
|
||||
continue
|
||||
|
||||
# Get user_id from first user message
|
||||
first_user_msg = next((m for m in messages if m.role == 'user'), None)
|
||||
user_id = first_user_msg.user_id if first_user_msg else None
|
||||
|
||||
# Extract first message content as preview
|
||||
first_message = None
|
||||
if first_user_msg and first_user_msg.content:
|
||||
content = first_user_msg.content
|
||||
if isinstance(content, str):
|
||||
first_message = content[:200]
|
||||
elif isinstance(content, list):
|
||||
text_parts = [b.get('text', '') for b in content if isinstance(b, dict)]
|
||||
first_message = ' '.join(text_parts)[:200]
|
||||
|
||||
# Get user info
|
||||
user_name = None
|
||||
if user_id:
|
||||
user_info = await Users.get_user_by_id(user_id, db=db)
|
||||
user_name = user_info.name if user_info else None
|
||||
|
||||
# Timestamps from messages
|
||||
updated_at = max(m.created_at for m in messages) if messages else 0
|
||||
|
||||
chats_data.append(
|
||||
ModelChatEntry(
|
||||
chat_id=chat_id,
|
||||
user_id=user_id,
|
||||
user_name=user_name,
|
||||
first_message=first_message,
|
||||
updated_at=updated_at,
|
||||
)
|
||||
)
|
||||
|
||||
return ModelChatsResponse(chats=chats_data, total=len(chats_data))
|
||||
return ModelChatsResponse(
|
||||
chats=[ModelChatEntry.model_validate(chat) for chat in result['items']],
|
||||
total=result['total'] or 0,
|
||||
)
|
||||
|
||||
|
||||
####################
|
||||
@@ -367,6 +349,12 @@ async def get_model_overview(
|
||||
):
|
||||
"""Get model overview with feedback history and chat tags."""
|
||||
|
||||
# Calculate start date for history
|
||||
now = datetime.now()
|
||||
start_dt = None
|
||||
if days > 0:
|
||||
start_dt = now - timedelta(days=days)
|
||||
|
||||
# Get chat IDs that used this model
|
||||
chat_ids = await ChatMessages.get_chat_ids_by_model_id(
|
||||
model_id=model_id,
|
||||
@@ -377,31 +365,18 @@ async def get_model_overview(
|
||||
db=db,
|
||||
)
|
||||
|
||||
# Get feedback history per day
|
||||
history_counts: dict[str, dict] = defaultdict(lambda: {'won': 0, 'lost': 0})
|
||||
|
||||
# Calculate start date for history
|
||||
now = datetime.now()
|
||||
start_dt = None
|
||||
if days > 0:
|
||||
start_dt = now - timedelta(days=days)
|
||||
|
||||
for chat_id in chat_ids:
|
||||
feedbacks = await Feedbacks.get_feedbacks_by_chat_id(chat_id, db=db)
|
||||
for fb in feedbacks:
|
||||
if fb.data and 'rating' in fb.data:
|
||||
rating = fb.data['rating']
|
||||
fb_date = datetime.fromtimestamp(fb.created_at)
|
||||
|
||||
# Filter by date range
|
||||
if start_dt and fb_date < start_dt:
|
||||
continue
|
||||
|
||||
date_str = fb_date.strftime('%Y-%m-%d')
|
||||
if rating == 1:
|
||||
history_counts[date_str]['won'] += 1
|
||||
elif rating == -1:
|
||||
history_counts[date_str]['lost'] += 1
|
||||
history_rows = await Feedbacks.get_model_feedback_counts_by_day(
|
||||
model_id=model_id,
|
||||
start_date=int(start_dt.timestamp()) if start_dt else None,
|
||||
db=db,
|
||||
)
|
||||
history_counts = {
|
||||
entry.date: {
|
||||
'won': entry.won,
|
||||
'lost': entry.lost,
|
||||
}
|
||||
for entry in history_rows
|
||||
}
|
||||
|
||||
# Fill in missing days
|
||||
history = []
|
||||
@@ -430,10 +405,14 @@ async def get_model_overview(
|
||||
|
||||
# Get chat tags
|
||||
tag_counts: dict[str, int] = defaultdict(int)
|
||||
for chat_id in chat_ids:
|
||||
chat = await Chats.get_chat_by_id(chat_id, db=db)
|
||||
if chat and chat.meta:
|
||||
for tag in chat.meta.get('tags', []):
|
||||
if chat_ids:
|
||||
chat_metas = await Chats.get_chat_metas_by_chat_ids(
|
||||
chat_ids,
|
||||
include_archived=True,
|
||||
db=db,
|
||||
)
|
||||
for meta in chat_metas:
|
||||
for tag in meta.get('tags', []):
|
||||
tag_counts[tag] += 1
|
||||
|
||||
# Sort by count and take top 10
|
||||
|
||||
+178
-169
@@ -28,6 +28,8 @@ from fastapi import (
|
||||
)
|
||||
from fastapi.responses import FileResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
# pydub needs stdlib audioop (gone in 3.13); keep requires-python capped < 3.13
|
||||
from pydub import AudioSegment
|
||||
from pydub.silence import split_on_silence
|
||||
from pydub.utils import mediainfo
|
||||
@@ -52,6 +54,8 @@ from open_webui.env import (
|
||||
ENABLE_FORWARD_USER_INFO_HEADERS,
|
||||
ENV,
|
||||
)
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
@@ -71,6 +75,50 @@ AZURE_MAX_FILE_SIZE: int = AZURE_MAX_FILE_SIZE_MB * 1024 * 1024
|
||||
SPEECH_CACHE_DIR = CACHE_DIR / 'audio' / 'speech'
|
||||
SPEECH_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
TTS_CONFIG_KEYS = {
|
||||
'OPENAI_API_BASE_URL': 'audio.tts.openai.api_base_url',
|
||||
'OPENAI_API_KEY': 'audio.tts.openai.api_key',
|
||||
'OPENAI_PARAMS': 'audio.tts.openai.params',
|
||||
'API_KEY': 'audio.tts.api_key',
|
||||
'ENGINE': 'audio.tts.engine',
|
||||
'MODEL': 'audio.tts.model',
|
||||
'VOICE': 'audio.tts.voice',
|
||||
'SPLIT_ON': 'audio.tts.split_on',
|
||||
'AZURE_SPEECH_REGION': 'audio.tts.azure.speech_region',
|
||||
'AZURE_SPEECH_BASE_URL': 'audio.tts.azure.speech_base_url',
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': 'audio.tts.azure.speech_output_format',
|
||||
'MISTRAL_API_KEY': 'audio.tts.mistral.api_key',
|
||||
'MISTRAL_API_BASE_URL': 'audio.tts.mistral.api_base_url',
|
||||
}
|
||||
|
||||
STT_CONFIG_KEYS = {
|
||||
'OPENAI_API_BASE_URL': 'audio.stt.openai.api_base_url',
|
||||
'OPENAI_API_KEY': 'audio.stt.openai.api_key',
|
||||
'ENGINE': 'audio.stt.engine',
|
||||
'MODEL': 'audio.stt.model',
|
||||
'SUPPORTED_CONTENT_TYPES': 'audio.stt.supported_content_types',
|
||||
'ALLOWED_EXTENSIONS': 'audio.stt.allowed_extensions',
|
||||
'WHISPER_MODEL': 'audio.stt.whisper_model',
|
||||
'DEEPGRAM_API_KEY': 'audio.stt.deepgram.api_key',
|
||||
'AZURE_API_KEY': 'audio.stt.azure.api_key',
|
||||
'AZURE_REGION': 'audio.stt.azure.region',
|
||||
'AZURE_LOCALES': 'audio.stt.azure.locales',
|
||||
'AZURE_BASE_URL': 'audio.stt.azure.base_url',
|
||||
'AZURE_MAX_SPEAKERS': 'audio.stt.azure.max_speakers',
|
||||
'MISTRAL_API_KEY': 'audio.stt.mistral.api_key',
|
||||
'MISTRAL_API_BASE_URL': 'audio.stt.mistral.api_base_url',
|
||||
'MISTRAL_USE_CHAT_COMPLETIONS': 'audio.stt.mistral.use_chat_completions',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
def is_audio_conversion_required(file_path):
|
||||
"""
|
||||
@@ -228,119 +276,39 @@ class AudioConfigUpdateForm(BaseModel):
|
||||
@router.get('/config')
|
||||
async def get_audio_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'tts': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.TTS_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.TTS_OPENAI_API_KEY,
|
||||
'OPENAI_PARAMS': request.app.state.config.TTS_OPENAI_PARAMS,
|
||||
'API_KEY': request.app.state.config.TTS_API_KEY,
|
||||
'ENGINE': request.app.state.config.TTS_ENGINE,
|
||||
'MODEL': request.app.state.config.TTS_MODEL,
|
||||
'VOICE': request.app.state.config.TTS_VOICE,
|
||||
'SPLIT_ON': request.app.state.config.TTS_SPLIT_ON,
|
||||
'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION,
|
||||
'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL,
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL,
|
||||
},
|
||||
'stt': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.STT_OPENAI_API_KEY,
|
||||
'ENGINE': request.app.state.config.STT_ENGINE,
|
||||
'MODEL': request.app.state.config.STT_MODEL,
|
||||
'SUPPORTED_CONTENT_TYPES': request.app.state.config.STT_SUPPORTED_CONTENT_TYPES,
|
||||
'ALLOWED_EXTENSIONS': request.app.state.config.STT_ALLOWED_EXTENSIONS,
|
||||
'WHISPER_MODEL': request.app.state.config.WHISPER_MODEL,
|
||||
'DEEPGRAM_API_KEY': request.app.state.config.DEEPGRAM_API_KEY,
|
||||
'AZURE_API_KEY': request.app.state.config.AUDIO_STT_AZURE_API_KEY,
|
||||
'AZURE_REGION': request.app.state.config.AUDIO_STT_AZURE_REGION,
|
||||
'AZURE_LOCALES': request.app.state.config.AUDIO_STT_AZURE_LOCALES,
|
||||
'AZURE_BASE_URL': request.app.state.config.AUDIO_STT_AZURE_BASE_URL,
|
||||
'AZURE_MAX_SPEAKERS': request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS,
|
||||
'MISTRAL_API_KEY': request.app.state.config.AUDIO_STT_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL,
|
||||
'MISTRAL_USE_CHAT_COMPLETIONS': request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS,
|
||||
},
|
||||
'tts': await get_config_values(TTS_CONFIG_KEYS),
|
||||
'stt': await get_config_values(STT_CONFIG_KEYS),
|
||||
}
|
||||
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_audio_config(request: Request, form_data: AudioConfigUpdateForm, user=Depends(get_admin_user)):
|
||||
# TTS settings
|
||||
request.app.state.config.TTS_OPENAI_API_BASE_URL = form_data.tts.OPENAI_API_BASE_URL
|
||||
request.app.state.config.TTS_OPENAI_API_KEY = form_data.tts.OPENAI_API_KEY
|
||||
request.app.state.config.TTS_OPENAI_PARAMS = form_data.tts.OPENAI_PARAMS
|
||||
request.app.state.config.TTS_API_KEY = form_data.tts.API_KEY
|
||||
request.app.state.config.TTS_ENGINE = form_data.tts.ENGINE
|
||||
request.app.state.config.TTS_MODEL = form_data.tts.MODEL
|
||||
request.app.state.config.TTS_VOICE = form_data.tts.VOICE
|
||||
request.app.state.config.TTS_SPLIT_ON = form_data.tts.SPLIT_ON
|
||||
request.app.state.config.TTS_AZURE_SPEECH_REGION = form_data.tts.AZURE_SPEECH_REGION
|
||||
request.app.state.config.TTS_AZURE_SPEECH_BASE_URL = form_data.tts.AZURE_SPEECH_BASE_URL
|
||||
request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT = form_data.tts.AZURE_SPEECH_OUTPUT_FORMAT
|
||||
request.app.state.config.TTS_MISTRAL_API_KEY = form_data.tts.MISTRAL_API_KEY
|
||||
request.app.state.config.TTS_MISTRAL_API_BASE_URL = form_data.tts.MISTRAL_API_BASE_URL
|
||||
await Config.upsert(
|
||||
{
|
||||
**config_updates(form_data.tts.model_dump(), TTS_CONFIG_KEYS),
|
||||
**config_updates(form_data.stt.model_dump(), STT_CONFIG_KEYS),
|
||||
}
|
||||
)
|
||||
|
||||
# STT settings
|
||||
request.app.state.config.STT_OPENAI_API_BASE_URL = form_data.stt.OPENAI_API_BASE_URL
|
||||
request.app.state.config.STT_OPENAI_API_KEY = form_data.stt.OPENAI_API_KEY
|
||||
request.app.state.config.STT_ENGINE = form_data.stt.ENGINE
|
||||
request.app.state.config.STT_MODEL = form_data.stt.MODEL
|
||||
request.app.state.config.STT_SUPPORTED_CONTENT_TYPES = form_data.stt.SUPPORTED_CONTENT_TYPES
|
||||
request.app.state.config.STT_ALLOWED_EXTENSIONS = form_data.stt.ALLOWED_EXTENSIONS
|
||||
request.app.state.config.WHISPER_MODEL = form_data.stt.WHISPER_MODEL
|
||||
request.app.state.config.DEEPGRAM_API_KEY = form_data.stt.DEEPGRAM_API_KEY
|
||||
request.app.state.config.AUDIO_STT_AZURE_API_KEY = form_data.stt.AZURE_API_KEY
|
||||
request.app.state.config.AUDIO_STT_AZURE_REGION = form_data.stt.AZURE_REGION
|
||||
request.app.state.config.AUDIO_STT_AZURE_LOCALES = form_data.stt.AZURE_LOCALES
|
||||
request.app.state.config.AUDIO_STT_AZURE_BASE_URL = form_data.stt.AZURE_BASE_URL
|
||||
request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS = form_data.stt.AZURE_MAX_SPEAKERS
|
||||
request.app.state.config.AUDIO_STT_MISTRAL_API_KEY = form_data.stt.MISTRAL_API_KEY
|
||||
request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL = form_data.stt.MISTRAL_API_BASE_URL
|
||||
request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS = form_data.stt.MISTRAL_USE_CHAT_COMPLETIONS
|
||||
|
||||
if request.app.state.config.STT_ENGINE == '':
|
||||
request.app.state.faster_whisper_model = set_faster_whisper_model(
|
||||
form_data.stt.WHISPER_MODEL, WHISPER_MODEL_AUTO_UPDATE
|
||||
if form_data.stt.ENGINE == '':
|
||||
request.app.state.faster_whisper_model = await asyncio.to_thread(
|
||||
set_faster_whisper_model, form_data.stt.WHISPER_MODEL, WHISPER_MODEL_AUTO_UPDATE
|
||||
)
|
||||
else:
|
||||
request.app.state.faster_whisper_model = None
|
||||
|
||||
return {
|
||||
'tts': {
|
||||
'ENGINE': request.app.state.config.TTS_ENGINE,
|
||||
'MODEL': request.app.state.config.TTS_MODEL,
|
||||
'VOICE': request.app.state.config.TTS_VOICE,
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.TTS_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.TTS_OPENAI_API_KEY,
|
||||
'OPENAI_PARAMS': request.app.state.config.TTS_OPENAI_PARAMS,
|
||||
'API_KEY': request.app.state.config.TTS_API_KEY,
|
||||
'SPLIT_ON': request.app.state.config.TTS_SPLIT_ON,
|
||||
'AZURE_SPEECH_REGION': request.app.state.config.TTS_AZURE_SPEECH_REGION,
|
||||
'AZURE_SPEECH_BASE_URL': request.app.state.config.TTS_AZURE_SPEECH_BASE_URL,
|
||||
'AZURE_SPEECH_OUTPUT_FORMAT': request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT,
|
||||
'MISTRAL_API_KEY': request.app.state.config.TTS_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.TTS_MISTRAL_API_BASE_URL,
|
||||
config = await get_audio_config(request, user)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_UPDATED,
|
||||
actor=user,
|
||||
subject_id='audio',
|
||||
data={
|
||||
'tts_engine': config.get('tts', {}).get('ENGINE'),
|
||||
'stt_engine': config.get('stt', {}).get('ENGINE'),
|
||||
},
|
||||
'stt': {
|
||||
'OPENAI_API_BASE_URL': request.app.state.config.STT_OPENAI_API_BASE_URL,
|
||||
'OPENAI_API_KEY': request.app.state.config.STT_OPENAI_API_KEY,
|
||||
'ENGINE': request.app.state.config.STT_ENGINE,
|
||||
'MODEL': request.app.state.config.STT_MODEL,
|
||||
'SUPPORTED_CONTENT_TYPES': request.app.state.config.STT_SUPPORTED_CONTENT_TYPES,
|
||||
'ALLOWED_EXTENSIONS': request.app.state.config.STT_ALLOWED_EXTENSIONS,
|
||||
'WHISPER_MODEL': request.app.state.config.WHISPER_MODEL,
|
||||
'DEEPGRAM_API_KEY': request.app.state.config.DEEPGRAM_API_KEY,
|
||||
'AZURE_API_KEY': request.app.state.config.AUDIO_STT_AZURE_API_KEY,
|
||||
'AZURE_REGION': request.app.state.config.AUDIO_STT_AZURE_REGION,
|
||||
'AZURE_LOCALES': request.app.state.config.AUDIO_STT_AZURE_LOCALES,
|
||||
'AZURE_BASE_URL': request.app.state.config.AUDIO_STT_AZURE_BASE_URL,
|
||||
'AZURE_MAX_SPEAKERS': request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS,
|
||||
'MISTRAL_API_KEY': request.app.state.config.AUDIO_STT_MISTRAL_API_KEY,
|
||||
'MISTRAL_API_BASE_URL': request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL,
|
||||
'MISTRAL_USE_CHAT_COMPLETIONS': request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS,
|
||||
},
|
||||
}
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
def load_speech_pipeline(request):
|
||||
@@ -388,14 +356,16 @@ async def _write_tts_cache(
|
||||
|
||||
async def _tts_openai(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via an OpenAI-compatible TTS endpoint."""
|
||||
payload['model'] = request.app.state.config.TTS_MODEL
|
||||
payload['model'] = await Config.get('audio.tts.model')
|
||||
if not payload.get('voice'):
|
||||
payload['voice'] = request.app.state.config.TTS_VOICE
|
||||
payload = {**payload, **(request.app.state.config.TTS_OPENAI_PARAMS or {})}
|
||||
payload['voice'] = await Config.get('audio.tts.voice')
|
||||
payload = {**payload, **(await Config.get('audio.tts.openai.params') or {})}
|
||||
api_key = await Config.get('audio.tts.openai.api_key')
|
||||
api_base_url = await Config.get('audio.tts.openai.api_base_url')
|
||||
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': f'Bearer {request.app.state.config.TTS_OPENAI_API_KEY}',
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
}
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
@@ -404,7 +374,7 @@ async def _tts_openai(request, payload, file_path, file_body_path, user):
|
||||
try:
|
||||
session = await get_session()
|
||||
r = await session.post(
|
||||
url=f'{request.app.state.config.TTS_OPENAI_API_BASE_URL}/audio/speech',
|
||||
url=f'{api_base_url}/audio/speech',
|
||||
json=payload,
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -429,8 +399,12 @@ async def _tts_openai(request, payload, file_path, file_body_path, user):
|
||||
|
||||
async def _tts_elevenlabs(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via the ElevenLabs TTS API."""
|
||||
voice_id = payload.get('voice', '')
|
||||
if voice_id not in await get_available_voices(request):
|
||||
voice_id = (payload.get('voice') or '').strip()
|
||||
if not voice_id:
|
||||
raise HTTPException(status_code=400, detail='Invalid voice id')
|
||||
|
||||
available_voices = await get_available_voices(request)
|
||||
if available_voices and voice_id not in available_voices:
|
||||
raise HTTPException(status_code=400, detail='Invalid voice id')
|
||||
|
||||
r = None
|
||||
@@ -440,13 +414,13 @@ async def _tts_elevenlabs(request, payload, file_path, file_body_path, user):
|
||||
f'{ELEVENLABS_API_BASE_URL}/v1/text-to-speech/{voice_id}',
|
||||
json={
|
||||
'text': payload['input'],
|
||||
'model_id': request.app.state.config.TTS_MODEL,
|
||||
'model_id': await Config.get('audio.tts.model'),
|
||||
'voice_settings': {'stability': 0.5, 'similarity_boost': 0.5},
|
||||
},
|
||||
headers={
|
||||
'Accept': 'audio/mpeg',
|
||||
'Content-Type': 'application/json',
|
||||
'xi-api-key': request.app.state.config.TTS_API_KEY,
|
||||
'xi-api-key': await Config.get('audio.tts.api_key'),
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
@@ -460,15 +434,15 @@ async def _tts_elevenlabs(request, payload, file_path, file_body_path, user):
|
||||
|
||||
async def _tts_azure(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via Azure Cognitive Services TTS."""
|
||||
az_region = request.app.state.config.TTS_AZURE_SPEECH_REGION or 'eastus'
|
||||
az_base = request.app.state.config.TTS_AZURE_SPEECH_BASE_URL
|
||||
language = payload.get('voice') or request.app.state.config.TTS_VOICE
|
||||
az_region = await Config.get('audio.tts.azure.speech_region') or 'eastus'
|
||||
az_base = await Config.get('audio.tts.azure.speech_base_url')
|
||||
language = payload.get('voice') or await Config.get('audio.tts.voice')
|
||||
locale = '-'.join(language.split('-')[:2])
|
||||
output_format = request.app.state.config.TTS_AZURE_SPEECH_OUTPUT_FORMAT
|
||||
output_format = await Config.get('audio.tts.azure.speech_output_format')
|
||||
|
||||
ssml = (
|
||||
f'<speak version="1.0" xmlns="http://www.w3.org/2001/10/synthesis" xml:lang="{locale}">'
|
||||
f'<voice name="{language}">{html.escape(payload["input"])}</voice>'
|
||||
f'<speak version="1.0" xmlns="http://www.w3.org/2001/10/synthesis" xml:lang="{html.escape(locale)}">'
|
||||
f'<voice name="{html.escape(language)}">{html.escape(payload["input"])}</voice>'
|
||||
f'</speak>'
|
||||
)
|
||||
|
||||
@@ -478,7 +452,7 @@ async def _tts_azure(request, payload, file_path, file_body_path, user):
|
||||
async with session.post(
|
||||
(az_base or f'https://{az_region}.tts.speech.microsoft.com') + '/cognitiveservices/v1',
|
||||
headers={
|
||||
'Ocp-Apim-Subscription-Key': request.app.state.config.TTS_API_KEY,
|
||||
'Ocp-Apim-Subscription-Key': await Config.get('audio.tts.api_key'),
|
||||
'Content-Type': 'application/ssml+xml',
|
||||
'X-Microsoft-OutputFormat': output_format,
|
||||
},
|
||||
@@ -498,10 +472,10 @@ async def _tts_transformers(request, payload, file_path, file_body_path, user):
|
||||
import soundfile as sf
|
||||
import torch
|
||||
|
||||
load_speech_pipeline(request)
|
||||
await asyncio.to_thread(load_speech_pipeline, request)
|
||||
|
||||
embeddings = request.app.state.speech_speaker_embeddings_dataset
|
||||
model_name = request.app.state.config.TTS_MODEL
|
||||
model_name = await Config.get('audio.tts.model')
|
||||
|
||||
idx = 6799
|
||||
try:
|
||||
@@ -529,8 +503,8 @@ async def _tts_transformers(request, payload, file_path, file_body_path, user):
|
||||
|
||||
async def _tts_mistral(request, payload, file_path, file_body_path, user):
|
||||
"""Generate speech via the Mistral TTS API."""
|
||||
api_key = request.app.state.config.TTS_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
api_key = await Config.get('audio.tts.mistral.api_key')
|
||||
api_base_url = await Config.get('audio.tts.mistral.api_base_url') or 'https://api.mistral.ai/v1'
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=400, detail='Mistral API key is required for Mistral TTS')
|
||||
@@ -542,7 +516,7 @@ async def _tts_mistral(request, payload, file_path, file_body_path, user):
|
||||
url=f'{api_base_url}/audio/speech',
|
||||
json={
|
||||
'input': payload.get('input', ''), # text to synthesize
|
||||
'model': request.app.state.config.TTS_MODEL or 'voxtral-mini-tts-2603',
|
||||
'model': await Config.get('audio.tts.model') or 'voxtral-mini-tts-2603',
|
||||
'voice_id': payload.get('voice', ''),
|
||||
'response_format': 'mp3',
|
||||
},
|
||||
@@ -578,16 +552,14 @@ _TTS_ENGINES = {
|
||||
|
||||
@router.post('/speech')
|
||||
async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
engine = request.app.state.config.TTS_ENGINE
|
||||
engine = await Config.get('audio.tts.engine')
|
||||
if engine == '':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.tts', request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'chat.tts', await Config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -595,7 +567,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
body = await request.body()
|
||||
name = hashlib.sha256(
|
||||
body + str(engine).encode('utf-8') + str(request.app.state.config.TTS_MODEL).encode('utf-8')
|
||||
body + str(engine).encode('utf-8') + str(await Config.get('audio.tts.model')).encode('utf-8')
|
||||
).hexdigest()
|
||||
|
||||
file_path = SPEECH_CACHE_DIR.joinpath(f'{name}.mp3')
|
||||
@@ -603,6 +575,13 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
# Return cached result if available
|
||||
if file_path.is_file():
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUDIO_SPEECH_REQUESTED,
|
||||
actor=user,
|
||||
subject_id=name,
|
||||
data={'engine': engine, 'cached': True},
|
||||
)
|
||||
return FileResponse(file_path)
|
||||
|
||||
try:
|
||||
@@ -615,12 +594,27 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
if handler is None:
|
||||
raise HTTPException(status_code=400, detail=f'Unsupported TTS engine: {engine}')
|
||||
|
||||
return await handler(request, payload, file_path, file_body_path, user)
|
||||
response = await handler(request, payload, file_path, file_body_path, user)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUDIO_SPEECH_REQUESTED,
|
||||
actor=user,
|
||||
subject_id=name,
|
||||
data={
|
||||
'engine': engine,
|
||||
'model': payload.get('model'),
|
||||
'input_preview': str(payload.get('input', ''))[:300],
|
||||
'cached': False,
|
||||
},
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
async def _transcribe_whisper(request, file_path, languages, file_dir, id):
|
||||
if request.app.state.faster_whisper_model is None:
|
||||
request.app.state.faster_whisper_model = set_faster_whisper_model(request.app.state.config.WHISPER_MODEL)
|
||||
request.app.state.faster_whisper_model = await asyncio.to_thread(
|
||||
set_faster_whisper_model, await Config.get('audio.stt.whisper_model')
|
||||
)
|
||||
|
||||
model = request.app.state.faster_whisper_model
|
||||
|
||||
@@ -651,11 +645,13 @@ async def _transcribe_openai(request, file_path, filename, languages, file_dir,
|
||||
try:
|
||||
session = await get_session()
|
||||
for language in languages:
|
||||
payload = {'model': request.app.state.config.STT_MODEL}
|
||||
payload = {'model': await Config.get('audio.stt.model')}
|
||||
if language:
|
||||
payload['language'] = language
|
||||
api_key = await Config.get('audio.stt.openai.api_key')
|
||||
api_base_url = await Config.get('audio.stt.openai.api_base_url')
|
||||
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.STT_OPENAI_API_KEY}'}
|
||||
headers = {'Authorization': f'Bearer {api_key}'}
|
||||
if user and ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
@@ -665,7 +661,7 @@ async def _transcribe_openai(request, file_path, filename, languages, file_dir,
|
||||
form_data.add_field('file', open(file_path, 'rb'), filename=filename)
|
||||
|
||||
r = await session.post(
|
||||
url=f'{request.app.state.config.STT_OPENAI_API_BASE_URL}/audio/transcriptions',
|
||||
url=f'{api_base_url}/audio/transcriptions',
|
||||
headers=headers,
|
||||
data=form_data,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -699,8 +695,8 @@ async def _transcribe_deepgram(request, file_path, languages, file_dir, id):
|
||||
async with aiofiles.open(file_path, 'rb') as f:
|
||||
audio_bytes = await f.read()
|
||||
|
||||
api_key = request.app.state.config.DEEPGRAM_API_KEY
|
||||
stt_model = request.app.state.config.STT_MODEL
|
||||
api_key = await Config.get('audio.stt.deepgram.api_key')
|
||||
stt_model = await Config.get('audio.stt.model')
|
||||
|
||||
r = None
|
||||
try:
|
||||
@@ -767,11 +763,11 @@ async def _transcribe_azure(request, file_path, filename, file_dir, id):
|
||||
detail=f'File size ({audio_size // (1024 * 1024)}MB) exceeds Azure limit of {AZURE_MAX_FILE_SIZE_MB}MB',
|
||||
)
|
||||
|
||||
api_key = request.app.state.config.AUDIO_STT_AZURE_API_KEY
|
||||
region = request.app.state.config.AUDIO_STT_AZURE_REGION or 'eastus'
|
||||
locale_str = request.app.state.config.AUDIO_STT_AZURE_LOCALES
|
||||
base_url = request.app.state.config.AUDIO_STT_AZURE_BASE_URL
|
||||
max_speakers = request.app.state.config.AUDIO_STT_AZURE_MAX_SPEAKERS or 3
|
||||
api_key = await Config.get('audio.stt.azure.api_key')
|
||||
region = await Config.get('audio.stt.azure.region') or 'eastus'
|
||||
locale_str = await Config.get('audio.stt.azure.locales')
|
||||
base_url = await Config.get('audio.stt.azure.base_url')
|
||||
max_speakers = await Config.get('audio.stt.azure.max_speakers') or 3
|
||||
|
||||
# Default to a broad set of locales when none are configured
|
||||
if len(locale_str) < 2:
|
||||
@@ -881,16 +877,16 @@ async def transcription_handler(request, file_path, metadata, user=None):
|
||||
None, # Always fallback to None in case transcription fails
|
||||
]
|
||||
|
||||
if request.app.state.config.STT_ENGINE == '':
|
||||
if await Config.get('audio.stt.engine') == '':
|
||||
return await _transcribe_whisper(request, file_path, languages, file_dir, id)
|
||||
elif request.app.state.config.STT_ENGINE == 'openai':
|
||||
elif await Config.get('audio.stt.engine') == 'openai':
|
||||
return await _transcribe_openai(request, file_path, filename, languages, file_dir, id, user)
|
||||
elif request.app.state.config.STT_ENGINE == 'deepgram':
|
||||
elif await Config.get('audio.stt.engine') == 'deepgram':
|
||||
return await _transcribe_deepgram(request, file_path, languages, file_dir, id)
|
||||
elif request.app.state.config.STT_ENGINE == 'azure':
|
||||
elif await Config.get('audio.stt.engine') == 'azure':
|
||||
return await _transcribe_azure(request, file_path, filename, file_dir, id)
|
||||
|
||||
elif request.app.state.config.STT_ENGINE == 'mistral':
|
||||
elif await Config.get('audio.stt.engine') == 'mistral':
|
||||
return await _transcribe_mistral(request, file_path, filename, metadata, file_dir, id)
|
||||
|
||||
|
||||
@@ -903,16 +899,16 @@ async def _transcribe_mistral(request, file_path, filename, metadata, file_dir,
|
||||
if file_size > MAX_FILE_SIZE:
|
||||
raise HTTPException(status_code=400, detail=f'File size exceeds limit of {MAX_FILE_SIZE_MB}MB')
|
||||
|
||||
api_key = request.app.state.config.AUDIO_STT_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.AUDIO_STT_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
use_chat_completions = request.app.state.config.AUDIO_STT_MISTRAL_USE_CHAT_COMPLETIONS
|
||||
api_key = await Config.get('audio.stt.mistral.api_key')
|
||||
api_base_url = await Config.get('audio.stt.mistral.api_base_url') or 'https://api.mistral.ai/v1'
|
||||
use_chat_completions = await Config.get('audio.stt.mistral.use_chat_completions')
|
||||
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=400, detail='Mistral API key is required for Mistral STT')
|
||||
|
||||
r = None
|
||||
try:
|
||||
model = request.app.state.config.STT_MODEL or 'voxtral-mini-latest'
|
||||
model = await Config.get('audio.stt.model') or 'voxtral-mini-latest'
|
||||
log.info(
|
||||
f'Mistral STT - model: {model}, method: {"chat_completions" if use_chat_completions else "transcriptions"}'
|
||||
)
|
||||
@@ -1056,7 +1052,7 @@ async def transcribe(request: Request, file_path: str, metadata: Optional[dict]
|
||||
log.exception(e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error processing audio file'),
|
||||
)
|
||||
|
||||
results = []
|
||||
@@ -1153,15 +1149,13 @@ async def transcription(
|
||||
language: Optional[str] = Form(None),
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.stt', request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'chat.stt', await Config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
log.info(f'file.content_type: {file.content_type}')
|
||||
stt_supported_content_types = getattr(request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', [])
|
||||
stt_supported_content_types = await Config.get('audio.stt.supported_content_types', [])
|
||||
|
||||
if not strict_match_mime_type(stt_supported_content_types, file.content_type):
|
||||
raise HTTPException(
|
||||
@@ -1173,7 +1167,7 @@ async def transcription(
|
||||
safe_name = os.path.basename(file.filename) if file.filename else ''
|
||||
ext = safe_name.rsplit('.', 1)[-1].lower() if '.' in safe_name else ''
|
||||
|
||||
allowed_extensions = getattr(request.app.state.config, 'STT_ALLOWED_EXTENSIONS', [])
|
||||
allowed_extensions = await Config.get('audio.stt.allowed_extensions', [])
|
||||
if allowed_extensions and ext not in allowed_extensions:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
@@ -1193,8 +1187,12 @@ async def transcription(
|
||||
if not os.path.realpath(file_path).startswith(os.path.realpath(file_dir)):
|
||||
raise ValueError('Invalid file path detected')
|
||||
|
||||
with open(file_path, 'wb') as f:
|
||||
f.write(contents)
|
||||
def _write_upload():
|
||||
with open(file_path, 'wb') as f:
|
||||
f.write(contents)
|
||||
|
||||
# Audio uploads can be large; write to disk off the event loop.
|
||||
await asyncio.to_thread(_write_upload)
|
||||
|
||||
try:
|
||||
metadata = None
|
||||
@@ -1204,6 +1202,17 @@ async def transcription(
|
||||
|
||||
result = await transcribe(request, file_path, metadata, user)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUDIO_TRANSCRIPTION_REQUESTED,
|
||||
actor=user,
|
||||
subject_id=str(id),
|
||||
data={
|
||||
'filename': safe_name,
|
||||
'content_type': file.content_type,
|
||||
'language': language,
|
||||
},
|
||||
)
|
||||
return {
|
||||
**result,
|
||||
'filename': os.path.basename(file_path),
|
||||
@@ -1233,11 +1242,11 @@ async def transcription(
|
||||
async def get_available_models(request: Request) -> list[dict]:
|
||||
"""Return the list of available TTS models for the configured engine."""
|
||||
available_models = []
|
||||
engine = request.app.state.config.TTS_ENGINE
|
||||
engine = await Config.get('audio.tts.engine')
|
||||
_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST)
|
||||
|
||||
if engine == 'openai':
|
||||
base_url = request.app.state.config.TTS_OPENAI_API_BASE_URL
|
||||
base_url = await Config.get('audio.tts.openai.api_base_url')
|
||||
if not base_url.startswith('https://api.openai.com'):
|
||||
session = await get_session()
|
||||
try:
|
||||
@@ -1272,7 +1281,7 @@ async def get_available_models(request: Request) -> list[dict]:
|
||||
async with session.get(
|
||||
f'{ELEVENLABS_API_BASE_URL}/v1/models',
|
||||
headers={
|
||||
'xi-api-key': request.app.state.config.TTS_API_KEY,
|
||||
'xi-api-key': await Config.get('audio.tts.api_key'),
|
||||
'Content-Type': 'application/json',
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -1307,11 +1316,11 @@ _OPENAI_DEFAULT_VOICES = {
|
||||
|
||||
async def get_available_voices(request) -> dict:
|
||||
"""Return ``{voice_id: voice_name}`` for the configured TTS engine."""
|
||||
engine = request.app.state.config.TTS_ENGINE
|
||||
engine = await Config.get('audio.tts.engine')
|
||||
_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST)
|
||||
|
||||
if engine == 'openai':
|
||||
base_url = request.app.state.config.TTS_OPENAI_API_BASE_URL
|
||||
base_url = await Config.get('audio.tts.openai.api_base_url')
|
||||
if not base_url.startswith('https://api.openai.com'):
|
||||
try:
|
||||
session = await get_session()
|
||||
@@ -1334,7 +1343,7 @@ async def get_available_voices(request) -> dict:
|
||||
async with session.get(
|
||||
f'{ELEVENLABS_API_BASE_URL}/v1/voices',
|
||||
headers={
|
||||
'xi-api-key': request.app.state.config.TTS_API_KEY,
|
||||
'xi-api-key': await Config.get('audio.tts.api_key'),
|
||||
'Content-Type': 'application/json',
|
||||
},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -1344,19 +1353,19 @@ async def get_available_voices(request) -> dict:
|
||||
voices_data = await resp.json()
|
||||
return {v['voice_id']: v['name'] for v in voices_data.get('voices', [])}
|
||||
except Exception as e:
|
||||
log.error(f'Error fetching ElevenLabs voices: {e}')
|
||||
log.warning(f'Error fetching ElevenLabs voices: {e}')
|
||||
return {}
|
||||
|
||||
if engine == 'azure':
|
||||
try:
|
||||
region = request.app.state.config.TTS_AZURE_SPEECH_REGION
|
||||
base_url = request.app.state.config.TTS_AZURE_SPEECH_BASE_URL
|
||||
region = await Config.get('audio.tts.azure.speech_region')
|
||||
base_url = await Config.get('audio.tts.azure.speech_base_url')
|
||||
url = (base_url or f'https://{region}.tts.speech.microsoft.com') + '/cognitiveservices/voices/list'
|
||||
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url,
|
||||
headers={'Ocp-Apim-Subscription-Key': request.app.state.config.TTS_API_KEY},
|
||||
headers={'Ocp-Apim-Subscription-Key': await Config.get('audio.tts.api_key')},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
timeout=_timeout,
|
||||
) as resp:
|
||||
@@ -1368,8 +1377,8 @@ async def get_available_voices(request) -> dict:
|
||||
return {}
|
||||
|
||||
if engine == 'mistral':
|
||||
api_key = request.app.state.config.TTS_MISTRAL_API_KEY
|
||||
api_base_url = request.app.state.config.TTS_MISTRAL_API_BASE_URL or 'https://api.mistral.ai/v1'
|
||||
api_key = await Config.get('audio.tts.mistral.api_key')
|
||||
api_base_url = await Config.get('audio.tts.mistral.api_base_url') or 'https://api.mistral.ai/v1'
|
||||
if api_key:
|
||||
try:
|
||||
session = await get_session()
|
||||
|
||||
+402
-244
File diff suppressed because it is too large
Load Diff
@@ -4,6 +4,7 @@ from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.automations import (
|
||||
AutomationForm,
|
||||
@@ -14,6 +15,7 @@ from open_webui.models.automations import (
|
||||
AutomationRuns,
|
||||
Automations,
|
||||
)
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.automations import (
|
||||
@@ -38,13 +40,14 @@ PAGE_ITEM_COUNT = 30
|
||||
|
||||
|
||||
async def check_automations_permission(request, user):
|
||||
if not request.app.state.config.ENABLE_AUTOMATIONS:
|
||||
config = await Config.get_many('automations.enable', 'user.permissions')
|
||||
if not config.get('automations.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.automations', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.automations', config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -72,7 +75,7 @@ async def check_automation_limits(request, user, rrule_str: str, db, is_create:
|
||||
|
||||
# Max count (create only)
|
||||
if is_create:
|
||||
max_count = request.app.state.config.AUTOMATION_MAX_COUNT
|
||||
max_count = await Config.get('automations.max_count')
|
||||
if max_count:
|
||||
max_count = int(max_count)
|
||||
if max_count > 0 and await Automations.count_by_user(user.id, db=db) >= max_count:
|
||||
@@ -82,7 +85,7 @@ async def check_automation_limits(request, user, rrule_str: str, db, is_create:
|
||||
)
|
||||
|
||||
# Min interval (create + update)
|
||||
min_interval = request.app.state.config.AUTOMATION_MIN_INTERVAL
|
||||
min_interval = await Config.get('automations.min_interval')
|
||||
if min_interval:
|
||||
min_interval = int(min_interval)
|
||||
if min_interval > 0:
|
||||
@@ -173,7 +176,15 @@ async def create_new_automation(
|
||||
|
||||
tz = user.timezone
|
||||
automation = await Automations.insert(user.id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
|
||||
return await enrich_automation(automation, db, tz=tz)
|
||||
response = await enrich_automation(automation, db, tz=tz)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUTOMATION_CREATED,
|
||||
actor=user,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name, 'is_active': automation.is_active},
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
############################
|
||||
@@ -223,7 +234,15 @@ async def update_automation_by_id(
|
||||
|
||||
tz = user.timezone
|
||||
updated = await Automations.update_by_id(id, form_data, next_run_ns(form_data.data.rrule, tz=tz), db=db)
|
||||
return await enrich_automation(updated, db, tz=tz)
|
||||
response = await enrich_automation(updated, db, tz=tz)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUTOMATION_UPDATED,
|
||||
actor=user,
|
||||
subject_id=updated.id,
|
||||
data={'name': updated.name, 'is_active': updated.is_active},
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
############################
|
||||
@@ -242,7 +261,16 @@ async def toggle_automation_by_id(
|
||||
automation = await Automations.get_by_id(id, db=db)
|
||||
check_automation_access(automation, user)
|
||||
toggled = await Automations.toggle(id, next_run_ns(automation.data['rrule'], tz=user.timezone), db=db)
|
||||
return await enrich_automation(toggled, db, tz=user.timezone)
|
||||
response = await enrich_automation(toggled, db, tz=user.timezone)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUTOMATION_ENABLED if toggled.is_active else EVENTS.AUTOMATION_DISABLED,
|
||||
actor=user,
|
||||
subject_id=toggled.id,
|
||||
subject_type='automation',
|
||||
data={'name': toggled.name},
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
############################
|
||||
@@ -261,6 +289,13 @@ async def run_automation_by_id(
|
||||
automation = await Automations.get_by_id(id, db=db)
|
||||
check_automation_access(automation, user)
|
||||
asyncio.create_task(execute_automation(request.app, automation))
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUTOMATION_RUN_STARTED,
|
||||
actor=user,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name},
|
||||
)
|
||||
return await enrich_automation(automation, db, tz=user.timezone)
|
||||
|
||||
|
||||
@@ -280,7 +315,16 @@ async def delete_automation_by_id(
|
||||
automation = await Automations.get_by_id(id, db=db)
|
||||
check_automation_access(automation, user)
|
||||
await AutomationRuns.delete_by_automation(id, db=db)
|
||||
return await Automations.delete(id, db=db)
|
||||
result = await Automations.delete(id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUTOMATION_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': automation.name},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
############################
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.calendar import (
|
||||
CalendarEventAttendees,
|
||||
@@ -19,6 +20,7 @@ from open_webui.models.calendar import (
|
||||
CalendarUpdateForm,
|
||||
RSVPForm,
|
||||
)
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import UserModel
|
||||
from open_webui.utils.access_control import filter_allowed_access_grants, has_permission
|
||||
@@ -34,14 +36,13 @@ SCHEDULED_TASKS_CALENDAR_ID = '__scheduled_tasks__'
|
||||
|
||||
async def check_calendar_permission(request: Request, user):
|
||||
"""Check global feature flag AND per-user permission for calendar access."""
|
||||
if not request.app.state.config.ENABLE_CALENDAR:
|
||||
config = await Config.get_many('calendar.enable', 'user.permissions')
|
||||
if not config.get('calendar.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.calendar', request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.calendar', config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
@@ -50,11 +51,12 @@ async def check_calendar_permission(request: Request, user):
|
||||
|
||||
async def _user_has_automations(request: Request, user) -> bool:
|
||||
"""Check if automations feature is available to this user."""
|
||||
if not getattr(request.app.state.config, 'ENABLE_AUTOMATIONS', False):
|
||||
config = await Config.get_many('automations.enable', 'user.permissions')
|
||||
if not config.get('automations.enable', False):
|
||||
return False
|
||||
if user.role == 'admin':
|
||||
return True
|
||||
return await has_permission(user.id, 'features.automations', request.app.state.config.USER_PERMISSIONS)
|
||||
return await has_permission(user.id, 'features.automations', config.get('user.permissions'))
|
||||
|
||||
|
||||
async def _check_calendar_access(calendar_id: str, user: UserModel, permission: str = 'write') -> CalendarModel:
|
||||
@@ -116,13 +118,21 @@ async def create_calendar(request: Request, form_data: CalendarForm, user: UserM
|
||||
# could create a calendar with `principal_id='*' permission='read'|'write'`,
|
||||
# making their events readable or writable by any other verified user.
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
'sharing.public_calendars',
|
||||
)
|
||||
return await Calendars.insert_new_calendar(user.id, form_data)
|
||||
calendar = await Calendars.insert_new_calendar(user.id, form_data)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_CREATED,
|
||||
actor=user,
|
||||
subject_id=calendar.id,
|
||||
data={'name': calendar.name},
|
||||
)
|
||||
return calendar
|
||||
|
||||
|
||||
####################
|
||||
@@ -263,7 +273,15 @@ async def get_events(
|
||||
async def create_event(request: Request, form_data: CalendarEventForm, user: UserModel = Depends(get_verified_user)):
|
||||
await check_calendar_permission(request, user)
|
||||
await _check_calendar_access(form_data.calendar_id, user, 'write')
|
||||
return await CalendarEvents.insert_new_event(user.id, form_data)
|
||||
event = await CalendarEvents.insert_new_event(user.id, form_data)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_EVENT_CREATED,
|
||||
actor=user,
|
||||
subject_id=event.id,
|
||||
data={'calendar_id': event.calendar_id, 'title': event.title},
|
||||
)
|
||||
return event
|
||||
|
||||
|
||||
@router.get('/events/search', response_model=CalendarEventListResponse)
|
||||
@@ -310,6 +328,13 @@ async def update_event(
|
||||
updated = await CalendarEvents.update_event_by_id(event_id, form_data)
|
||||
if not updated:
|
||||
raise HTTPException(status_code=500, detail='Failed to update')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_EVENT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=updated.id,
|
||||
data={'calendar_id': updated.calendar_id, 'title': updated.title},
|
||||
)
|
||||
return updated
|
||||
|
||||
|
||||
@@ -325,6 +350,13 @@ async def delete_event(request: Request, event_id: str, user: UserModel = Depend
|
||||
result = await CalendarEvents.delete_event_by_id(event_id)
|
||||
if not result:
|
||||
raise HTTPException(status_code=500, detail='Failed to delete')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_EVENT_DELETED,
|
||||
actor=user,
|
||||
subject_id=event_id,
|
||||
data={'calendar_id': event.calendar_id, 'title': event.title},
|
||||
)
|
||||
return {'status': True}
|
||||
|
||||
|
||||
@@ -340,6 +372,13 @@ async def rsvp_event(
|
||||
result = await CalendarEventAttendees.update_rsvp(event_id, user.id, form_data.status)
|
||||
if not result:
|
||||
raise HTTPException(status_code=404, detail='Not an attendee of this event')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_EVENT_RSVP_UPDATED,
|
||||
actor=user,
|
||||
subject_id=event_id,
|
||||
data={'status': result.status},
|
||||
)
|
||||
return {'status': True, 'rsvp': result.status}
|
||||
|
||||
|
||||
@@ -373,7 +412,7 @@ async def update_calendar(
|
||||
# publicly readable/writable without the corresponding sharing permission.
|
||||
if form_data.access_grants is not None:
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -383,6 +422,13 @@ async def update_calendar(
|
||||
updated = await Calendars.update_calendar_by_id(calendar_id, form_data)
|
||||
if not updated:
|
||||
raise HTTPException(status_code=500, detail='Failed to update')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_UPDATED,
|
||||
actor=user,
|
||||
subject_id=updated.id,
|
||||
data={'name': updated.name},
|
||||
)
|
||||
return updated
|
||||
|
||||
|
||||
@@ -407,6 +453,13 @@ async def delete_calendar(request: Request, calendar_id: str, user: UserModel =
|
||||
result = await Calendars.delete_calendar_by_id(calendar_id)
|
||||
if not result:
|
||||
raise HTTPException(status_code=500, detail='Failed to delete')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_DELETED,
|
||||
actor=user,
|
||||
subject_id=calendar_id,
|
||||
data={'name': cal.name},
|
||||
)
|
||||
return {'status': True}
|
||||
|
||||
|
||||
@@ -416,4 +469,11 @@ async def set_default_calendar(request: Request, calendar_id: str, user: UserMod
|
||||
cal = await Calendars.set_default_calendar(user.id, calendar_id)
|
||||
if not cal:
|
||||
raise HTTPException(status_code=404, detail='Calendar not found')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CALENDAR_DEFAULT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=cal.id,
|
||||
data={'name': cal.name},
|
||||
)
|
||||
return cal
|
||||
|
||||
@@ -8,9 +8,11 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request,
|
||||
from fastapi.responses import FileResponse, Response, StreamingResponse
|
||||
from open_webui.config import ENABLE_ADMIN_CHAT_ACCESS, ENABLE_ADMIN_EXPORT
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import STATIC_DIR
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants, has_public_read_access_grant, has_public_write_access_grant
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.channels import (
|
||||
ChannelForm,
|
||||
ChannelModel,
|
||||
@@ -31,9 +33,7 @@ from open_webui.models.messages import (
|
||||
from open_webui.models.users import (
|
||||
UserIdNameResponse,
|
||||
UserIdNameStatusResponse,
|
||||
UserListResponse,
|
||||
UserModel,
|
||||
UserModelResponse,
|
||||
UserNameResponse,
|
||||
Users,
|
||||
)
|
||||
@@ -125,7 +125,7 @@ def get_channel_permitted_group_and_user_ids(
|
||||
|
||||
async def check_channels_access(request: Request, user: Optional[UserModel] = None):
|
||||
"""Dependency to ensure channels are globally enabled."""
|
||||
if not request.app.state.config.ENABLE_CHANNELS:
|
||||
if not await Config.get('channels.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Channels'),
|
||||
@@ -133,7 +133,7 @@ async def check_channels_access(request: Request, user: Optional[UserModel] = No
|
||||
|
||||
if user:
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.channels', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.channels', await Config.get('user.permissions')
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -294,7 +294,7 @@ async def create_new_channel(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -316,6 +316,13 @@ async def create_new_channel(
|
||||
await enter_room_for_users(f'channel:{existing_channel.id}', participant_ids)
|
||||
|
||||
await Channels.update_member_active_status(existing_channel.id, user.id, True, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_MEMBER_ACTIVE_UPDATED,
|
||||
actor=user,
|
||||
subject_id=existing_channel.id,
|
||||
data={'is_active': True},
|
||||
)
|
||||
return ChannelModel(**existing_channel.model_dump())
|
||||
|
||||
channel = await Channels.insert_new_channel(form_data, user.id, db=db)
|
||||
@@ -330,6 +337,13 @@ async def create_new_channel(
|
||||
)
|
||||
await enter_room_for_users(f'channel:{channel.id}', participant_ids)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_CREATED,
|
||||
actor=user,
|
||||
subject_id=channel.id,
|
||||
data={'type': channel.type, 'name': channel.name},
|
||||
)
|
||||
return ChannelModel(**channel.model_dump())
|
||||
else:
|
||||
raise Exception('Error creating channel')
|
||||
@@ -440,7 +454,40 @@ async def get_channel_by_id(
|
||||
PAGE_ITEM_COUNT = 30
|
||||
|
||||
|
||||
@router.get('/{id}/members', response_model=UserListResponse)
|
||||
class ChannelMemberResponse(BaseModel):
|
||||
id: str
|
||||
email: str
|
||||
name: str
|
||||
role: str
|
||||
profile_image_url: str | None = None
|
||||
presence_state: str | None = None
|
||||
status_emoji: str | None = None
|
||||
status_message: str | None = None
|
||||
status_expires_at: int | None = None
|
||||
is_active: bool = False
|
||||
|
||||
|
||||
class ChannelMemberListResponse(BaseModel):
|
||||
users: list[ChannelMemberResponse]
|
||||
total: int
|
||||
|
||||
|
||||
def serialize_channel_member(user: UserModel) -> ChannelMemberResponse:
|
||||
return ChannelMemberResponse(
|
||||
id=user.id,
|
||||
email=user.email,
|
||||
name=user.name,
|
||||
role=user.role,
|
||||
profile_image_url=user.profile_image_url,
|
||||
presence_state=user.presence_state,
|
||||
status_emoji=user.status_emoji,
|
||||
status_message=user.status_message,
|
||||
status_expires_at=user.status_expires_at,
|
||||
is_active=Users.is_active(user),
|
||||
)
|
||||
|
||||
|
||||
@router.get('/{id}/members', response_model=ChannelMemberListResponse)
|
||||
async def get_channel_members_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
@@ -475,7 +522,7 @@ async def get_channel_members_by_id(
|
||||
total = len(fetched_users)
|
||||
|
||||
return {
|
||||
'users': [UserModelResponse(**u.model_dump(), is_active=Users.is_active(u)) for u in fetched_users],
|
||||
'users': [serialize_channel_member(u) for u in fetched_users],
|
||||
'total': total,
|
||||
}
|
||||
else:
|
||||
@@ -503,7 +550,7 @@ async def get_channel_members_by_id(
|
||||
total = result['total']
|
||||
|
||||
return {
|
||||
'users': [UserModelResponse(**u.model_dump(), is_active=Users.is_active(u)) for u in fetched_users],
|
||||
'users': [serialize_channel_member(u) for u in fetched_users],
|
||||
'total': total,
|
||||
}
|
||||
|
||||
@@ -534,6 +581,13 @@ async def update_is_active_member_by_id_and_user_id(
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
await Channels.update_member_active_status(channel.id, user.id, form_data.is_active, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_MEMBER_ACTIVE_UPDATED,
|
||||
actor=user,
|
||||
subject_id=channel.id,
|
||||
data={'is_active': form_data.is_active},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
@@ -568,6 +622,13 @@ async def add_members_by_id(
|
||||
channel.id, user.id, form_data.user_ids, form_data.group_ids, db=db
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_MEMBER_ADDED,
|
||||
actor=user,
|
||||
subject_id=channel.id,
|
||||
data={'user_ids': form_data.user_ids, 'group_ids': form_data.group_ids},
|
||||
)
|
||||
return memberships
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -603,6 +664,13 @@ async def remove_members_by_id(
|
||||
try:
|
||||
deleted = await Channels.remove_members_from_channel(channel.id, form_data.user_ids, db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_MEMBER_REMOVED,
|
||||
actor=user,
|
||||
subject_id=channel.id,
|
||||
data={'user_ids': form_data.user_ids},
|
||||
)
|
||||
return deleted
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -632,7 +700,7 @@ async def update_channel_by_id(
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -641,6 +709,13 @@ async def update_channel_by_id(
|
||||
|
||||
try:
|
||||
channel = await Channels.update_channel_by_id(id, form_data, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': channel.name, 'type': channel.type},
|
||||
)
|
||||
return ChannelModel(**channel.model_dump())
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -670,6 +745,13 @@ async def delete_channel_by_id(
|
||||
|
||||
try:
|
||||
await Channels.delete_channel_by_id(id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': channel.name, 'type': channel.type},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -726,10 +808,14 @@ async def get_channel_messages(
|
||||
user_ids = list(set(m.user_id for m in message_list))
|
||||
fetched_users = {u.id: u for u in await Users.get_users_by_user_ids(user_ids, db=db)}
|
||||
|
||||
# Batch fetch reactions and reply counts in 2 queries (fixes N+1)
|
||||
message_ids = [m.id for m in message_list]
|
||||
all_reactions = await Messages.get_reactions_by_message_ids(message_ids, db=db)
|
||||
all_reply_counts = await Messages.get_thread_reply_counts_by_message_ids(message_ids, db=db)
|
||||
|
||||
messages = []
|
||||
for message in message_list:
|
||||
thread_replies = await Messages.get_thread_replies_by_message_id(message.id, db=db)
|
||||
latest_thread_reply_at = thread_replies[0].created_at if thread_replies else None
|
||||
reply_count, latest_reply_at = all_reply_counts.get(message.id, (0, None))
|
||||
|
||||
# Use message.user if present (for webhooks), otherwise look up by user_id
|
||||
user_info = message.user
|
||||
@@ -740,9 +826,9 @@ async def get_channel_messages(
|
||||
MessageUserResponse(
|
||||
**{
|
||||
**message.model_dump(),
|
||||
'reply_count': len(thread_replies),
|
||||
'latest_reply_at': latest_thread_reply_at,
|
||||
'reactions': await Messages.get_reactions_by_message_id(message.id, db=db),
|
||||
'reply_count': reply_count,
|
||||
'latest_reply_at': latest_reply_at,
|
||||
'reactions': all_reactions.get(message.id, []),
|
||||
'user': user_info,
|
||||
}
|
||||
)
|
||||
@@ -791,6 +877,10 @@ async def get_pinned_channel_messages(
|
||||
user_ids = list(set(m.user_id for m in message_list))
|
||||
fetched_users = {u.id: u for u in await Users.get_users_by_user_ids(user_ids, db=db)}
|
||||
|
||||
# Batch fetch reactions in 1 query (fixes N+1)
|
||||
message_ids = [m.id for m in message_list]
|
||||
all_reactions = await Messages.get_reactions_by_message_ids(message_ids, db=db)
|
||||
|
||||
messages = []
|
||||
for message in message_list:
|
||||
# Check for webhook identity in meta
|
||||
@@ -810,7 +900,7 @@ async def get_pinned_channel_messages(
|
||||
MessageWithReactionsResponse(
|
||||
**{
|
||||
**message.model_dump(),
|
||||
'reactions': await Messages.get_reactions_by_message_id(message.id, db=db),
|
||||
'reactions': all_reactions.get(message.id, []),
|
||||
'user': user_info,
|
||||
}
|
||||
)
|
||||
@@ -826,13 +916,16 @@ async def get_pinned_channel_messages(
|
||||
|
||||
async def send_notification(request, channel, message, active_user_ids, db=None):
|
||||
name = request.app.state.WEBUI_NAME
|
||||
webui_url = request.app.state.config.WEBUI_URL
|
||||
enable_user_webhooks = request.app.state.config.ENABLE_USER_WEBHOOKS
|
||||
webui_url = await Config.get('webui.url')
|
||||
enable_user_webhooks = await Config.get('ui.enable_user_webhooks')
|
||||
|
||||
users = await get_channel_users_with_access(channel, 'read', db=db)
|
||||
|
||||
# Batch fetch channel members in 1 query (fixes N+1)
|
||||
member_ids = {m.user_id for m in await Channels.get_members_by_channel_id(channel.id, db=db)}
|
||||
|
||||
for u in users:
|
||||
if (u.id not in active_user_ids) and await Channels.is_user_channel_member(channel.id, u.id, db=db):
|
||||
if (u.id not in active_user_ids) and u.id in member_ids:
|
||||
if enable_user_webhooks and u.settings:
|
||||
webhook_url = u.settings.ui.get('notifications', {}).get('webhook_url', None)
|
||||
if webhook_url:
|
||||
@@ -978,7 +1071,7 @@ async def model_response_handler(request, channel, message, user, db=None):
|
||||
)
|
||||
|
||||
tool_ids = _resolve_model_tool_ids(request.app, model_id)
|
||||
features = _resolve_model_features(request.app, model_id)
|
||||
features = await _resolve_model_features(request.app, model_id)
|
||||
filter_ids = _resolve_model_filter_ids(request.app, model_id)
|
||||
|
||||
# Build full form_data — same shape as frontend POST.
|
||||
@@ -1036,6 +1129,13 @@ async def new_message_handler(request: Request, id: str, form_data: MessageForm,
|
||||
):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
# Thread parent / reply target must belong to this channel (no cross-channel binding).
|
||||
for ref_id in (form_data.parent_id, form_data.reply_to_id):
|
||||
if ref_id:
|
||||
ref = await Messages.get_message_by_id(ref_id, include_thread_replies=False, db=db)
|
||||
if not ref or ref.channel_id != channel.id:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
try:
|
||||
message = await Messages.insert_new_message(form_data, channel.id, user.id, db=db)
|
||||
if message:
|
||||
@@ -1128,6 +1228,16 @@ async def post_new_message(
|
||||
|
||||
background_tasks.add_task(background_handler)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_CREATED,
|
||||
actor=user,
|
||||
subject_id=message.id,
|
||||
data={
|
||||
'channel_id': channel.id,
|
||||
'content_preview': message.content[:300],
|
||||
},
|
||||
)
|
||||
return message
|
||||
|
||||
except HTTPException as e:
|
||||
@@ -1255,12 +1365,37 @@ async def pin_channel_message(
|
||||
await Messages.update_is_pinned_by_id(message_id, form_data.is_pinned, user.id, db=db)
|
||||
message = await Messages.get_message_by_id(message_id, db=db)
|
||||
message_user = await Users.get_user_by_id(message.user_id, db=db)
|
||||
return MessageUserResponse(
|
||||
message_data = MessageUserResponse(
|
||||
**{
|
||||
**message.model_dump(),
|
||||
'user': UserNameResponse(**message_user.model_dump()) if message_user else None,
|
||||
}
|
||||
)
|
||||
|
||||
await sio.emit(
|
||||
'events:channel',
|
||||
{
|
||||
'channel_id': channel.id,
|
||||
'message_id': message.id,
|
||||
'data': {
|
||||
'type': 'message:update',
|
||||
'data': message_data.model_dump(),
|
||||
},
|
||||
'user': UserNameResponse(**user.model_dump()).model_dump(),
|
||||
'channel': channel.model_dump(),
|
||||
},
|
||||
to=f'channel:{channel.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_PINNED if form_data.is_pinned else EVENTS.MESSAGE_UNPINNED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
subject_type='message',
|
||||
data={'channel_id': id},
|
||||
)
|
||||
return message_data
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT())
|
||||
@@ -1302,6 +1437,10 @@ async def get_channel_thread_messages(
|
||||
user_ids = list(set(m.user_id for m in message_list))
|
||||
fetched_users = {u.id: u for u in await Users.get_users_by_user_ids(user_ids, db=db)}
|
||||
|
||||
# Batch fetch reactions in 1 query (fixes N+1)
|
||||
message_ids = [m.id for m in message_list]
|
||||
all_reactions = await Messages.get_reactions_by_message_ids(message_ids, db=db)
|
||||
|
||||
messages = []
|
||||
for message in message_list:
|
||||
# Use message.user if present (for webhooks), otherwise look up by user_id
|
||||
@@ -1315,7 +1454,7 @@ async def get_channel_thread_messages(
|
||||
**message.model_dump(),
|
||||
'reply_count': 0,
|
||||
'latest_reply_at': None,
|
||||
'reactions': await Messages.get_reactions_by_message_id(message.id, db=db),
|
||||
'reactions': all_reactions.get(message.id, []),
|
||||
'user': user_info,
|
||||
}
|
||||
)
|
||||
@@ -1384,6 +1523,13 @@ async def update_message_by_id(
|
||||
to=f'channel:{channel.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_UPDATED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'channel_id': id, 'content_preview': form_data.content[:300]},
|
||||
)
|
||||
return MessageModel(**message.model_dump())
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -1455,6 +1601,13 @@ async def add_reaction_to_message(
|
||||
to=f'channel:{channel.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_REACTION_ADDED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'channel_id': id, 'reaction': form_data.name},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -1523,6 +1676,13 @@ async def remove_reaction_by_id_and_user_id_and_name(
|
||||
to=f'channel:{channel.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_REACTION_REMOVED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'channel_id': id, 'reaction': form_data.name},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -1614,6 +1774,13 @@ async def delete_message_by_id(
|
||||
to=f'channel:{channel.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_DELETED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'channel_id': id},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -1699,6 +1866,13 @@ async def create_channel_webhook(
|
||||
if not webhook:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_WEBHOOK_CREATED,
|
||||
actor=user,
|
||||
subject_id=webhook.id,
|
||||
data={'channel_id': id, 'name': webhook.name},
|
||||
)
|
||||
return webhook
|
||||
|
||||
|
||||
@@ -1728,6 +1902,13 @@ async def update_channel_webhook(
|
||||
if not updated:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_WEBHOOK_UPDATED,
|
||||
actor=user,
|
||||
subject_id=webhook_id,
|
||||
data={'channel_id': id, 'name': updated.name},
|
||||
)
|
||||
return updated
|
||||
|
||||
|
||||
@@ -1752,7 +1933,16 @@ async def delete_channel_webhook(
|
||||
if not webhook or webhook.channel_id != id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
return await Channels.delete_webhook_by_id(webhook_id, db=db)
|
||||
deleted = await Channels.delete_webhook_by_id(webhook_id, db=db)
|
||||
if deleted:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHANNEL_WEBHOOK_DELETED,
|
||||
actor=user,
|
||||
subject_id=webhook_id,
|
||||
data={'channel_id': id},
|
||||
)
|
||||
return deleted
|
||||
|
||||
|
||||
############################
|
||||
@@ -1835,4 +2025,12 @@ async def post_webhook_message(
|
||||
to=f'channel:{channel.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_CREATED,
|
||||
actor={'id': webhook.id, 'name': webhook.name, 'role': 'webhook', 'type': 'webhook'},
|
||||
subject_id=message.id,
|
||||
source='channel_webhook',
|
||||
data={'channel_id': channel.id, 'content_preview': form_data.content[:300]},
|
||||
)
|
||||
return {'success': True, 'message_id': message.id}
|
||||
|
||||
@@ -10,8 +10,10 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import StreamingResponse
|
||||
from open_webui.config import ENABLE_ADMIN_CHAT_ACCESS, ENABLE_ADMIN_EXPORT
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.chats import (
|
||||
AggregateChatStats,
|
||||
ChatBody,
|
||||
@@ -30,11 +32,13 @@ from open_webui.models.folders import Folders
|
||||
from open_webui.models.shared_chats import SharedChatResponse, SharedChats
|
||||
from open_webui.models.tags import TagModel, Tags
|
||||
from open_webui.socket.main import get_event_emitter
|
||||
from open_webui.tasks import stop_item_tasks
|
||||
from open_webui.tasks import has_active_tasks, stop_item_tasks
|
||||
from open_webui.utils.access_control import filter_allowed_access_grants, has_permission
|
||||
from open_webui.utils.access_control.folders import has_folder_access
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.middleware import serialize_output
|
||||
from open_webui.utils.context_compaction import compact_chat_branch
|
||||
from open_webui.utils.misc import get_message_list
|
||||
from open_webui.utils.models import get_all_models
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
@@ -42,6 +46,81 @@ log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
SEARCH_FILTER_PREFIXES = ('tag:', 'folder:', 'pinned:', 'archived:', 'shared:')
|
||||
|
||||
CHAT_CONFIG_KEYS = {
|
||||
'ENABLE_CONTEXT_COMPACTION': 'chat.context_compaction.enable',
|
||||
'CONTEXT_COMPACTION_TOKEN_THRESHOLD': 'chat.context_compaction.token_threshold',
|
||||
'CONTEXT_COMPACTION_PROMPT_TEMPLATE': 'chat.context_compaction.prompt_template',
|
||||
}
|
||||
|
||||
|
||||
class ChatConfigForm(BaseModel):
|
||||
ENABLE_CONTEXT_COMPACTION: bool
|
||||
CONTEXT_COMPACTION_TOKEN_THRESHOLD: int
|
||||
CONTEXT_COMPACTION_PROMPT_TEMPLATE: str
|
||||
|
||||
|
||||
class CompactChatForm(BaseModel):
|
||||
model: str | None = None
|
||||
|
||||
|
||||
def chat_search_content_text(text: str) -> str:
|
||||
words = text.lower().strip().split(' ')
|
||||
return ' '.join(word for word in words if not word.startswith(SEARCH_FILTER_PREFIXES)).strip()
|
||||
|
||||
|
||||
def chat_search_snippet(chat: dict, search_text: str, max_length: int = 200) -> str | None:
|
||||
if not search_text:
|
||||
return None
|
||||
|
||||
messages = chat.get('messages', [])
|
||||
if isinstance(messages, dict):
|
||||
messages = messages.values()
|
||||
|
||||
for message in messages:
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
|
||||
content = message.get('content')
|
||||
if not isinstance(content, str):
|
||||
continue
|
||||
|
||||
index = content.lower().find(search_text)
|
||||
if index == -1:
|
||||
continue
|
||||
|
||||
start = max(index - max_length // 2, 0)
|
||||
end = min(start + max_length, len(content))
|
||||
if index + len(search_text) > end:
|
||||
end = min(index + len(search_text), len(content))
|
||||
start = max(end - max_length, 0)
|
||||
|
||||
snippet = ' '.join(content[start:end].split())
|
||||
return f'{"..." if start else ""}{snippet}{"..." if end < len(content) else ""}'
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def get_chat_config_values() -> dict:
|
||||
values = await Config.get_many(*CHAT_CONFIG_KEYS.values())
|
||||
return {field: values[storage_key] for field, storage_key in CHAT_CONFIG_KEYS.items() if storage_key in values}
|
||||
|
||||
|
||||
def chat_config_updates(data: dict) -> dict:
|
||||
return {CHAT_CONFIG_KEYS[field]: value for field, value in data.items() if field in CHAT_CONFIG_KEYS}
|
||||
|
||||
|
||||
async def require_chat_import_permission(request: Request, user, db: AsyncSession):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.import', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# GetChatList
|
||||
# Let the record outlive the session, so that what was
|
||||
@@ -400,7 +479,7 @@ async def export_chat_stats(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
# Check if the user has permission to share/export chats
|
||||
if (user.role != 'admin') and (not request.app.state.config.ENABLE_COMMUNITY_SHARING):
|
||||
if (user.role != 'admin') and (not await Config.get('ui.enable_community_sharing')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -449,7 +528,7 @@ async def export_single_chat_stats(
|
||||
Returns ChatStatsExport for the specified chat.
|
||||
"""
|
||||
# Check if the user has permission to share/export chats
|
||||
if (user.role != 'admin') and (not request.app.state.config.ENABLE_COMMUNITY_SHARING):
|
||||
if (user.role != 'admin') and (not await Config.get('ui.enable_community_sharing')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -495,15 +574,21 @@ async def delete_all_user_chats(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role == 'user' and not await has_permission(
|
||||
user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if user.role == 'user' and not await has_permission(user.id, 'chat.delete', await Config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
result = await Chats.delete_chats_by_user_id(user.id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_DELETED_ALL,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
subject_type='user',
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@@ -550,6 +635,7 @@ async def get_user_chat_list_by_user_id(
|
||||
|
||||
@router.post('/new', response_model=ChatResponse | None)
|
||||
async def create_new_chat(
|
||||
request: Request,
|
||||
form_data: ChatForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
@@ -561,13 +647,23 @@ async def create_new_chat(
|
||||
# to assume the column is clean. Also catches non-UUID / nonexistent IDs.
|
||||
if form_data.folder_id is not None:
|
||||
if not await Folders.get_folder_by_id_and_user_id(form_data.folder_id, user.id, db=db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
# Check shared folder write access
|
||||
shared_folder = await Folders.get_folder_by_id(form_data.folder_id, db=db)
|
||||
if not shared_folder or not await has_folder_access(user.id, shared_folder, 'write', db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
try:
|
||||
chat = await Chats.insert_new_chat(str(uuid4()), user.id, form_data, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_CREATED,
|
||||
actor=user,
|
||||
subject_id=chat.id,
|
||||
data={'title': chat.title, 'folder_id': chat.folder_id},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -581,18 +677,52 @@ async def create_new_chat(
|
||||
|
||||
@router.post('/import', response_model=list[ChatResponse])
|
||||
async def import_chats(
|
||||
request: Request,
|
||||
form_data: ChatsImportForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await require_chat_import_permission(request, user, db)
|
||||
|
||||
try:
|
||||
chats = await Chats.import_chats(user.id, form_data.chats, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_IMPORTED,
|
||||
actor=user,
|
||||
subject_type='chat.import',
|
||||
data={'count': len(chats), 'chat_ids': [chat.id for chat in chats]},
|
||||
)
|
||||
return chats
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
|
||||
############################
|
||||
# ChatConfig
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/config', response_model=ChatConfigForm)
|
||||
async def get_chat_config(user=Depends(get_admin_user)):
|
||||
return await get_chat_config_values()
|
||||
|
||||
|
||||
@router.post('/config', response_model=ChatConfigForm)
|
||||
async def set_chat_config(form_data: ChatConfigForm, user=Depends(get_admin_user)):
|
||||
threshold = max(1, int(form_data.CONTEXT_COMPACTION_TOKEN_THRESHOLD))
|
||||
await Config.upsert(
|
||||
chat_config_updates(
|
||||
{
|
||||
**form_data.model_dump(),
|
||||
'CONTEXT_COMPACTION_TOKEN_THRESHOLD': threshold,
|
||||
}
|
||||
)
|
||||
)
|
||||
return await get_chat_config_values()
|
||||
|
||||
|
||||
############################
|
||||
# GetChats
|
||||
############################
|
||||
@@ -611,10 +741,10 @@ async def search_user_chats(
|
||||
limit = 60
|
||||
skip = (page - 1) * limit
|
||||
|
||||
chat_list = [
|
||||
ChatTitleIdResponse(**chat.model_dump())
|
||||
for chat in await Chats.get_chats_by_user_id_and_search_text(user.id, text, skip=skip, limit=limit, db=db)
|
||||
]
|
||||
search_text = chat_search_content_text(text)
|
||||
chat_list = []
|
||||
for chat in await Chats.get_chats_by_user_id_and_search_text(user.id, text, skip=skip, limit=limit, db=db):
|
||||
chat_list.append(ChatTitleIdResponse(**chat.model_dump(), snippet=chat_search_snippet(chat.chat, search_text)))
|
||||
|
||||
# Delete tag if no chat is found
|
||||
words = text.strip().split(' ')
|
||||
@@ -800,14 +930,32 @@ async def get_archived_session_user_chat_list(
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# GetArchivedChatsCount
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/archived/count', response_model=int)
|
||||
async def get_archived_session_user_chat_count(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
return await Chats.count_archived_chats_by_user_id(user.id, db=db)
|
||||
|
||||
|
||||
############################
|
||||
# ArchiveAllChats
|
||||
############################
|
||||
|
||||
|
||||
@router.post('/archive/all', response_model=bool)
|
||||
async def archive_all_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
return await Chats.archive_all_chats_by_user_id(user.id, db=db)
|
||||
async def archive_all_chats(
|
||||
request: Request, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
result = await Chats.archive_all_chats_by_user_id(user.id, db=db)
|
||||
if result:
|
||||
await publish_event(request, EVENTS.CHAT_ARCHIVED, actor=user, subject_id=user.id, subject_type='user')
|
||||
return result
|
||||
|
||||
|
||||
############################
|
||||
@@ -816,15 +964,48 @@ async def archive_all_chats(user=Depends(get_verified_user), db: AsyncSession =
|
||||
|
||||
|
||||
@router.post('/unarchive/all', response_model=bool)
|
||||
async def unarchive_all_chats(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
return await Chats.unarchive_all_chats_by_user_id(user.id, db=db)
|
||||
async def unarchive_all_chats(
|
||||
request: Request, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
result = await Chats.unarchive_all_chats_by_user_id(user.id, db=db)
|
||||
if result:
|
||||
await publish_event(request, EVENTS.CHAT_UNARCHIVED, actor=user, subject_id=user.id, subject_type='user')
|
||||
return result
|
||||
|
||||
|
||||
############################
|
||||
# GetSharedChats
|
||||
# UnshareAllChats
|
||||
############################
|
||||
|
||||
|
||||
@router.delete('/share/all', response_model=bool)
|
||||
async def unshare_all_chats(
|
||||
request: Request, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
# Collect chat_ids that have shares so we can clear share_id and access grants
|
||||
shared_list = await SharedChats.get_by_user_id(user.id, db=db)
|
||||
chat_ids = [s.chat_id for s in shared_list]
|
||||
|
||||
# Delete all shared_chat rows for this user
|
||||
result = await SharedChats.delete_all_by_user_id(user.id, db=db)
|
||||
|
||||
# Clear share_id on the original chats and remove access grants
|
||||
for chat_id in chat_ids:
|
||||
await Chats.update_chat_share_id_by_id(chat_id, None, db=db)
|
||||
await AccessGrants.set_access_grants('shared_chat', chat_id, [], db=db)
|
||||
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_UNSHARED,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
subject_type='user',
|
||||
data={'count': len(chat_ids), 'chat_ids': chat_ids},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@router.get('/shared', response_model=list[SharedChatResponse])
|
||||
async def get_shared_session_user_chat_list(
|
||||
page: int | None = None,
|
||||
@@ -927,6 +1108,58 @@ async def get_user_chat_list_by_tag_name(
|
||||
return chats
|
||||
|
||||
|
||||
############################
|
||||
# CompactChat
|
||||
############################
|
||||
|
||||
|
||||
@router.post('/{id}/compact')
|
||||
async def compact_chat_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: CompactChatForm | None = None,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if not chat:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
if await has_active_tasks(request.app.state.redis, id):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail='Wait for the current response to finish before compacting.',
|
||||
)
|
||||
|
||||
if not request.app.state.MODELS:
|
||||
await get_all_models(request, user=user)
|
||||
|
||||
history = (chat.chat or {}).get('history') or {}
|
||||
messages_map = await Chats.get_messages_map_by_chat_id(id)
|
||||
message_list = get_message_list(messages_map or history.get('messages') or {}, history.get('currentId'))
|
||||
model_id = (form_data.model if form_data else None) or next(
|
||||
(message.get('model') for message in reversed(message_list) if message.get('model')),
|
||||
None,
|
||||
)
|
||||
|
||||
if not model_id:
|
||||
chat_models = (chat.chat or {}).get('models') or []
|
||||
model_id = chat_models[0] if chat_models else None
|
||||
if not model_id:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='No model found for context compaction.')
|
||||
|
||||
result = await compact_chat_branch(request, user, chat, model_id, request.app.state.MODELS)
|
||||
if result.get('compacted'):
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_COMPACTED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'dropped_messages': result.get('dropped_messages')},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
############################
|
||||
# GetChatById
|
||||
############################
|
||||
@@ -951,6 +1184,14 @@ async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSess
|
||||
if has_grant:
|
||||
chat = await Chats.get_chat_by_id(id, db=db)
|
||||
|
||||
# Check folder-based access (shared folders)
|
||||
if not chat:
|
||||
candidate = await Chats.get_chat_by_id(id, db=db)
|
||||
if candidate and candidate.folder_id:
|
||||
folder = await Folders.get_folder_by_id(candidate.folder_id, db=db)
|
||||
if folder and await has_folder_access(user.id, folder, 'read', db):
|
||||
chat = candidate
|
||||
|
||||
if chat:
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
@@ -964,6 +1205,7 @@ async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSess
|
||||
|
||||
@router.post('/{id}', response_model=ChatResponse | None)
|
||||
async def update_chat_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: ChatForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -972,26 +1214,27 @@ async def update_chat_by_id(
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
updated_chat = {**chat.chat, **form_data.chat}
|
||||
|
||||
# Re-derive content from output for assistant messages so that frontend
|
||||
# edits to output items are reflected in content. Only when output
|
||||
# actually changed — otherwise content set independently of output
|
||||
# (e.g. a `replace` event or an outlet filter footer) would be reverted.
|
||||
existing_messages = (chat.chat.get('history') or {}).get('messages') or {}
|
||||
for msg_id, msg in updated_chat.get('history', {}).get('messages', {}).items():
|
||||
if msg.get('role') == 'assistant' and msg.get('output'):
|
||||
if msg.get('output') != existing_messages.get(msg_id, {}).get('output'):
|
||||
msg['content'] = serialize_output(msg['output'])
|
||||
if 'history' in form_data.chat:
|
||||
updated_chat['history'] = Chats.merge_history(
|
||||
chat.chat.get('history'),
|
||||
form_data.chat.get('history'),
|
||||
)
|
||||
|
||||
chat = await Chats.update_chat_by_id(id, updated_chat, db=db)
|
||||
|
||||
# Reconcile chat_message rows with the committed blob.
|
||||
# This is the only caller where the frontend pushes a full
|
||||
# history with potential edits, deletions, or new branches.
|
||||
# Reconcile chat_message rows without inferring deletes from missing IDs.
|
||||
# Message deletion has its own endpoint below.
|
||||
messages = (updated_chat.get('history') or {}).get('messages') or {}
|
||||
if messages:
|
||||
await Chats.reconcile_messages_by_chat_id(id, user.id, messages)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'title': chat.title},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -1009,6 +1252,7 @@ class MessageForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/messages/{message_id}', response_model=ChatResponse | None)
|
||||
async def update_chat_message_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
message_id: str,
|
||||
form_data: MessageForm,
|
||||
@@ -1058,6 +1302,52 @@ async def update_chat_message_by_id(
|
||||
}
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_UPDATED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'chat_id': id, 'content_preview': form_data.content[:300]},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
|
||||
@router.delete('/{id}/messages/{message_id}', response_model=ChatResponse | None)
|
||||
async def delete_chat_message_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
message_id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
chat = await Chats.get_chat_by_id(id, db=db)
|
||||
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if chat.user_id != user.id and user.role != 'admin':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
chat = await Chats.delete_message_from_chat_by_id_and_message_id(id, message_id)
|
||||
if not chat:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_DELETED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'chat_id': id},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
|
||||
@@ -1071,6 +1361,7 @@ class EventForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/messages/{message_id}/event', response_model=bool | None)
|
||||
async def send_chat_message_event_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
message_id: str,
|
||||
form_data: EventForm,
|
||||
@@ -1104,6 +1395,13 @@ async def send_chat_message_event_by_id(
|
||||
await event_emitter(form_data.model_dump())
|
||||
else:
|
||||
return False
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MESSAGE_EVENT_RECEIVED,
|
||||
actor=user,
|
||||
subject_id=message_id,
|
||||
data={'chat_id': id, 'event_type': form_data.type},
|
||||
)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
@@ -1136,9 +1434,17 @@ async def delete_chat_by_id(
|
||||
|
||||
result = await Chats.delete_chat_by_id(id, db=db)
|
||||
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'owner_id': chat.user_id},
|
||||
)
|
||||
return result
|
||||
else:
|
||||
if not await has_permission(user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS):
|
||||
if not await has_permission(user.id, 'chat.delete', await Config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -1153,6 +1459,14 @@ async def delete_chat_by_id(
|
||||
await Chats.delete_orphan_tags_for_user(chat.meta.get('tags', []), user.id, threshold=1, db=db)
|
||||
|
||||
result = await Chats.delete_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'owner_id': user.id},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@@ -1178,10 +1492,19 @@ async def get_pinned_status_by_id(
|
||||
|
||||
|
||||
@router.post('/{id}/pin', response_model=ChatResponse | None)
|
||||
async def pin_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def pin_chat_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
chat = await Chats.toggle_chat_pinned_by_id(id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_PINNED if chat.pinned else EVENTS.CHAT_UNPINNED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
subject_type='chat',
|
||||
)
|
||||
return chat
|
||||
else:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
|
||||
@@ -1198,11 +1521,14 @@ class CloneForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/clone', response_model=ChatResponse | None)
|
||||
async def clone_chat_by_id(
|
||||
request: Request,
|
||||
form_data: CloneForm,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await require_chat_import_permission(request, user, db)
|
||||
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
updated_chat = {
|
||||
@@ -1229,6 +1555,13 @@ async def clone_chat_by_id(
|
||||
|
||||
if chats:
|
||||
chat = chats[0]
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_CLONED,
|
||||
actor=user,
|
||||
subject_id=chat.id,
|
||||
data={'original_chat_id': id},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -1246,8 +1579,13 @@ async def clone_chat_by_id(
|
||||
|
||||
@router.post('/{id}/clone/shared', response_model=ChatResponse | None)
|
||||
async def clone_shared_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await require_chat_import_permission(request, user, db)
|
||||
|
||||
chat = await Chats.get_chat_by_share_id(id, db=db)
|
||||
|
||||
# Fallback: admins can also access any chat directly by chat ID
|
||||
@@ -1334,6 +1672,13 @@ async def archive_chat_by_id(
|
||||
# Unarchived — ensure tag rows exist
|
||||
await Tags.ensure_tags_exist(tag_ids, user.id, db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_ARCHIVED if chat.archived else EVENTS.CHAT_UNARCHIVED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
subject_type='chat',
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
|
||||
@@ -1349,9 +1694,7 @@ async def share_chat_by_id(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'chat.share', request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'chat.share', await Config.get('user.permissions')):
|
||||
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
@@ -1363,6 +1706,13 @@ async def share_chat_by_id(
|
||||
shared = await SharedChats.update(chat.share_id, db=db)
|
||||
if shared:
|
||||
chat = await Chats.get_chat_by_id(id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_SHARED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'share_id': chat.share_id, 'updated': True},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
# Create a new share
|
||||
@@ -1374,6 +1724,13 @@ async def share_chat_by_id(
|
||||
if not chat:
|
||||
raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_SHARED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'share_id': shared.id},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
|
||||
|
||||
@@ -1382,19 +1739,26 @@ async def share_chat_by_id(
|
||||
|
||||
@router.delete('/{id}/share', response_model=bool | None)
|
||||
async def delete_shared_chat_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if not chat:
|
||||
raise HTTPException(status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
|
||||
if not chat.share_id:
|
||||
return False
|
||||
|
||||
await SharedChats.delete_by_chat_id(id, db=db)
|
||||
await Chats.update_chat_share_id_by_id(id, None, db=db)
|
||||
|
||||
if chat.share_id:
|
||||
await Chats.update_chat_share_id_by_id(id, None, db=db)
|
||||
|
||||
await AccessGrants.set_access_grants('shared_chat', id, [], db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_UNSHARED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'share_id': chat.share_id},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
@@ -1426,7 +1790,7 @@ async def update_shared_chat_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -1482,6 +1846,7 @@ class ChatFolderIdForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/folder', response_model=ChatResponse | None)
|
||||
async def update_chat_folder_id_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: ChatFolderIdForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -1493,12 +1858,22 @@ async def update_chat_folder_id_by_id(
|
||||
# folder_id values. None is allowed (moves the chat out of any folder).
|
||||
if form_data.folder_id is not None:
|
||||
if not await Folders.get_folder_by_id_and_user_id(form_data.folder_id, user.id, db=db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
# Check shared folder write access
|
||||
shared_folder = await Folders.get_folder_by_id(form_data.folder_id, db=db)
|
||||
if not shared_folder or not await has_folder_access(user.id, shared_folder, 'write', db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
chat = await Chats.update_chat_folder_id_by_id_and_user_id(id, user.id, form_data.folder_id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_FOLDER_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'folder_id': form_data.folder_id},
|
||||
)
|
||||
return ChatResponse(**chat.model_dump())
|
||||
else:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.DEFAULT())
|
||||
@@ -1526,6 +1901,7 @@ async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user), db: Asyn
|
||||
|
||||
@router.post('/{id}/tags', response_model=list[TagModel])
|
||||
async def add_tag_by_id_and_tag_name(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: TagForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -1544,6 +1920,13 @@ async def add_tag_by_id_and_tag_name(
|
||||
|
||||
if tag_id not in tags:
|
||||
await Chats.add_chat_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_TAG_ADDED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'tag': form_data.name},
|
||||
)
|
||||
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
tags = chat.meta.get('tags', [])
|
||||
@@ -1559,6 +1942,7 @@ async def add_tag_by_id_and_tag_name(
|
||||
|
||||
@router.delete('/{id}/tags', response_model=list[TagModel])
|
||||
async def delete_tag_by_id_and_tag_name(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: TagForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -1567,6 +1951,13 @@ async def delete_tag_by_id_and_tag_name(
|
||||
chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db)
|
||||
if chat:
|
||||
await Chats.delete_tag_by_id_and_user_id_and_tag_name(id, user.id, form_data.name, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CHAT_TAG_REMOVED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'tag': form_data.name},
|
||||
)
|
||||
|
||||
if await Chats.count_chats_by_tag_name_and_user_id(form_data.name, user.id, db=db) == 0:
|
||||
await Tags.delete_tag_by_name_and_user_id(form_data.name, user.id, db=db)
|
||||
|
||||
@@ -7,22 +7,27 @@ from typing import Optional
|
||||
import aiohttp
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from mcp.shared.auth import OAuthMetadata
|
||||
from open_webui.config import BannerModel, async_save_config, get_config, save_config
|
||||
from open_webui.config import BannerModel
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import get_custom_headers
|
||||
from open_webui.utils.mcp.client import MCPClient
|
||||
from open_webui.utils.oauth import (
|
||||
OAuthClientInformationFull,
|
||||
apply_connection_oauth_options,
|
||||
decrypt_data,
|
||||
encrypt_data,
|
||||
get_discovery_urls,
|
||||
get_oauth_client_info_with_dynamic_client_registration,
|
||||
get_oauth_client_info_with_static_credentials,
|
||||
recover_static_oauth_client_metadata,
|
||||
resolve_oauth_client_info,
|
||||
)
|
||||
from open_webui.utils.tools import (
|
||||
bearer_auth_header,
|
||||
get_tool_server_data,
|
||||
get_tool_server_url,
|
||||
set_terminal_servers,
|
||||
@@ -34,6 +39,44 @@ router = APIRouter()
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
CONNECTIONS_CONFIG_KEYS = {
|
||||
'ENABLE_DIRECT_CONNECTIONS': 'direct.enable',
|
||||
'ENABLE_BASE_MODELS_CACHE': 'models.base_models_cache',
|
||||
}
|
||||
CODE_EXECUTION_CONFIG_KEYS = {
|
||||
'ENABLE_CODE_EXECUTION': 'code_execution.enable',
|
||||
'CODE_EXECUTION_ENGINE': 'code_execution.engine',
|
||||
'CODE_EXECUTION_JUPYTER_URL': 'code_execution.jupyter.url',
|
||||
'CODE_EXECUTION_JUPYTER_AUTH': 'code_execution.jupyter.auth',
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_TOKEN': 'code_execution.jupyter.auth_token',
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD': 'code_execution.jupyter.auth_password',
|
||||
'CODE_EXECUTION_JUPYTER_TIMEOUT': 'code_execution.jupyter.timeout',
|
||||
'ENABLE_CODE_INTERPRETER': 'code_interpreter.enable',
|
||||
'CODE_INTERPRETER_ENGINE': 'code_interpreter.engine',
|
||||
'CODE_INTERPRETER_PROMPT_TEMPLATE': 'code_interpreter.prompt_template',
|
||||
'CODE_INTERPRETER_JUPYTER_URL': 'code_interpreter.jupyter.url',
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH': 'code_interpreter.jupyter.auth',
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN': 'code_interpreter.jupyter.auth_token',
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD': 'code_interpreter.jupyter.auth_password',
|
||||
'CODE_INTERPRETER_JUPYTER_TIMEOUT': 'code_interpreter.jupyter.timeout',
|
||||
}
|
||||
MODELS_CONFIG_KEYS = {
|
||||
'DEFAULT_MODELS': 'ui.default_models',
|
||||
'DEFAULT_PINNED_MODELS': 'ui.default_pinned_models',
|
||||
'MODEL_ORDER_LIST': 'ui.model_order_list',
|
||||
'DEFAULT_MODEL_METADATA': 'models.default_metadata',
|
||||
'DEFAULT_MODEL_PARAMS': 'models.default_params',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
############################
|
||||
# ImportConfig
|
||||
@@ -48,9 +91,15 @@ class ImportConfigForm(BaseModel):
|
||||
|
||||
@router.post('/import', response_model=dict)
|
||||
async def import_config(request: Request, form_data: ImportConfigForm, user=Depends(get_admin_user)):
|
||||
await async_save_config(form_data.config)
|
||||
request.app.state.config._sync_to_redis()
|
||||
return get_config()
|
||||
await Config.upsert(form_data.config)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_IMPORTED,
|
||||
actor=user,
|
||||
subject_id='import',
|
||||
data={'keys': list(form_data.config.keys())},
|
||||
)
|
||||
return await Config.get_all()
|
||||
|
||||
|
||||
############################
|
||||
@@ -60,7 +109,12 @@ async def import_config(request: Request, form_data: ImportConfigForm, user=Depe
|
||||
|
||||
@router.get('/export', response_model=dict)
|
||||
async def export_config(user=Depends(get_admin_user)):
|
||||
return get_config()
|
||||
return await Config.get_all()
|
||||
|
||||
|
||||
@router.get('/namespace/{namespace}', response_model=dict)
|
||||
async def get_config_namespace(namespace: str, user=Depends(get_admin_user)):
|
||||
return await Config.get_namespace(namespace)
|
||||
|
||||
|
||||
############################
|
||||
@@ -75,10 +129,7 @@ class ConnectionsConfigForm(BaseModel):
|
||||
|
||||
@router.get('/connections', response_model=ConnectionsConfigForm)
|
||||
async def get_connections_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_DIRECT_CONNECTIONS': request.app.state.config.ENABLE_DIRECT_CONNECTIONS,
|
||||
'ENABLE_BASE_MODELS_CACHE': request.app.state.config.ENABLE_BASE_MODELS_CACHE,
|
||||
}
|
||||
return await get_config_values(CONNECTIONS_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/connections', response_model=ConnectionsConfigForm)
|
||||
@@ -87,13 +138,17 @@ async def set_connections_config(
|
||||
form_data: ConnectionsConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
request.app.state.config.ENABLE_DIRECT_CONNECTIONS = form_data.ENABLE_DIRECT_CONNECTIONS
|
||||
request.app.state.config.ENABLE_BASE_MODELS_CACHE = form_data.ENABLE_BASE_MODELS_CACHE
|
||||
|
||||
return {
|
||||
'ENABLE_DIRECT_CONNECTIONS': request.app.state.config.ENABLE_DIRECT_CONNECTIONS,
|
||||
'ENABLE_BASE_MODELS_CACHE': request.app.state.config.ENABLE_BASE_MODELS_CACHE,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), CONNECTIONS_CONFIG_KEYS))
|
||||
values = await get_config_values(CONNECTIONS_CONFIG_KEYS)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_CONNECTIONS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='connections',
|
||||
subject_type='config',
|
||||
data=values,
|
||||
)
|
||||
return values
|
||||
|
||||
|
||||
class OAuthClientRegistrationForm(BaseModel):
|
||||
@@ -102,6 +157,7 @@ class OAuthClientRegistrationForm(BaseModel):
|
||||
client_name: str | None = None
|
||||
client_secret: str | None = None
|
||||
oauth_server_url: str | None = None
|
||||
oauth_scope: str | None = None
|
||||
|
||||
|
||||
@router.post('/oauth/clients/register')
|
||||
@@ -126,10 +182,11 @@ async def register_oauth_client(
|
||||
oauth_server_url,
|
||||
oauth_client_id=form_data.client_id,
|
||||
oauth_client_secret=form_data.client_secret,
|
||||
oauth_scope=form_data.oauth_scope,
|
||||
)
|
||||
else:
|
||||
oauth_client_info = await get_oauth_client_info_with_dynamic_client_registration(
|
||||
request, oauth_client_id, oauth_server_url
|
||||
request, oauth_client_id, oauth_server_url, oauth_scope=form_data.oauth_scope
|
||||
)
|
||||
return {
|
||||
'status': True,
|
||||
@@ -167,9 +224,7 @@ class ToolServersConfigForm(BaseModel):
|
||||
|
||||
@router.get('/tool_servers', response_model=ToolServersConfigForm)
|
||||
async def get_tool_servers_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'TOOL_SERVER_CONNECTIONS': request.app.state.config.TOOL_SERVER_CONNECTIONS,
|
||||
}
|
||||
return {'TOOL_SERVER_CONNECTIONS': await Config.get('tool_server.connections')}
|
||||
|
||||
|
||||
@router.post('/tool_servers', response_model=ToolServersConfigForm)
|
||||
@@ -178,13 +233,14 @@ async def set_tool_servers_config(
|
||||
form_data: ToolServersConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
for connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
existing_connections = await Config.get('tool_server.connections', []) or []
|
||||
for connection in existing_connections:
|
||||
server_type = connection.get('type', 'openapi')
|
||||
auth_type = connection.get('auth_type', 'none')
|
||||
|
||||
if auth_type in ('oauth_2.1', 'oauth_2.1_static'):
|
||||
# Remove existing OAuth clients for tool servers
|
||||
server_id = connection.get('info', {}).get('id')
|
||||
server_id = (connection.get('info') or {}).get('id')
|
||||
client_key = f'{server_type}:{server_id}'
|
||||
|
||||
try:
|
||||
@@ -193,21 +249,22 @@ async def set_tool_servers_config(
|
||||
pass
|
||||
|
||||
# Set new tool server connections
|
||||
request.app.state.config.TOOL_SERVER_CONNECTIONS = [
|
||||
connection.model_dump() for connection in form_data.TOOL_SERVER_CONNECTIONS
|
||||
]
|
||||
connections = [connection.model_dump() for connection in form_data.TOOL_SERVER_CONNECTIONS]
|
||||
await Config.upsert({'tool_server.connections': connections})
|
||||
|
||||
await set_tool_servers(request)
|
||||
|
||||
for connection in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
for connection in connections:
|
||||
server_type = connection.get('type', 'openapi')
|
||||
if server_type == 'mcp':
|
||||
server_id = connection.get('info', {}).get('id')
|
||||
server_id = (connection.get('info') or {}).get('id')
|
||||
auth_type = connection.get('auth_type', 'none')
|
||||
|
||||
if auth_type in ('oauth_2.1', 'oauth_2.1_static') and server_id:
|
||||
try:
|
||||
oauth_client_info = resolve_oauth_client_info(connection)
|
||||
oauth_client_info = await recover_static_oauth_client_metadata(connection, oauth_client_info)
|
||||
oauth_client_info = apply_connection_oauth_options(connection, oauth_client_info)
|
||||
request.app.state.oauth_client_manager.add_client(
|
||||
f'{server_type}:{server_id}',
|
||||
OAuthClientInformationFull(**oauth_client_info),
|
||||
@@ -216,9 +273,15 @@ async def set_tool_servers_config(
|
||||
log.debug(f'Failed to add OAuth client for MCP tool server: {e}')
|
||||
continue
|
||||
|
||||
return {
|
||||
'TOOL_SERVER_CONNECTIONS': request.app.state.config.TOOL_SERVER_CONNECTIONS,
|
||||
}
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_TOOL_SERVERS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='tool_server.connections',
|
||||
subject_type='config',
|
||||
data={'count': len(connections), 'types': [connection.get('type', 'openapi') for connection in connections]},
|
||||
)
|
||||
return {'TOOL_SERVER_CONNECTIONS': connections}
|
||||
|
||||
|
||||
class TerminalServerConnection(BaseModel):
|
||||
@@ -249,9 +312,7 @@ class TerminalServersConfigForm(BaseModel):
|
||||
|
||||
@router.get('/terminal_servers')
|
||||
async def get_terminal_servers_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'TERMINAL_SERVER_CONNECTIONS': request.app.state.config.TERMINAL_SERVER_CONNECTIONS,
|
||||
}
|
||||
return {'TERMINAL_SERVER_CONNECTIONS': await Config.get('terminal_server.connections')}
|
||||
|
||||
|
||||
@router.post('/terminal_servers')
|
||||
@@ -260,15 +321,20 @@ async def set_terminal_servers_config(
|
||||
form_data: TerminalServersConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
request.app.state.config.TERMINAL_SERVER_CONNECTIONS = [
|
||||
connection.model_dump() for connection in form_data.TERMINAL_SERVER_CONNECTIONS
|
||||
]
|
||||
connections = [connection.model_dump() for connection in form_data.TERMINAL_SERVER_CONNECTIONS]
|
||||
await Config.upsert({'terminal_server.connections': connections})
|
||||
|
||||
await set_terminal_servers(request)
|
||||
|
||||
return {
|
||||
'TERMINAL_SERVER_CONNECTIONS': request.app.state.config.TERMINAL_SERVER_CONNECTIONS,
|
||||
}
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_TERMINAL_SERVERS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='terminal_server.connections',
|
||||
subject_type='config',
|
||||
data={'count': len(connections)},
|
||||
)
|
||||
return {'TERMINAL_SERVER_CONNECTIONS': connections}
|
||||
|
||||
|
||||
@router.post('/terminal_servers/verify')
|
||||
@@ -287,7 +353,7 @@ async def verify_terminal_server_connection(
|
||||
|
||||
headers = {}
|
||||
if form_data.auth_type == 'bearer' and form_data.key:
|
||||
headers['Authorization'] = f'Bearer {form_data.key}'
|
||||
headers.update(bearer_auth_header(form_data.key))
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
@@ -328,6 +394,24 @@ class TerminalServerPolicyForm(BaseModel):
|
||||
policy_data: dict
|
||||
|
||||
|
||||
class TerminalServerLifecycleForm(BaseModel):
|
||||
url: str
|
||||
key: str | None = ''
|
||||
auth_type: str | None = 'bearer'
|
||||
policy_id: str
|
||||
lifecycle_data: dict
|
||||
|
||||
|
||||
class TerminalServerRefreshForm(BaseModel):
|
||||
url: str
|
||||
key: str | None = ''
|
||||
auth_type: str | None = 'bearer'
|
||||
user_id: str | None = None
|
||||
policy_id: str | None = None
|
||||
only_idle: bool = True
|
||||
reset: bool = False
|
||||
|
||||
|
||||
@router.post('/terminal_servers/policy')
|
||||
async def put_terminal_server_policy(
|
||||
request: Request, form_data: TerminalServerPolicyForm, user=Depends(get_admin_user)
|
||||
@@ -341,7 +425,7 @@ async def put_terminal_server_policy(
|
||||
|
||||
headers = {'Content-Type': 'application/json'}
|
||||
if form_data.auth_type == 'bearer' and form_data.key:
|
||||
headers['Authorization'] = f'Bearer {form_data.key}'
|
||||
headers.update(bearer_auth_header(form_data.key))
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
@@ -363,6 +447,91 @@ async def put_terminal_server_policy(
|
||||
raise HTTPException(status_code=400, detail='Failed to save policy to terminal server')
|
||||
|
||||
|
||||
@router.post('/terminal_servers/lifecycle')
|
||||
async def put_terminal_server_lifecycle(
|
||||
request: Request, form_data: TerminalServerLifecycleForm, user=Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Proxy a policy lifecycle PUT to an orchestrator terminal server.
|
||||
"""
|
||||
base_url = (form_data.url or '').rstrip('/')
|
||||
if not base_url:
|
||||
raise HTTPException(status_code=400, detail='Terminal server URL is required')
|
||||
|
||||
headers = {'Content-Type': 'application/json'}
|
||||
if form_data.auth_type == 'bearer' and form_data.key:
|
||||
headers.update(bearer_auth_header(form_data.key))
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
trust_env=True,
|
||||
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
|
||||
) as session:
|
||||
lifecycle_url = f'{base_url}/api/v1/policies/{form_data.policy_id}/lifecycle'
|
||||
async with session.put(
|
||||
lifecycle_url,
|
||||
headers=headers,
|
||||
json=form_data.lifecycle_data,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as resp:
|
||||
if resp.ok:
|
||||
return await resp.json()
|
||||
detail = await resp.text()
|
||||
raise HTTPException(status_code=resp.status, detail=detail)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.debug(f'Failed to save lifecycle to terminal server: {e}')
|
||||
raise HTTPException(status_code=400, detail='Failed to save lifecycle to terminal server')
|
||||
|
||||
|
||||
@router.post('/terminal_servers/refresh')
|
||||
async def refresh_terminal_server_terminals(
|
||||
request: Request, form_data: TerminalServerRefreshForm, user=Depends(get_admin_user)
|
||||
):
|
||||
"""
|
||||
Proxy a terminal refresh request to an orchestrator terminal server.
|
||||
"""
|
||||
base_url = (form_data.url or '').rstrip('/')
|
||||
if not base_url:
|
||||
raise HTTPException(status_code=400, detail='Terminal server URL is required')
|
||||
|
||||
headers = {'Content-Type': 'application/json'}
|
||||
if form_data.auth_type == 'bearer' and form_data.key:
|
||||
headers.update(bearer_auth_header(form_data.key))
|
||||
|
||||
body = {
|
||||
'only_idle': form_data.only_idle,
|
||||
'reset': form_data.reset,
|
||||
}
|
||||
if form_data.user_id:
|
||||
body['user_id'] = form_data.user_id
|
||||
if form_data.policy_id:
|
||||
body['policy_id'] = form_data.policy_id
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession(
|
||||
trust_env=True,
|
||||
timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT),
|
||||
) as session:
|
||||
refresh_url = f'{base_url}/api/v1/terminals/refresh'
|
||||
async with session.post(
|
||||
refresh_url,
|
||||
headers=headers,
|
||||
json=body,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as resp:
|
||||
if resp.ok:
|
||||
return await resp.json()
|
||||
detail = await resp.text()
|
||||
raise HTTPException(status_code=resp.status, detail=detail)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.debug(f'Failed to refresh terminals: {e}')
|
||||
raise HTTPException(status_code=400, detail='Failed to refresh terminals')
|
||||
|
||||
|
||||
@router.post('/tool_servers/verify')
|
||||
async def verify_tool_servers_config(request: Request, form_data: ToolServerConnection, user=Depends(get_admin_user)):
|
||||
"""
|
||||
@@ -518,67 +687,29 @@ class CodeInterpreterConfigForm(BaseModel):
|
||||
|
||||
@router.get('/code_execution', response_model=CodeInterpreterConfigForm)
|
||||
async def get_code_execution_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_CODE_EXECUTION': request.app.state.config.ENABLE_CODE_EXECUTION,
|
||||
'CODE_EXECUTION_ENGINE': request.app.state.config.CODE_EXECUTION_ENGINE,
|
||||
'CODE_EXECUTION_JUPYTER_URL': request.app.state.config.CODE_EXECUTION_JUPYTER_URL,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_EXECUTION_JUPYTER_TIMEOUT': request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT,
|
||||
'ENABLE_CODE_INTERPRETER': request.app.state.config.ENABLE_CODE_INTERPRETER,
|
||||
'CODE_INTERPRETER_ENGINE': request.app.state.config.CODE_INTERPRETER_ENGINE,
|
||||
'CODE_INTERPRETER_PROMPT_TEMPLATE': request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE,
|
||||
'CODE_INTERPRETER_JUPYTER_URL': request.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_INTERPRETER_JUPYTER_TIMEOUT': request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
}
|
||||
return await get_config_values(CODE_EXECUTION_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/code_execution', response_model=CodeInterpreterConfigForm)
|
||||
async def set_code_execution_config(
|
||||
request: Request, form_data: CodeInterpreterConfigForm, user=Depends(get_admin_user)
|
||||
):
|
||||
request.app.state.config.ENABLE_CODE_EXECUTION = form_data.ENABLE_CODE_EXECUTION
|
||||
|
||||
request.app.state.config.CODE_EXECUTION_ENGINE = form_data.CODE_EXECUTION_ENGINE
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_URL = form_data.CODE_EXECUTION_JUPYTER_URL
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH = form_data.CODE_EXECUTION_JUPYTER_AUTH
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN = form_data.CODE_EXECUTION_JUPYTER_AUTH_TOKEN
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD = form_data.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT = form_data.CODE_EXECUTION_JUPYTER_TIMEOUT
|
||||
|
||||
request.app.state.config.ENABLE_CODE_INTERPRETER = form_data.ENABLE_CODE_INTERPRETER
|
||||
request.app.state.config.CODE_INTERPRETER_ENGINE = form_data.CODE_INTERPRETER_ENGINE
|
||||
request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE = form_data.CODE_INTERPRETER_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_URL = form_data.CODE_INTERPRETER_JUPYTER_URL
|
||||
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH = form_data.CODE_INTERPRETER_JUPYTER_AUTH
|
||||
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN = form_data.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD = form_data.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD
|
||||
request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT = form_data.CODE_INTERPRETER_JUPYTER_TIMEOUT
|
||||
|
||||
return {
|
||||
'ENABLE_CODE_EXECUTION': request.app.state.config.ENABLE_CODE_EXECUTION,
|
||||
'CODE_EXECUTION_ENGINE': request.app.state.config.CODE_EXECUTION_ENGINE,
|
||||
'CODE_EXECUTION_JUPYTER_URL': request.app.state.config.CODE_EXECUTION_JUPYTER_URL,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_EXECUTION_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_EXECUTION_JUPYTER_TIMEOUT': request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT,
|
||||
'ENABLE_CODE_INTERPRETER': request.app.state.config.ENABLE_CODE_INTERPRETER,
|
||||
'CODE_INTERPRETER_ENGINE': request.app.state.config.CODE_INTERPRETER_ENGINE,
|
||||
'CODE_INTERPRETER_PROMPT_TEMPLATE': request.app.state.config.CODE_INTERPRETER_PROMPT_TEMPLATE,
|
||||
'CODE_INTERPRETER_JUPYTER_URL': request.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_TOKEN': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN,
|
||||
'CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD': request.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD,
|
||||
'CODE_INTERPRETER_JUPYTER_TIMEOUT': request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), CODE_EXECUTION_CONFIG_KEYS))
|
||||
values = await get_config_values(CODE_EXECUTION_CONFIG_KEYS)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_CODE_EXECUTION_UPDATED,
|
||||
actor=user,
|
||||
subject_id='code_execution',
|
||||
subject_type='config',
|
||||
data={
|
||||
'code_execution_enabled': values.get('ENABLE_CODE_EXECUTION'),
|
||||
'code_execution_engine': values.get('CODE_EXECUTION_ENGINE'),
|
||||
'code_interpreter_enabled': values.get('ENABLE_CODE_INTERPRETER'),
|
||||
'code_interpreter_engine': values.get('CODE_INTERPRETER_ENGINE'),
|
||||
},
|
||||
)
|
||||
return values
|
||||
|
||||
|
||||
############################
|
||||
@@ -595,35 +726,32 @@ class ModelsConfigForm(BaseModel):
|
||||
@router.get('/models/defaults')
|
||||
async def get_models_defaults(request: Request, user=Depends(get_verified_user)):
|
||||
return {
|
||||
'DEFAULT_MODEL_METADATA': request.app.state.config.DEFAULT_MODEL_METADATA,
|
||||
'DEFAULT_MODEL_METADATA': await Config.get('models.default_metadata'),
|
||||
}
|
||||
|
||||
|
||||
@router.get('/models', response_model=ModelsConfigForm)
|
||||
async def get_models_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'DEFAULT_MODELS': request.app.state.config.DEFAULT_MODELS,
|
||||
'DEFAULT_PINNED_MODELS': request.app.state.config.DEFAULT_PINNED_MODELS,
|
||||
'MODEL_ORDER_LIST': request.app.state.config.MODEL_ORDER_LIST,
|
||||
'DEFAULT_MODEL_METADATA': request.app.state.config.DEFAULT_MODEL_METADATA,
|
||||
'DEFAULT_MODEL_PARAMS': request.app.state.config.DEFAULT_MODEL_PARAMS,
|
||||
}
|
||||
return await get_config_values(MODELS_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/models', response_model=ModelsConfigForm)
|
||||
async def set_models_config(request: Request, form_data: ModelsConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.DEFAULT_MODELS = form_data.DEFAULT_MODELS
|
||||
request.app.state.config.DEFAULT_PINNED_MODELS = form_data.DEFAULT_PINNED_MODELS
|
||||
request.app.state.config.MODEL_ORDER_LIST = form_data.MODEL_ORDER_LIST
|
||||
request.app.state.config.DEFAULT_MODEL_METADATA = form_data.DEFAULT_MODEL_METADATA
|
||||
request.app.state.config.DEFAULT_MODEL_PARAMS = form_data.DEFAULT_MODEL_PARAMS
|
||||
return {
|
||||
'DEFAULT_MODELS': request.app.state.config.DEFAULT_MODELS,
|
||||
'DEFAULT_PINNED_MODELS': request.app.state.config.DEFAULT_PINNED_MODELS,
|
||||
'MODEL_ORDER_LIST': request.app.state.config.MODEL_ORDER_LIST,
|
||||
'DEFAULT_MODEL_METADATA': request.app.state.config.DEFAULT_MODEL_METADATA,
|
||||
'DEFAULT_MODEL_PARAMS': request.app.state.config.DEFAULT_MODEL_PARAMS,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), MODELS_CONFIG_KEYS))
|
||||
values = await get_config_values(MODELS_CONFIG_KEYS)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_MODELS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='models',
|
||||
subject_type='config',
|
||||
data={
|
||||
'default_models': values.get('DEFAULT_MODELS'),
|
||||
'default_pinned_models': values.get('DEFAULT_PINNED_MODELS'),
|
||||
'model_order_count': len(values.get('MODEL_ORDER_LIST') or []),
|
||||
},
|
||||
)
|
||||
return values
|
||||
|
||||
|
||||
class PromptSuggestion(BaseModel):
|
||||
@@ -642,8 +770,17 @@ async def set_default_suggestions(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
data = form_data.model_dump()
|
||||
request.app.state.config.DEFAULT_PROMPT_SUGGESTIONS = data['suggestions']
|
||||
return request.app.state.config.DEFAULT_PROMPT_SUGGESTIONS
|
||||
await Config.upsert({'ui.prompt_suggestions': data['suggestions']})
|
||||
suggestions = await Config.get('ui.prompt_suggestions')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_SUGGESTIONS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='ui.prompt_suggestions',
|
||||
subject_type='config',
|
||||
data={'count': len(suggestions or [])},
|
||||
)
|
||||
return suggestions
|
||||
|
||||
|
||||
############################
|
||||
@@ -662,8 +799,17 @@ async def set_banners(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
data = form_data.model_dump()
|
||||
request.app.state.config.BANNERS = data['banners']
|
||||
return request.app.state.config.BANNERS
|
||||
await Config.upsert({'ui.banners': data['banners']})
|
||||
banners = await Config.get('ui.banners')
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_BANNERS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='ui.banners',
|
||||
subject_type='config',
|
||||
data={'count': len(banners or [])},
|
||||
)
|
||||
return banners
|
||||
|
||||
|
||||
@router.get('/banners', response_model=list[BannerModel])
|
||||
@@ -671,4 +817,4 @@ async def get_banners(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
return request.app.state.config.BANNERS
|
||||
return await Config.get('ui.banners')
|
||||
|
||||
@@ -4,7 +4,9 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.concurrency import run_in_threadpool
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.feedbacks import (
|
||||
FeedbackForm,
|
||||
FeedbackIdResponse,
|
||||
@@ -25,6 +27,16 @@ log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
EVALUATION_CONFIG_KEYS = {
|
||||
'ENABLE_EVALUATION_ARENA_MODELS': 'evaluation.arena.enable',
|
||||
'EVALUATION_ARENA_MODELS': 'evaluation.arena.models',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
# Leaderboard Elo Rating Computation
|
||||
# The judgment has already been rendered with grace;
|
||||
@@ -255,10 +267,7 @@ async def get_model_history(
|
||||
|
||||
@router.get('/config')
|
||||
async def get_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_EVALUATION_ARENA_MODELS': request.app.state.config.ENABLE_EVALUATION_ARENA_MODELS,
|
||||
'EVALUATION_ARENA_MODELS': request.app.state.config.EVALUATION_ARENA_MODELS,
|
||||
}
|
||||
return await get_config_values(EVALUATION_CONFIG_KEYS)
|
||||
|
||||
|
||||
############################
|
||||
@@ -277,15 +286,25 @@ async def update_config(
|
||||
form_data: UpdateConfigForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
config = request.app.state.config
|
||||
updates = {}
|
||||
if form_data.ENABLE_EVALUATION_ARENA_MODELS is not None:
|
||||
config.ENABLE_EVALUATION_ARENA_MODELS = form_data.ENABLE_EVALUATION_ARENA_MODELS
|
||||
updates['evaluation.arena.enable'] = form_data.ENABLE_EVALUATION_ARENA_MODELS
|
||||
if form_data.EVALUATION_ARENA_MODELS is not None:
|
||||
config.EVALUATION_ARENA_MODELS = form_data.EVALUATION_ARENA_MODELS
|
||||
return {
|
||||
'ENABLE_EVALUATION_ARENA_MODELS': config.ENABLE_EVALUATION_ARENA_MODELS,
|
||||
'EVALUATION_ARENA_MODELS': config.EVALUATION_ARENA_MODELS,
|
||||
}
|
||||
updates['evaluation.arena.models'] = form_data.EVALUATION_ARENA_MODELS
|
||||
await Config.upsert(updates)
|
||||
values = await get_config_values(EVALUATION_CONFIG_KEYS)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_UPDATED,
|
||||
actor=user,
|
||||
subject_id='evaluation',
|
||||
data={
|
||||
'keys': list(updates.keys()),
|
||||
'arena_enabled': values.get('ENABLE_EVALUATION_ARENA_MODELS'),
|
||||
'arena_model_count': len(values.get('EVALUATION_ARENA_MODELS') or []),
|
||||
},
|
||||
)
|
||||
return values
|
||||
|
||||
|
||||
@router.get('/feedbacks/models', response_model=list[str])
|
||||
@@ -299,8 +318,19 @@ async def get_all_feedback_ids(user=Depends(get_admin_user), db: AsyncSession =
|
||||
|
||||
|
||||
@router.delete('/feedbacks/all')
|
||||
async def delete_all_feedbacks(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_all_feedbacks(
|
||||
request: Request,
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
success = await Feedbacks.delete_all_feedbacks(db=db)
|
||||
if success:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FEEDBACK_DELETED_ALL,
|
||||
actor=user,
|
||||
subject_id='all',
|
||||
)
|
||||
return success
|
||||
|
||||
|
||||
@@ -332,8 +362,20 @@ async def get_user_feedbacks(
|
||||
|
||||
|
||||
@router.delete('/feedbacks', response_model=bool)
|
||||
async def delete_feedbacks(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_feedbacks(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
success = await Feedbacks.delete_feedbacks_by_user_id(user.id, db=db)
|
||||
if success:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FEEDBACK_DELETED_ALL,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
subject_type='user',
|
||||
)
|
||||
return success
|
||||
|
||||
|
||||
@@ -377,6 +419,13 @@ async def create_feedback(
|
||||
detail=ERROR_MESSAGES.DEFAULT(),
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FEEDBACK_CREATED,
|
||||
actor=user,
|
||||
subject_id=feedback.id,
|
||||
data={'rating': getattr(feedback, 'rating', None)},
|
||||
)
|
||||
return feedback
|
||||
|
||||
|
||||
@@ -395,6 +444,7 @@ async def get_feedback_by_id(id: str, user=Depends(get_verified_user), db: Async
|
||||
|
||||
@router.post('/feedback/{id}', response_model=FeedbackModel)
|
||||
async def update_feedback_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: FeedbackForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -408,12 +458,22 @@ async def update_feedback_by_id(
|
||||
if not feedback:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FEEDBACK_UPDATED,
|
||||
actor=user,
|
||||
subject_id=feedback.id,
|
||||
data={'rating': getattr(feedback, 'rating', None)},
|
||||
)
|
||||
return feedback
|
||||
|
||||
|
||||
@router.delete('/feedback/{id}')
|
||||
async def delete_feedback_by_id(
|
||||
id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role == 'admin':
|
||||
success = await Feedbacks.delete_feedback_by_id(id=id, db=db)
|
||||
@@ -423,4 +483,10 @@ async def delete_feedback_by_id(
|
||||
if not success:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FEEDBACK_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
)
|
||||
return success
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import errno
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
@@ -23,9 +24,11 @@ from fastapi import (
|
||||
from fastapi.responses import FileResponse, StreamingResponse
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STORAGE_LOCAL_CACHE, STORAGE_PROVIDER, UPLOAD_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_db_context, get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.channels import Channels
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.files import (
|
||||
FileForm,
|
||||
@@ -123,7 +126,7 @@ async def process_uploaded_file(
|
||||
if _is_text_file(file_path):
|
||||
content_type = 'text/plain'
|
||||
|
||||
stt_supported = getattr(request.app.state.config, 'STT_SUPPORTED_CONTENT_TYPES', [])
|
||||
stt_supported = await Config.get('audio.stt.supported_content_types', [])
|
||||
|
||||
if content_type and strict_match_mime_type(stt_supported, content_type):
|
||||
# Audio / STT-supported files → transcribe then index
|
||||
@@ -144,7 +147,7 @@ async def process_uploaded_file(
|
||||
elif (
|
||||
content_type
|
||||
and content_type.startswith(('image/', 'video/'))
|
||||
and request.app.state.config.CONTENT_EXTRACTION_ENGINE != 'external'
|
||||
and await Config.get('rag.content_extraction_engine') != 'external'
|
||||
):
|
||||
# Media files without an external extraction engine
|
||||
if content_type.startswith('video/'):
|
||||
@@ -178,19 +181,39 @@ async def process_uploaded_file(
|
||||
knowledge_id = file_metadata.get('knowledge_id')
|
||||
if knowledge_id:
|
||||
try:
|
||||
await Knowledges.add_file_to_knowledge_by_id(
|
||||
knowledge_id=knowledge_id,
|
||||
file_id=file_item.id,
|
||||
user_id=user.id,
|
||||
directory_id=file_metadata.get('directory_id'),
|
||||
# Gate like POST /knowledge/{id}/file/add: a client-supplied
|
||||
# metadata.knowledge_id must not let a non-writer attach files (CWE-862/863).
|
||||
knowledge = await Knowledges.get_knowledge_by_id(id=knowledge_id, db=db_session)
|
||||
can_write = bool(knowledge) and (
|
||||
knowledge.user_id == user.id
|
||||
or user.role == 'admin'
|
||||
or await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge.id,
|
||||
permission='write',
|
||||
db=db_session,
|
||||
)
|
||||
)
|
||||
await process_file(
|
||||
request,
|
||||
ProcessFileForm(file_id=file_item.id, collection_name=knowledge_id),
|
||||
user=user,
|
||||
db=db_session,
|
||||
)
|
||||
log.info(f'Linked file {file_item.id} to knowledge {knowledge_id}')
|
||||
if not can_write:
|
||||
log.warning(
|
||||
f'Refusing to auto-link file {file_item.id} to knowledge '
|
||||
f'{knowledge_id}: user {user.id} lacks write access'
|
||||
)
|
||||
else:
|
||||
await Knowledges.add_file_to_knowledge_by_id(
|
||||
knowledge_id=knowledge_id,
|
||||
file_id=file_item.id,
|
||||
user_id=user.id,
|
||||
directory_id=file_metadata.get('directory_id'),
|
||||
)
|
||||
await process_file(
|
||||
request,
|
||||
ProcessFileForm(file_id=file_item.id, collection_name=knowledge_id),
|
||||
user=user,
|
||||
db=db_session,
|
||||
)
|
||||
log.info(f'Linked file {file_item.id} to knowledge {knowledge_id}')
|
||||
except Exception as e:
|
||||
log.warning(f'Failed to link file {file_item.id} to knowledge {knowledge_id}: {e}')
|
||||
|
||||
@@ -226,7 +249,7 @@ async def upload_file(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
return await upload_file_handler(
|
||||
result = await upload_file_handler(
|
||||
request,
|
||||
file=file,
|
||||
metadata=metadata,
|
||||
@@ -237,6 +260,27 @@ async def upload_file(
|
||||
db=db,
|
||||
)
|
||||
|
||||
if isinstance(result, dict):
|
||||
result_id = result.get('id')
|
||||
result_filename = result.get('filename')
|
||||
result_meta = result.get('meta') or {}
|
||||
else:
|
||||
result_id = result.id
|
||||
result_filename = result.filename
|
||||
result_meta = result.meta or {}
|
||||
|
||||
result_content_type = (
|
||||
result_meta.get('content_type') if isinstance(result_meta, dict) else getattr(result_meta, 'content_type', None)
|
||||
)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FILE_UPLOADED,
|
||||
actor=user,
|
||||
subject_id=result_id,
|
||||
data={'filename': result_filename, 'content_type': result_content_type},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
async def upload_file_handler(
|
||||
request: Request,
|
||||
@@ -268,36 +312,59 @@ async def upload_file_handler(
|
||||
# Remove the leading dot from the file extension and lowercase it
|
||||
file_extension = file_extension[1:].lower() if file_extension else ''
|
||||
|
||||
if process and request.app.state.config.ALLOWED_FILE_EXTENSIONS:
|
||||
request.app.state.config.ALLOWED_FILE_EXTENSIONS = [
|
||||
ext for ext in request.app.state.config.ALLOWED_FILE_EXTENSIONS if ext
|
||||
]
|
||||
allowed_file_extensions = await Config.get('rag.file.allowed_extensions')
|
||||
if process and allowed_file_extensions:
|
||||
allowed_file_extensions = [ext for ext in allowed_file_extensions if ext]
|
||||
|
||||
if file_extension not in request.app.state.config.ALLOWED_FILE_EXTENSIONS:
|
||||
if file_extension not in allowed_file_extensions:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(f'File type {file_extension} is not allowed'),
|
||||
)
|
||||
|
||||
# replace filename with uuid
|
||||
# Prefer readable storage names for admins, but fall back if the filesystem rejects it.
|
||||
id = str(uuid.uuid4())
|
||||
name = filename
|
||||
filename = f'{id}_{filename}'
|
||||
contents, file_path = await asyncio.to_thread(
|
||||
Storage.upload_file,
|
||||
file.file,
|
||||
filename,
|
||||
{
|
||||
'OpenWebUI-User-Email': user.email,
|
||||
'OpenWebUI-User-Id': user.id,
|
||||
'OpenWebUI-User-Name': user.name,
|
||||
'OpenWebUI-File-Id': id,
|
||||
},
|
||||
)
|
||||
tags = {
|
||||
'OpenWebUI-User-Email': user.email,
|
||||
'OpenWebUI-User-Id': user.id,
|
||||
'OpenWebUI-User-Name': user.name,
|
||||
'OpenWebUI-File-Id': id,
|
||||
}
|
||||
try:
|
||||
contents, file_path = await asyncio.to_thread(Storage.upload_file, file.file, filename, tags)
|
||||
except OSError as e:
|
||||
if e.errno != errno.ENAMETOOLONG:
|
||||
log.exception(e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e.strerror or 'Error uploading file'),
|
||||
)
|
||||
|
||||
file.file.seek(0)
|
||||
filename = f'{id}.{file_extension}' if file_extension else id
|
||||
try:
|
||||
contents, file_path = await asyncio.to_thread(Storage.upload_file, file.file, filename, tags)
|
||||
except OSError as e:
|
||||
log.exception(e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e.strerror or 'Error uploading file'),
|
||||
)
|
||||
max_size = await Config.get('rag.file.max_size')
|
||||
if max_size and len(contents) > int(max_size) * 1024 * 1024:
|
||||
await asyncio.to_thread(Storage.delete_file, file_path)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=ERROR_MESSAGES.FILE_TOO_LARGE(size=f'{max_size} MB'),
|
||||
)
|
||||
|
||||
# SHA-256 of raw uploaded bytes for incremental sync diffing.
|
||||
# If the client pre-computed and sent file_hash, use that.
|
||||
file_hash = file_metadata.get('file_hash') or hashlib.sha256(contents).hexdigest()
|
||||
file_hash = file_metadata.get('file_hash') or await asyncio.to_thread(
|
||||
lambda: hashlib.sha256(contents).hexdigest()
|
||||
)
|
||||
|
||||
file_item = await Files.insert_new_file(
|
||||
user.id,
|
||||
@@ -443,13 +510,29 @@ async def search_files(
|
||||
return files
|
||||
|
||||
|
||||
############################
|
||||
# Count Files
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/count', response_model=int)
|
||||
async def count_files(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
user_id = None if (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL) else user.id
|
||||
return await Files.count_files_by_user_id(user_id=user_id, db=db)
|
||||
|
||||
|
||||
############################
|
||||
# Delete All Files
|
||||
############################
|
||||
|
||||
|
||||
@router.delete('/all')
|
||||
async def delete_all_files(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_all_files(
|
||||
request: Request, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
result = await Files.delete_all_files(db=db)
|
||||
if result:
|
||||
try:
|
||||
@@ -462,6 +545,7 @@ async def delete_all_files(user=Depends(get_admin_user), db: AsyncSession = Depe
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error deleting files'),
|
||||
)
|
||||
await publish_event(request, EVENTS.FILE_DELETED_ALL, actor=user, subject_type='file')
|
||||
return {'message': 'All files deleted successfully'}
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -605,6 +689,12 @@ async def update_file_data_content_by_id(
|
||||
)
|
||||
|
||||
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db):
|
||||
max_size = await Config.get('rag.file.max_size')
|
||||
if max_size and len(form_data.content.encode('utf-8')) > int(max_size) * 1024 * 1024:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=ERROR_MESSAGES.FILE_TOO_LARGE(size=f'{max_size} MB'),
|
||||
)
|
||||
try:
|
||||
await process_file(
|
||||
request,
|
||||
@@ -623,18 +713,29 @@ async def update_file_data_content_by_id(
|
||||
knowledges = await Knowledges.get_knowledges_by_file_id(id, db=db)
|
||||
for knowledge in knowledges:
|
||||
try:
|
||||
# Remove old embeddings for this file from the KB collection
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, filter={'file_id': id})
|
||||
# Re-add from the now-updated file-{file_id} collection
|
||||
old_vectors = await ASYNC_VECTOR_DB_CLIENT.query(collection_name=knowledge.id, filter={'file_id': id})
|
||||
old_vector_ids = old_vectors.ids[0] if old_vectors and old_vectors.ids else []
|
||||
|
||||
# Re-add from the now-updated file-{file_id} collection before
|
||||
# removing old vectors, so a failed reindex keeps the KB usable.
|
||||
await process_file(
|
||||
request,
|
||||
ProcessFileForm(file_id=id, collection_name=knowledge.id),
|
||||
user=user,
|
||||
db=db,
|
||||
)
|
||||
if old_vector_ids:
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=knowledge.id, ids=old_vector_ids)
|
||||
except Exception as e:
|
||||
log.warning(f'Failed to update knowledge {knowledge.id} after content change for file {id}: {e}')
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FILE_CONTENT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'content_preview': form_data.content[:300]},
|
||||
)
|
||||
return {'content': file.data.get('content', '')}
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -824,6 +925,7 @@ class FileRenameForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/rename')
|
||||
async def rename_file_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: FileRenameForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -840,6 +942,13 @@ async def rename_file_by_id(
|
||||
if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'write', user, db=db):
|
||||
result = await Files.update_file_name_by_id(id, form_data.filename, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FILE_RENAMED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'filename': form_data.filename},
|
||||
)
|
||||
return result
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -859,7 +968,9 @@ async def rename_file_by_id(
|
||||
|
||||
|
||||
@router.delete('/{id}')
|
||||
async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_file_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
file = await Files.get_file_by_id(id, db=db)
|
||||
|
||||
if not file:
|
||||
@@ -894,6 +1005,13 @@ async def delete_file_by_id(id: str, user=Depends(get_verified_user), db: AsyncS
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error deleting files'),
|
||||
)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FILE_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'filename': file.filename},
|
||||
)
|
||||
return {'message': 'File deleted successfully'}
|
||||
else:
|
||||
raise HTTPException(
|
||||
|
||||
@@ -10,7 +10,9 @@ from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile
|
||||
from fastapi.responses import FileResponse, StreamingResponse
|
||||
from open_webui.config import UPLOAD_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.folders import (
|
||||
FolderForm,
|
||||
@@ -19,7 +21,13 @@ from open_webui.models.folders import (
|
||||
Folders,
|
||||
FolderUpdateForm,
|
||||
)
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.access_control import (
|
||||
filter_allowed_access_grants,
|
||||
)
|
||||
from open_webui.utils.access_control.files import get_accessible_folder_files
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from pydantic import BaseModel
|
||||
@@ -31,6 +39,29 @@ log = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
from open_webui.utils.access_control.folders import has_folder_access as _has_folder_access
|
||||
|
||||
|
||||
async def check_folders_permission(request: Request, user, db=None):
|
||||
"""Verify the folders feature is enabled and the user has permission."""
|
||||
config = await Config.get_many('folders.enable', 'user.permissions')
|
||||
if config.get('folders.enable') is False:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'features.folders',
|
||||
config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# Get Folders
|
||||
############################
|
||||
@@ -42,22 +73,7 @@ async def get_folders(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if request.app.state.config.ENABLE_FOLDERS is False:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'features.folders',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_folders_permission(request, user, db=db)
|
||||
|
||||
folders = await Folders.get_folders_by_user_id(user.id, db=db)
|
||||
|
||||
@@ -87,10 +103,12 @@ async def get_folders(
|
||||
|
||||
@router.post('/')
|
||||
async def create_folder(
|
||||
request: Request,
|
||||
form_data: FolderForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(
|
||||
form_data.parent_id, user.id, form_data.name, db=db
|
||||
)
|
||||
@@ -101,8 +119,43 @@ async def create_folder(
|
||||
detail=ERROR_MESSAGES.DEFAULT('Folder already exists'),
|
||||
)
|
||||
|
||||
# Check if creating a subfolder in a shared folder
|
||||
if form_data.parent_id:
|
||||
parent = await Folders.get_folder_by_id(form_data.parent_id, db=db)
|
||||
if parent and parent.user_id != user.id:
|
||||
# Creating subfolder in someone else's shared folder
|
||||
if user.role != 'admin' and not await _has_folder_access(user.id, parent, 'write', db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
# Create as the folder owner's subfolder (keep tree consistent)
|
||||
try:
|
||||
folder = await Folders.insert_new_folder(parent.user_id, form_data, form_data.parent_id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FOLDER_CREATED,
|
||||
actor=user,
|
||||
subject_id=folder.id,
|
||||
data={'name': folder.name, 'parent_id': folder.parent_id, 'owner_id': folder.user_id},
|
||||
)
|
||||
return folder
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error creating folder'),
|
||||
)
|
||||
|
||||
try:
|
||||
folder = await Folders.insert_new_folder(user.id, form_data, form_data.parent_id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FOLDER_CREATED,
|
||||
actor=user,
|
||||
subject_id=folder.id,
|
||||
data={'name': folder.name, 'parent_id': folder.parent_id, 'owner_id': folder.user_id},
|
||||
)
|
||||
return folder
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -113,21 +166,90 @@ async def create_folder(
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# Get Shared Folders
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/shared')
|
||||
async def get_shared_folders(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
"""Get all folders shared with the current user (not owned by them)."""
|
||||
await check_folders_permission(request, user, db=db)
|
||||
groups = await Groups.get_groups_by_member_id(user.id, db=db)
|
||||
group_ids = {g.id for g in groups}
|
||||
|
||||
folder_perms = await Folders.get_shared_folder_ids_for_user(user.id, group_ids, db=db)
|
||||
|
||||
# Filter out folders owned by the user
|
||||
results = []
|
||||
owner_cache = {}
|
||||
for folder_id, permission in folder_perms.items():
|
||||
folder = await Folders.get_folder_by_id(folder_id, db=db)
|
||||
if not folder or folder.user_id == user.id:
|
||||
continue
|
||||
|
||||
# Get owner name (cached)
|
||||
if folder.user_id not in owner_cache:
|
||||
owner = await Users.get_user_by_id(folder.user_id, db=db)
|
||||
owner_cache[folder.user_id] = owner.name if owner else 'Unknown'
|
||||
|
||||
results.append(
|
||||
{
|
||||
**folder.model_dump(),
|
||||
'owner_name': owner_cache[folder.user_id],
|
||||
'permission': permission,
|
||||
}
|
||||
)
|
||||
|
||||
# Also include child folders of shared folders (inheritance)
|
||||
shared_root_ids = {r['id'] for r in results}
|
||||
for root_id in list(shared_root_ids):
|
||||
root_folder = await Folders.get_folder_by_id(root_id, db=db)
|
||||
if root_folder:
|
||||
children = await Folders.get_children_folders_by_id_and_user_id(root_id, root_folder.user_id, db=db)
|
||||
if children:
|
||||
for child in children:
|
||||
if child.id not in {r['id'] for r in results}:
|
||||
results.append(
|
||||
{
|
||||
**child.model_dump(),
|
||||
'owner_name': owner_cache.get(child.user_id, 'Unknown'),
|
||||
'permission': folder_perms.get(root_id, 'read'),
|
||||
}
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
############################
|
||||
# Get Folders By Id
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/{id}', response_model=Optional[FolderModel])
|
||||
async def get_folder_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
@router.get('/{id}', response_model=None)
|
||||
async def get_folder_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
|
||||
if folder:
|
||||
return folder
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
grants = await AccessGrants.get_grants_by_resource('folder', id, db=db)
|
||||
return {**folder.model_dump(), 'access_grants': [g.model_dump() for g in grants]}
|
||||
|
||||
# Check shared access
|
||||
folder = await Folders.get_folder_by_id(id, db=db)
|
||||
if folder and (user.role == 'admin' or await _has_folder_access(user.id, folder, 'read', db)):
|
||||
grants = await AccessGrants.get_grants_by_resource('folder', id, db=db)
|
||||
return {**folder.model_dump(), 'access_grants': [g.model_dump() for g in grants]}
|
||||
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
@@ -137,17 +259,28 @@ async def get_folder_by_id(id: str, user=Depends(get_verified_user), db: AsyncSe
|
||||
|
||||
@router.post('/{id}/update')
|
||||
async def update_folder_name_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: FolderUpdateForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
|
||||
if not folder:
|
||||
# Check shared write access
|
||||
folder = await Folders.get_folder_by_id(id, db=db)
|
||||
if not folder or (user.role != 'admin' and not await _has_folder_access(user.id, folder, 'write', db)):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if folder:
|
||||
if form_data.name is not None:
|
||||
# Check if folder with same name exists
|
||||
existing_folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(
|
||||
folder.parent_id, user.id, form_data.name, db=db
|
||||
folder.parent_id, folder.user_id, form_data.name, db=db
|
||||
)
|
||||
if existing_folder and existing_folder.id != id:
|
||||
raise HTTPException(
|
||||
@@ -166,7 +299,14 @@ async def update_folder_name_by_id(
|
||||
)
|
||||
|
||||
try:
|
||||
folder = await Folders.update_folder_by_id_and_user_id(id, user.id, form_data, db=db)
|
||||
folder = await Folders.update_folder_by_id_and_user_id(id, folder.user_id, form_data, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FOLDER_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': folder.name},
|
||||
)
|
||||
return folder
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -175,11 +315,6 @@ async def update_folder_name_by_id(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error updating folder'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
@@ -193,11 +328,13 @@ class FolderParentIdForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/update/parent')
|
||||
async def update_folder_parent_id_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: FolderParentIdForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
|
||||
if folder:
|
||||
existing_folder = await Folders.get_folder_by_parent_id_and_user_id_and_name(
|
||||
@@ -212,6 +349,13 @@ async def update_folder_parent_id_by_id(
|
||||
|
||||
try:
|
||||
folder = await Folders.update_folder_parent_id_by_id_and_user_id(id, user.id, form_data.parent_id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FOLDER_PARENT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'parent_id': form_data.parent_id},
|
||||
)
|
||||
return folder
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -238,11 +382,13 @@ class FolderIsExpandedForm(BaseModel):
|
||||
|
||||
@router.post('/{id}/update/expanded')
|
||||
async def update_folder_is_expanded_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: FolderIsExpandedForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
|
||||
if folder:
|
||||
try:
|
||||
@@ -264,6 +410,113 @@ async def update_folder_is_expanded_by_id(
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# Update Folder Access By Id
|
||||
############################
|
||||
|
||||
|
||||
class FolderAccessGrantsForm(BaseModel):
|
||||
access_grants: list[dict]
|
||||
|
||||
|
||||
@router.post('/{id}/access/update')
|
||||
async def update_folder_access_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: FolderAccessGrantsForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id(id, db=db)
|
||||
if not folder:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
# Only owner, admin, or write-granted user can update access
|
||||
if user.role != 'admin' and user.id != folder.user_id:
|
||||
if not await _has_folder_access(user.id, folder, 'write', db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
None,
|
||||
db=db,
|
||||
)
|
||||
|
||||
await AccessGrants.set_access_grants('folder', id, form_data.access_grants, db=db)
|
||||
|
||||
grants = await AccessGrants.get_grants_by_resource('folder', id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FOLDER_ACCESS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'grant_count': len(grants)},
|
||||
)
|
||||
return {
|
||||
**folder.model_dump(),
|
||||
'access_grants': [g.model_dump() for g in grants],
|
||||
}
|
||||
|
||||
|
||||
############################
|
||||
# Get Shared Folder Chats
|
||||
############################
|
||||
|
||||
|
||||
@router.get('/{id}/shared/chats')
|
||||
async def get_shared_folder_chats(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
"""Get chats within a shared folder. Returns readonly flag based on permission."""
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id(id, db=db)
|
||||
if not folder:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
is_owner = user.id == folder.user_id
|
||||
is_admin = user.role == 'admin'
|
||||
has_write = is_owner or is_admin or await _has_folder_access(user.id, folder, 'write', db)
|
||||
has_read = has_write or await _has_folder_access(user.id, folder, 'read', db)
|
||||
|
||||
if not has_read:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
chats = await Chats.get_all_chats_by_folder_id(id, db=db)
|
||||
|
||||
# Resolve owner names for display (avatar URLs are constructed client-side)
|
||||
owner_cache: dict[str, str] = {}
|
||||
for chat in chats:
|
||||
uid = chat['user_id']
|
||||
if uid not in owner_cache:
|
||||
u = await Users.get_user_by_id(uid, db=db)
|
||||
owner_cache[uid] = u.name if u else 'Unknown'
|
||||
chat['owner_name'] = owner_cache[uid]
|
||||
|
||||
return {
|
||||
'chats': [{**chat, 'readonly': chat['user_id'] != user.id} for chat in chats],
|
||||
'folder_permission': 'write' if has_write else 'read',
|
||||
}
|
||||
|
||||
|
||||
############################
|
||||
# Delete Folder By Id
|
||||
############################
|
||||
@@ -277,9 +530,37 @@ async def delete_folder_by_id(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if await Chats.count_chats_by_folder_id_and_user_id(id, user.id, db=db):
|
||||
await check_folders_permission(request, user, db=db)
|
||||
folder = await Folders.get_folder_by_id_and_user_id(id, user.id, db=db)
|
||||
|
||||
if not folder:
|
||||
# Check if it's a shared subfolder with write access
|
||||
folder = await Folders.get_folder_by_id(id, db=db)
|
||||
if folder and folder.parent_id:
|
||||
if user.role != 'admin' and not await _has_folder_access(user.id, folder, 'write', db):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
elif folder and not folder.parent_id:
|
||||
# Root shared folders can only be deleted by owner/admin
|
||||
if user.role != 'admin':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
folder_owner_id = folder.user_id
|
||||
|
||||
folder_ids = await Folders.get_folder_ids_by_id_and_user_id_in_subtree(id, folder_owner_id, db=db)
|
||||
if await Chats.count_chats_by_folder_ids_and_user_id(folder_ids, folder_owner_id, db=db):
|
||||
chat_delete_permission = await has_permission(
|
||||
user.id, 'chat.delete', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'chat.delete', await Config.get('user.permissions'), db=db
|
||||
)
|
||||
if user.role != 'admin' and not chat_delete_permission:
|
||||
raise HTTPException(
|
||||
@@ -288,19 +569,29 @@ async def delete_folder_by_id(
|
||||
)
|
||||
|
||||
folders = []
|
||||
folders.append(await Folders.get_folder_by_id_and_user_id(id, user.id, db=db))
|
||||
folders.append(folder)
|
||||
while folders:
|
||||
folder = folders.pop()
|
||||
if folder:
|
||||
try:
|
||||
folder_ids = await Folders.delete_folder_by_id_and_user_id(folder.id, user.id, db=db)
|
||||
folder_ids = await Folders.delete_folder_by_id_and_user_id(folder.id, folder_owner_id, db=db)
|
||||
|
||||
for folder_id in folder_ids:
|
||||
if delete_contents:
|
||||
await Chats.delete_chats_by_user_id_and_folder_id(user.id, folder_id, db=db)
|
||||
await Chats.delete_chats_by_user_id_and_folder_id(folder_owner_id, folder_id, db=db)
|
||||
else:
|
||||
await Chats.move_chats_by_user_id_and_folder_id(user.id, folder_id, None, db=db)
|
||||
await Chats.move_chats_by_user_id_and_folder_id(folder_owner_id, folder_id, None, db=db)
|
||||
|
||||
# Clean up access grants for this folder
|
||||
await AccessGrants.revoke_all_access('folder', folder_id, db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FOLDER_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'folder_ids': folder_ids, 'delete_contents': delete_contents},
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -311,7 +602,7 @@ async def delete_folder_by_id(
|
||||
)
|
||||
finally:
|
||||
# Get all subfolders
|
||||
subfolders = await Folders.get_folders_by_parent_id_and_user_id(folder.id, user.id, db=db)
|
||||
subfolders = await Folders.get_folders_by_parent_id_and_user_id(folder.id, folder_owner_id, db=db)
|
||||
folders.extend(subfolders)
|
||||
|
||||
else:
|
||||
|
||||
@@ -11,6 +11,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.config import CACHE_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.functions import (
|
||||
FunctionForm,
|
||||
@@ -22,6 +23,7 @@ from open_webui.models.functions import (
|
||||
)
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.plugin import (
|
||||
get_functions_cache,
|
||||
get_function_module_from_cache,
|
||||
load_function_module_by_id,
|
||||
replace_imports,
|
||||
@@ -130,8 +132,13 @@ async def load_function_from_url(request: Request, form_data: LoadUrlForm, user=
|
||||
'name': function_name,
|
||||
'content': data,
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.DEFAULT(e))
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error fetching function'),
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
@@ -171,7 +178,7 @@ async def sync_functions(
|
||||
log.exception(f'Failed to load a function: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error loading function'),
|
||||
)
|
||||
|
||||
|
||||
@@ -205,7 +212,7 @@ async def create_new_function(
|
||||
)
|
||||
form_data.meta.manifest = frontmatter
|
||||
|
||||
FUNCTIONS = request.app.state.FUNCTIONS
|
||||
FUNCTIONS = get_functions_cache(request)
|
||||
FUNCTIONS[form_data.id] = function_module
|
||||
|
||||
function = await Functions.insert_new_function(user.id, function_type, form_data, db=db)
|
||||
@@ -217,17 +224,26 @@ async def create_new_function(
|
||||
await Functions.update_function_metadata_by_id(form_data.id, {'toggle': True}, db=db)
|
||||
|
||||
if function:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_CREATED,
|
||||
actor=user,
|
||||
subject_id=function.id,
|
||||
data={'type': function.type, 'name': function.name},
|
||||
)
|
||||
return function
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error creating function'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to create a new function: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error creating function'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -260,12 +276,25 @@ async def get_function_by_id(id: str, user=Depends(get_admin_user), db: AsyncSes
|
||||
|
||||
|
||||
@router.post('/id/{id}/toggle', response_model=FunctionModel | None)
|
||||
async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def toggle_function_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
function = await Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function = await Functions.update_function_by_id(id, {'is_active': not function.is_active}, db=db)
|
||||
|
||||
if function:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_ENABLED if function.is_active else EVENTS.FUNCTION_DISABLED,
|
||||
actor=user,
|
||||
subject_id=function.id,
|
||||
subject_type='function',
|
||||
data={'type': function.type, 'name': function.name},
|
||||
)
|
||||
return function
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -285,12 +314,24 @@ async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: Async
|
||||
|
||||
|
||||
@router.post('/id/{id}/toggle/global', response_model=FunctionModel | None)
|
||||
async def toggle_global_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def toggle_global_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
function = await Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function = await Functions.update_function_by_id(id, {'is_global': not function.is_global}, db=db)
|
||||
|
||||
if function:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_UPDATED,
|
||||
actor=user,
|
||||
subject_id=function.id,
|
||||
data={'type': function.type, 'name': function.name, 'is_global': function.is_global},
|
||||
)
|
||||
return function
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -322,7 +363,7 @@ async def update_function_by_id(
|
||||
function_module, function_type, frontmatter = await load_function_module_by_id(id, content=form_data.content)
|
||||
form_data.meta.manifest = frontmatter
|
||||
|
||||
FUNCTIONS = request.app.state.FUNCTIONS
|
||||
FUNCTIONS = get_functions_cache(request)
|
||||
FUNCTIONS[id] = function_module
|
||||
|
||||
updated = {**form_data.model_dump(exclude={'id'}), 'type': function_type}
|
||||
@@ -334,6 +375,13 @@ async def update_function_by_id(
|
||||
await Functions.update_function_metadata_by_id(id, {'toggle': True}, db=db)
|
||||
|
||||
if function:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_UPDATED,
|
||||
actor=user,
|
||||
subject_id=function.id,
|
||||
data={'type': function.type, 'name': function.name},
|
||||
)
|
||||
return function
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -341,10 +389,12 @@ async def update_function_by_id(
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error updating function'),
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating function'),
|
||||
)
|
||||
|
||||
|
||||
@@ -363,9 +413,14 @@ async def delete_function_by_id(
|
||||
result = await Functions.delete_function_by_id(id, db=db)
|
||||
|
||||
if result:
|
||||
FUNCTIONS = request.app.state.FUNCTIONS
|
||||
if id in FUNCTIONS:
|
||||
del FUNCTIONS[id]
|
||||
FUNCTIONS = get_functions_cache(request)
|
||||
FUNCTIONS.pop(id, None)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@@ -387,7 +442,7 @@ async def get_function_valves_by_id(
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error getting function valves'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -452,12 +507,18 @@ async def update_function_valves_by_id(
|
||||
|
||||
valves_dict = valves.model_dump(exclude_unset=True)
|
||||
await Functions.update_function_valves_by_id(id, valves_dict, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_VALVES_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
)
|
||||
return valves_dict
|
||||
except Exception as e:
|
||||
log.exception(f'Error updating function values by id {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating function valves'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -489,7 +550,7 @@ async def get_function_user_valves_by_id(
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error getting function user valves'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -544,12 +605,19 @@ async def update_function_user_valves_by_id(
|
||||
user_valves = UserValves(**form_data)
|
||||
user_valves_dict = user_valves.model_dump(exclude_unset=True)
|
||||
await Functions.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.FUNCTION_VALVES_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'scope': 'user'},
|
||||
)
|
||||
return user_valves_dict
|
||||
except Exception as e:
|
||||
log.exception(f'Error updating function user valves by id {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating function user valves'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.config import CACHE_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.groups import (
|
||||
@@ -58,6 +59,7 @@ async def get_groups(
|
||||
|
||||
@router.post('/create', response_model=Optional[GroupResponse])
|
||||
async def create_new_group(
|
||||
request: Request,
|
||||
form_data: GroupForm,
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
@@ -65,6 +67,13 @@ async def create_new_group(
|
||||
try:
|
||||
group = await Groups.insert_new_group(user.id, form_data, db=db)
|
||||
if group:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_CREATED,
|
||||
actor=user,
|
||||
subject_id=group.id,
|
||||
data={'name': group.name},
|
||||
)
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=await Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
@@ -74,11 +83,13 @@ async def create_new_group(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error creating group'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Error creating a new group: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error creating group'),
|
||||
)
|
||||
|
||||
|
||||
@@ -157,7 +168,7 @@ async def get_users_in_group(id: str, user=Depends(get_admin_user), db: AsyncSes
|
||||
log.exception(f'Error adding users to group {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error getting group members'),
|
||||
)
|
||||
|
||||
|
||||
@@ -168,6 +179,7 @@ async def get_users_in_group(id: str, user=Depends(get_admin_user), db: AsyncSes
|
||||
|
||||
@router.post('/id/{id}/update', response_model=Optional[GroupResponse])
|
||||
async def update_group_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: GroupUpdateForm,
|
||||
user=Depends(get_admin_user),
|
||||
@@ -176,6 +188,13 @@ async def update_group_by_id(
|
||||
try:
|
||||
group = await Groups.update_group_by_id(id, form_data, db=db)
|
||||
if group:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': group.name},
|
||||
)
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=await Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
@@ -185,11 +204,13 @@ async def update_group_by_id(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error updating group'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Error updating group {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating group'),
|
||||
)
|
||||
|
||||
|
||||
@@ -200,6 +221,7 @@ async def update_group_by_id(
|
||||
|
||||
@router.post('/id/{id}/users/add', response_model=Optional[GroupResponse])
|
||||
async def add_user_to_group(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: UserIdsForm,
|
||||
user=Depends(get_admin_user),
|
||||
@@ -211,6 +233,13 @@ async def add_user_to_group(
|
||||
|
||||
group = await Groups.add_users_to_group(id, form_data.user_ids, db=db)
|
||||
if group:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_ADDED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'user_ids': form_data.user_ids},
|
||||
)
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=await Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
@@ -220,16 +249,19 @@ async def add_user_to_group(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error adding users to group'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Error adding users to group {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error adding users to group'),
|
||||
)
|
||||
|
||||
|
||||
@router.post('/id/{id}/users/remove', response_model=Optional[GroupResponse])
|
||||
async def remove_users_from_group(
|
||||
request: Request,
|
||||
id: str,
|
||||
form_data: UserIdsForm,
|
||||
user=Depends(get_admin_user),
|
||||
@@ -238,6 +270,13 @@ async def remove_users_from_group(
|
||||
try:
|
||||
group = await Groups.remove_users_from_group(id, form_data.user_ids, db=db)
|
||||
if group:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_REMOVED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'user_ids': form_data.user_ids},
|
||||
)
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=await Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
@@ -247,11 +286,13 @@ async def remove_users_from_group(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error removing users from group'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Error removing users from group {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error removing users from group'),
|
||||
)
|
||||
|
||||
|
||||
@@ -261,21 +302,31 @@ async def remove_users_from_group(
|
||||
|
||||
|
||||
@router.delete('/id/{id}/delete', response_model=bool)
|
||||
async def delete_group_by_id(id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_group_by_id(
|
||||
request: Request, id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
try:
|
||||
result = await Groups.delete_group_by_id(id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
)
|
||||
return result
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error deleting group'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Error deleting group {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error deleting group'),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -9,22 +9,27 @@ import mimetypes
|
||||
import re
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Optional
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
import aiohttp
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
from PIL import Image, ImageOps
|
||||
from open_webui.config import (
|
||||
CACHE_DIR,
|
||||
ENABLE_OPENAI_IMAGE_EDIT_NORMALIZATION,
|
||||
IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN,
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_ALLOW_REDIRECTS, AIOHTTP_CLIENT_SESSION_SSL, ENABLE_FORWARD_USER_INFO_HEADERS
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.retrieval.web.utils import validate_url
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.retrieval.web.utils import get_ssrf_safe_session, validate_url
|
||||
from open_webui.routers.files import get_file_content_by_id, upload_file_handler
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
@@ -50,17 +55,130 @@ IMAGE_CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
IMAGE_FILE_EXTENSIONS = {
|
||||
'image/jpeg': '.jpg',
|
||||
'image/jpg': '.jpg',
|
||||
'image/mpo': '.jpg',
|
||||
'image/png': '.png',
|
||||
'image/webp': '.webp',
|
||||
}
|
||||
|
||||
IMAGE_CONFIG_KEYS = {
|
||||
'ENABLE_IMAGE_GENERATION': 'image_generation.enable',
|
||||
'ENABLE_IMAGE_PROMPT_GENERATION': 'image_generation.prompt.enable',
|
||||
'IMAGE_GENERATION_ENGINE': 'image_generation.engine',
|
||||
'IMAGE_GENERATION_MODEL': 'image_generation.model',
|
||||
'IMAGE_SIZE': 'image_generation.size',
|
||||
'IMAGE_STEPS': 'image_generation.steps',
|
||||
'IMAGES_OPENAI_API_BASE_URL': 'image_generation.openai.api_base_url',
|
||||
'IMAGES_OPENAI_API_KEY': 'image_generation.openai.api_key',
|
||||
'IMAGES_OPENAI_API_VERSION': 'image_generation.openai.api_version',
|
||||
'IMAGES_OPENAI_API_PARAMS': 'image_generation.openai.params',
|
||||
'AUTOMATIC1111_BASE_URL': 'image_generation.automatic1111.base_url',
|
||||
'AUTOMATIC1111_API_AUTH': 'image_generation.automatic1111.api_auth',
|
||||
'AUTOMATIC1111_PARAMS': 'image_generation.automatic1111.api_params',
|
||||
'COMFYUI_BASE_URL': 'image_generation.comfyui.base_url',
|
||||
'COMFYUI_API_KEY': 'image_generation.comfyui.api_key',
|
||||
'COMFYUI_WORKFLOW': 'image_generation.comfyui.workflow',
|
||||
'COMFYUI_WORKFLOW_NODES': 'image_generation.comfyui.nodes',
|
||||
'IMAGES_GEMINI_API_BASE_URL': 'image_generation.gemini.api_base_url',
|
||||
'IMAGES_GEMINI_API_KEY': 'image_generation.gemini.api_key',
|
||||
'IMAGES_GEMINI_ENDPOINT_METHOD': 'image_generation.gemini.endpoint_method',
|
||||
'ENABLE_IMAGE_EDIT': 'images.edit.enable',
|
||||
'IMAGE_EDIT_ENGINE': 'images.edit.engine',
|
||||
'IMAGE_EDIT_MODEL': 'images.edit.model',
|
||||
'IMAGE_EDIT_SIZE': 'images.edit.size',
|
||||
'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_GEMINI_API_BASE_URL': 'images.edit.gemini.api_base_url',
|
||||
'IMAGES_EDIT_GEMINI_API_KEY': 'images.edit.gemini.api_key',
|
||||
'IMAGES_EDIT_COMFYUI_BASE_URL': 'images.edit.comfyui.base_url',
|
||||
'IMAGES_EDIT_COMFYUI_API_KEY': 'images.edit.comfyui.api_key',
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW': 'images.edit.comfyui.workflow',
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES': 'images.edit.comfyui.nodes',
|
||||
'USER_PERMISSIONS': 'user.permissions',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
async def get_image_config() -> SimpleNamespace:
|
||||
return SimpleNamespace(**await get_config_values(IMAGE_CONFIG_KEYS))
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
def normalize_openai_edit_image_data_url(data_url: str) -> str:
|
||||
if not data_url.startswith('data:') or ',' not in data_url:
|
||||
return data_url
|
||||
|
||||
header, encoded = data_url.split(',', 1)
|
||||
mime_type = header.split(';')[0].lstrip('data:').lower()
|
||||
if mime_type not in {'image/jpeg', 'image/jpg', 'image/mpo'}:
|
||||
return data_url
|
||||
|
||||
try:
|
||||
image_bytes = base64.b64decode(encoded)
|
||||
with Image.open(io.BytesIO(image_bytes)) as image:
|
||||
orientation = image.getexif().get(274)
|
||||
needs_normalization = (
|
||||
mime_type == 'image/mpo'
|
||||
or image.format == 'MPO'
|
||||
or getattr(image, 'n_frames', 1) > 1
|
||||
or orientation not in (None, 1)
|
||||
or image.mode not in ('RGB', 'L')
|
||||
)
|
||||
|
||||
if not needs_normalization:
|
||||
return data_url
|
||||
|
||||
image.seek(0)
|
||||
image = ImageOps.exif_transpose(image)
|
||||
if image.mode != 'RGB':
|
||||
image = image.convert('RGB')
|
||||
|
||||
output = io.BytesIO()
|
||||
image.save(output, format='JPEG', quality=95)
|
||||
normalized_image = base64.b64encode(output.getvalue()).decode('utf-8')
|
||||
return f'data:image/jpeg;base64,{normalized_image}'
|
||||
except Exception as e:
|
||||
log.debug(f'Image edit normalization skipped: {e}')
|
||||
|
||||
return data_url
|
||||
|
||||
|
||||
def get_image_file_item(base64_string, param_name='image'):
|
||||
header, encoded = base64_string.split(',', 1)
|
||||
mime_type = header.split(';')[0].lstrip('data:') or 'image/png'
|
||||
image_data = base64.b64decode(encoded)
|
||||
extension = IMAGE_FILE_EXTENSIONS.get(mime_type.lower()) or mimetypes.guess_extension(mime_type) or '.png'
|
||||
return (
|
||||
param_name,
|
||||
(
|
||||
f'{uuid.uuid4()}{extension}',
|
||||
io.BytesIO(image_data),
|
||||
mime_type,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def set_image_model(request: Request, model: str):
|
||||
log.info(f'Setting image model to {model}')
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL = model
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE in ['', 'automatic1111']:
|
||||
api_auth = get_automatic1111_api_auth(request)
|
||||
await Config.upsert({'image_generation.model': model})
|
||||
image_config = await get_image_config()
|
||||
if image_config.IMAGE_GENERATION_ENGINE in ['', 'automatic1111']:
|
||||
api_auth = get_automatic1111_api_auth(image_config)
|
||||
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': api_auth},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
@@ -68,7 +186,7 @@ async def set_image_model(request: Request, model: str):
|
||||
if model != options['sd_model_checkpoint']:
|
||||
options['sd_model_checkpoint'] = model
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
json=options,
|
||||
headers={'authorization': api_auth},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -77,41 +195,33 @@ async def set_image_model(request: Request, model: str):
|
||||
except Exception as e:
|
||||
log.debug(f'{e}')
|
||||
|
||||
return request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
return image_config.IMAGE_GENERATION_MODEL
|
||||
|
||||
|
||||
async def get_image_model(request):
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
return (
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
if request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
else 'dall-e-2'
|
||||
)
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
return (
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
if request.app.state.config.IMAGE_GENERATION_MODEL
|
||||
else 'imagen-3.0-generate-002'
|
||||
)
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
return (
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL if request.app.state.config.IMAGE_GENERATION_MODEL else ''
|
||||
)
|
||||
elif (
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
|
||||
):
|
||||
image_config = await get_image_config()
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
return image_config.IMAGE_GENERATION_MODEL if image_config.IMAGE_GENERATION_MODEL else 'dall-e-2'
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
return image_config.IMAGE_GENERATION_MODEL if image_config.IMAGE_GENERATION_MODEL else 'imagen-3.0-generate-002'
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
return image_config.IMAGE_GENERATION_MODEL if image_config.IMAGE_GENERATION_MODEL else ''
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'automatic1111' or image_config.IMAGE_GENERATION_ENGINE == '':
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
options = await r.json()
|
||||
return options['sd_model_checkpoint']
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e))
|
||||
log.exception(f'Failed to get default model from automatic1111: {e}')
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Failed to connect to the image generation engine'),
|
||||
)
|
||||
|
||||
|
||||
class ImagesConfig(BaseModel):
|
||||
@@ -159,52 +269,11 @@ class ImagesConfig(BaseModel):
|
||||
|
||||
@router.get('/config', response_model=ImagesConfig)
|
||||
async def get_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_IMAGE_GENERATION': request.app.state.config.ENABLE_IMAGE_GENERATION,
|
||||
'ENABLE_IMAGE_PROMPT_GENERATION': request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION,
|
||||
'IMAGE_GENERATION_ENGINE': request.app.state.config.IMAGE_GENERATION_ENGINE,
|
||||
'IMAGE_GENERATION_MODEL': request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
'IMAGE_SIZE': request.app.state.config.IMAGE_SIZE,
|
||||
'IMAGE_STEPS': request.app.state.config.IMAGE_STEPS,
|
||||
'IMAGES_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_OPENAI_API_BASE_URL,
|
||||
'IMAGES_OPENAI_API_KEY': request.app.state.config.IMAGES_OPENAI_API_KEY,
|
||||
'IMAGES_OPENAI_API_VERSION': request.app.state.config.IMAGES_OPENAI_API_VERSION,
|
||||
'IMAGES_OPENAI_API_PARAMS': request.app.state.config.IMAGES_OPENAI_API_PARAMS,
|
||||
'AUTOMATIC1111_BASE_URL': request.app.state.config.AUTOMATIC1111_BASE_URL,
|
||||
'AUTOMATIC1111_API_AUTH': request.app.state.config.AUTOMATIC1111_API_AUTH,
|
||||
'AUTOMATIC1111_PARAMS': request.app.state.config.AUTOMATIC1111_PARAMS,
|
||||
'COMFYUI_BASE_URL': request.app.state.config.COMFYUI_BASE_URL,
|
||||
'COMFYUI_API_KEY': request.app.state.config.COMFYUI_API_KEY,
|
||||
'COMFYUI_WORKFLOW': request.app.state.config.COMFYUI_WORKFLOW,
|
||||
'COMFYUI_WORKFLOW_NODES': request.app.state.config.COMFYUI_WORKFLOW_NODES,
|
||||
'IMAGES_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_GEMINI_API_BASE_URL,
|
||||
'IMAGES_GEMINI_API_KEY': request.app.state.config.IMAGES_GEMINI_API_KEY,
|
||||
'IMAGES_GEMINI_ENDPOINT_METHOD': request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD,
|
||||
'ENABLE_IMAGE_EDIT': request.app.state.config.ENABLE_IMAGE_EDIT,
|
||||
'IMAGE_EDIT_ENGINE': request.app.state.config.IMAGE_EDIT_ENGINE,
|
||||
'IMAGE_EDIT_MODEL': request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
'IMAGE_EDIT_SIZE': request.app.state.config.IMAGE_EDIT_SIZE,
|
||||
'IMAGES_EDIT_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL,
|
||||
'IMAGES_EDIT_OPENAI_API_KEY': request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY,
|
||||
'IMAGES_EDIT_OPENAI_API_VERSION': request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION,
|
||||
'IMAGES_EDIT_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL,
|
||||
'IMAGES_EDIT_GEMINI_API_KEY': request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_BASE_URL': request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
'IMAGES_EDIT_COMFYUI_API_KEY': request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
return await get_config_values(IMAGE_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_config(request: Request, form_data: ImagesConfig, user=Depends(get_admin_user)):
|
||||
request.app.state.config.ENABLE_IMAGE_GENERATION = form_data.ENABLE_IMAGE_GENERATION
|
||||
|
||||
# Create Image
|
||||
request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION = form_data.ENABLE_IMAGE_PROMPT_GENERATION
|
||||
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE = form_data.IMAGE_GENERATION_ENGINE
|
||||
await set_image_model(request, form_data.IMAGE_GENERATION_MODEL)
|
||||
if form_data.IMAGE_SIZE == 'auto' and not re.match(
|
||||
IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN, form_data.IMAGE_GENERATION_MODEL
|
||||
):
|
||||
@@ -216,100 +285,44 @@ async def update_config(request: Request, form_data: ImagesConfig, user=Depends(
|
||||
)
|
||||
|
||||
pattern = r'^\d+x\d+$'
|
||||
if form_data.IMAGE_SIZE == 'auto' or form_data.IMAGE_SIZE == '' or re.match(pattern, form_data.IMAGE_SIZE):
|
||||
request.app.state.config.IMAGE_SIZE = form_data.IMAGE_SIZE
|
||||
else:
|
||||
if not (form_data.IMAGE_SIZE == 'auto' or form_data.IMAGE_SIZE == '' or re.match(pattern, form_data.IMAGE_SIZE)):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.INCORRECT_FORMAT(' (e.g., 512x512).'),
|
||||
)
|
||||
|
||||
if form_data.IMAGE_STEPS >= 0:
|
||||
request.app.state.config.IMAGE_STEPS = form_data.IMAGE_STEPS
|
||||
else:
|
||||
if form_data.IMAGE_STEPS < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.INCORRECT_FORMAT(' (e.g., 50).'),
|
||||
)
|
||||
|
||||
request.app.state.config.IMAGES_OPENAI_API_BASE_URL = form_data.IMAGES_OPENAI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_OPENAI_API_KEY = form_data.IMAGES_OPENAI_API_KEY
|
||||
request.app.state.config.IMAGES_OPENAI_API_VERSION = form_data.IMAGES_OPENAI_API_VERSION
|
||||
request.app.state.config.IMAGES_OPENAI_API_PARAMS = form_data.IMAGES_OPENAI_API_PARAMS
|
||||
|
||||
request.app.state.config.AUTOMATIC1111_BASE_URL = form_data.AUTOMATIC1111_BASE_URL
|
||||
request.app.state.config.AUTOMATIC1111_API_AUTH = form_data.AUTOMATIC1111_API_AUTH
|
||||
request.app.state.config.AUTOMATIC1111_PARAMS = form_data.AUTOMATIC1111_PARAMS
|
||||
|
||||
request.app.state.config.COMFYUI_BASE_URL = form_data.COMFYUI_BASE_URL.strip('/')
|
||||
request.app.state.config.COMFYUI_API_KEY = form_data.COMFYUI_API_KEY
|
||||
request.app.state.config.COMFYUI_WORKFLOW = form_data.COMFYUI_WORKFLOW
|
||||
request.app.state.config.COMFYUI_WORKFLOW_NODES = form_data.COMFYUI_WORKFLOW_NODES
|
||||
|
||||
request.app.state.config.IMAGES_GEMINI_API_BASE_URL = form_data.IMAGES_GEMINI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_GEMINI_API_KEY = form_data.IMAGES_GEMINI_API_KEY
|
||||
request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD = form_data.IMAGES_GEMINI_ENDPOINT_METHOD
|
||||
|
||||
# Edit Image
|
||||
request.app.state.config.ENABLE_IMAGE_EDIT = form_data.ENABLE_IMAGE_EDIT
|
||||
request.app.state.config.IMAGE_EDIT_ENGINE = form_data.IMAGE_EDIT_ENGINE
|
||||
request.app.state.config.IMAGE_EDIT_MODEL = form_data.IMAGE_EDIT_MODEL
|
||||
request.app.state.config.IMAGE_EDIT_SIZE = form_data.IMAGE_EDIT_SIZE
|
||||
|
||||
request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL = form_data.IMAGES_EDIT_OPENAI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY = form_data.IMAGES_EDIT_OPENAI_API_KEY
|
||||
request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION = form_data.IMAGES_EDIT_OPENAI_API_VERSION
|
||||
|
||||
request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL = form_data.IMAGES_EDIT_GEMINI_API_BASE_URL
|
||||
request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY = form_data.IMAGES_EDIT_GEMINI_API_KEY
|
||||
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL = form_data.IMAGES_EDIT_COMFYUI_BASE_URL.strip('/')
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY = form_data.IMAGES_EDIT_COMFYUI_API_KEY
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW = form_data.IMAGES_EDIT_COMFYUI_WORKFLOW
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES = form_data.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES
|
||||
|
||||
return {
|
||||
'ENABLE_IMAGE_GENERATION': request.app.state.config.ENABLE_IMAGE_GENERATION,
|
||||
'ENABLE_IMAGE_PROMPT_GENERATION': request.app.state.config.ENABLE_IMAGE_PROMPT_GENERATION,
|
||||
'IMAGE_GENERATION_ENGINE': request.app.state.config.IMAGE_GENERATION_ENGINE,
|
||||
'IMAGE_GENERATION_MODEL': request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
'IMAGE_SIZE': request.app.state.config.IMAGE_SIZE,
|
||||
'IMAGE_STEPS': request.app.state.config.IMAGE_STEPS,
|
||||
'IMAGES_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_OPENAI_API_BASE_URL,
|
||||
'IMAGES_OPENAI_API_KEY': request.app.state.config.IMAGES_OPENAI_API_KEY,
|
||||
'IMAGES_OPENAI_API_VERSION': request.app.state.config.IMAGES_OPENAI_API_VERSION,
|
||||
'IMAGES_OPENAI_API_PARAMS': request.app.state.config.IMAGES_OPENAI_API_PARAMS,
|
||||
'AUTOMATIC1111_BASE_URL': request.app.state.config.AUTOMATIC1111_BASE_URL,
|
||||
'AUTOMATIC1111_API_AUTH': request.app.state.config.AUTOMATIC1111_API_AUTH,
|
||||
'AUTOMATIC1111_PARAMS': request.app.state.config.AUTOMATIC1111_PARAMS,
|
||||
'COMFYUI_BASE_URL': request.app.state.config.COMFYUI_BASE_URL,
|
||||
'COMFYUI_API_KEY': request.app.state.config.COMFYUI_API_KEY,
|
||||
'COMFYUI_WORKFLOW': request.app.state.config.COMFYUI_WORKFLOW,
|
||||
'COMFYUI_WORKFLOW_NODES': request.app.state.config.COMFYUI_WORKFLOW_NODES,
|
||||
'IMAGES_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_GEMINI_API_BASE_URL,
|
||||
'IMAGES_GEMINI_API_KEY': request.app.state.config.IMAGES_GEMINI_API_KEY,
|
||||
'IMAGES_GEMINI_ENDPOINT_METHOD': request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD,
|
||||
'ENABLE_IMAGE_EDIT': request.app.state.config.ENABLE_IMAGE_EDIT,
|
||||
'IMAGE_EDIT_ENGINE': request.app.state.config.IMAGE_EDIT_ENGINE,
|
||||
'IMAGE_EDIT_MODEL': request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
'IMAGE_EDIT_SIZE': request.app.state.config.IMAGE_EDIT_SIZE,
|
||||
'IMAGES_EDIT_OPENAI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL,
|
||||
'IMAGES_EDIT_OPENAI_API_KEY': request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY,
|
||||
'IMAGES_EDIT_OPENAI_API_VERSION': request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION,
|
||||
'IMAGES_EDIT_GEMINI_API_BASE_URL': request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL,
|
||||
'IMAGES_EDIT_GEMINI_API_KEY': request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_BASE_URL': request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
'IMAGES_EDIT_COMFYUI_API_KEY': request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'IMAGES_EDIT_COMFYUI_WORKFLOW_NODES': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
updates = config_updates(form_data.model_dump(), IMAGE_CONFIG_KEYS)
|
||||
updates['image_generation.comfyui.base_url'] = form_data.COMFYUI_BASE_URL.strip('/')
|
||||
updates['images.edit.comfyui.base_url'] = form_data.IMAGES_EDIT_COMFYUI_BASE_URL.strip('/')
|
||||
await Config.upsert(updates)
|
||||
await set_image_model(request, form_data.IMAGE_GENERATION_MODEL)
|
||||
values = await get_config_values(IMAGE_CONFIG_KEYS)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.CONFIG_UPDATED,
|
||||
actor=user,
|
||||
subject_id='images',
|
||||
data={
|
||||
'image_generation_enabled': values.get('ENABLE_IMAGE_GENERATION'),
|
||||
'image_edit_enabled': values.get('ENABLE_IMAGE_EDIT'),
|
||||
'image_generation_engine': values.get('IMAGE_GENERATION_ENGINE'),
|
||||
'image_edit_engine': values.get('IMAGE_EDIT_ENGINE'),
|
||||
},
|
||||
)
|
||||
return values
|
||||
|
||||
|
||||
def get_automatic1111_api_auth(request: Request):
|
||||
if request.app.state.config.AUTOMATIC1111_API_AUTH is None:
|
||||
def get_automatic1111_api_auth(image_config):
|
||||
if image_config.AUTOMATIC1111_API_AUTH is None:
|
||||
return ''
|
||||
else:
|
||||
auth1111_byte_string = request.app.state.config.AUTOMATIC1111_API_AUTH.encode('utf-8')
|
||||
auth1111_byte_string = image_config.AUTOMATIC1111_API_AUTH.encode('utf-8')
|
||||
auth1111_base64_encoded_bytes = base64.b64encode(auth1111_byte_string)
|
||||
auth1111_base64_encoded_string = auth1111_base64_encoded_bytes.decode('utf-8')
|
||||
return f'Basic {auth1111_base64_encoded_string}'
|
||||
@@ -317,26 +330,27 @@ def get_automatic1111_api_auth(request: Request):
|
||||
|
||||
@router.get('/config/url/verify')
|
||||
async def verify_url(request: Request, user=Depends(get_admin_user)):
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111':
|
||||
image_config = await get_image_config()
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'automatic1111':
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/options',
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
return True
|
||||
except Exception:
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.INVALID_URL)
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
headers = None
|
||||
if request.app.state.config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
|
||||
if image_config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {image_config.COMFYUI_API_KEY}'}
|
||||
try:
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.COMFYUI_BASE_URL}/object_info',
|
||||
url=f'{image_config.COMFYUI_BASE_URL}/object_info',
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
@@ -350,33 +364,34 @@ async def verify_url(request: Request, user=Depends(get_admin_user)):
|
||||
|
||||
@router.get('/models')
|
||||
async def get_models(request: Request, user=Depends(get_verified_user)):
|
||||
image_config = await get_image_config()
|
||||
try:
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
return [
|
||||
{'id': 'dall-e-2', 'name': 'DALL·E 2'},
|
||||
{'id': 'dall-e-3', 'name': 'DALL·E 3'},
|
||||
{'id': 'gpt-image-1', 'name': 'GPT-IMAGE 1'},
|
||||
{'id': 'gpt-image-1.5', 'name': 'GPT-IMAGE 1.5'},
|
||||
]
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
return [
|
||||
{'id': 'imagen-3.0-generate-002', 'name': 'imagen-3.0 generate-002'},
|
||||
]
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
# TODO - get models from comfyui
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
|
||||
headers = {'Authorization': f'Bearer {image_config.COMFYUI_API_KEY}'}
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.COMFYUI_BASE_URL}/object_info',
|
||||
url=f'{image_config.COMFYUI_BASE_URL}/object_info',
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
info = await r.json()
|
||||
|
||||
workflow = json.loads(request.app.state.config.COMFYUI_WORKFLOW)
|
||||
workflow = json.loads(image_config.COMFYUI_WORKFLOW)
|
||||
model_node_id = None
|
||||
|
||||
for node in request.app.state.config.COMFYUI_WORKFLOW_NODES:
|
||||
for node in image_config.COMFYUI_WORKFLOW_NODES:
|
||||
if node['type'] == 'model':
|
||||
if node['node_ids']:
|
||||
model_node_id = node['node_ids'][0]
|
||||
@@ -405,14 +420,11 @@ async def get_models(request: Request, user=Depends(get_verified_user)):
|
||||
info['CheckpointLoaderSimple']['input']['required']['ckpt_name'][0],
|
||||
)
|
||||
)
|
||||
elif (
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
|
||||
):
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'automatic1111' or image_config.IMAGE_GENERATION_ENGINE == '':
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/sd-models',
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/sd-models',
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
models = await r.json()
|
||||
@@ -423,7 +435,11 @@ async def get_models(request: Request, user=Depends(get_verified_user)):
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e))
|
||||
log.exception(f'Failed to list image generation models: {e}')
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Failed to retrieve image generation models'),
|
||||
)
|
||||
|
||||
|
||||
class CreateImageForm(BaseModel):
|
||||
@@ -468,12 +484,12 @@ async def get_image_data(data: str, headers=None, trusted_base_url: str | None =
|
||||
# ComfyUI on a private network), skip SSRF validation only when
|
||||
# the URL shares the exact same origin (scheme + host + port)
|
||||
# as the admin-configured base. This avoids both the global
|
||||
# ENABLE_RAG_LOCAL_WEB_FETCH hammer and a blanket trust flag
|
||||
# ENABLE_LOCAL_WEB_FETCH hammer and a blanket trust flag
|
||||
# that would follow arbitrary redirects.
|
||||
if trusted_base_url and _is_same_origin(data, trusted_base_url):
|
||||
log.debug(f'Skipping URL validation for trusted backend: {data}')
|
||||
else:
|
||||
validate_url(data)
|
||||
await asyncio.to_thread(validate_url, data)
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
data,
|
||||
@@ -540,21 +556,36 @@ async def upload_image(request, image_data, content_type, metadata, user, db=Non
|
||||
|
||||
@router.post('/generations')
|
||||
async def generate_images(request: Request, form_data: CreateImageForm, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_IMAGE_GENERATION:
|
||||
image_config = await get_image_config()
|
||||
if not image_config.ENABLE_IMAGE_GENERATION:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.image_generation', request.app.state.config.USER_PERMISSIONS
|
||||
user.id, 'features.image_generation', image_config.USER_PERMISSIONS
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
return await image_generations(request, form_data, user=user)
|
||||
result = await image_generations(request, form_data, user=user)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.IMAGE_GENERATED,
|
||||
actor=user,
|
||||
subject_id=None,
|
||||
subject_type='image',
|
||||
data={
|
||||
'model': form_data.model,
|
||||
'size': form_data.size,
|
||||
'n': form_data.n,
|
||||
'prompt_preview': form_data.prompt[:300],
|
||||
},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
async def image_generations(
|
||||
@@ -563,13 +594,14 @@ async def image_generations(
|
||||
metadata: dict | None = None,
|
||||
user=None,
|
||||
):
|
||||
image_config = await get_image_config()
|
||||
# if IMAGE_SIZE = 'auto', default WidthxHeight to the 512x512 default
|
||||
# This is only relevant when the user has set IMAGE_SIZE to 'auto' with an
|
||||
# image model other than gpt-image-1, which is warned about on settings save
|
||||
|
||||
size = '512x512'
|
||||
if request.app.state.config.IMAGE_SIZE and 'x' in request.app.state.config.IMAGE_SIZE:
|
||||
size = request.app.state.config.IMAGE_SIZE
|
||||
if image_config.IMAGE_SIZE and 'x' in image_config.IMAGE_SIZE:
|
||||
size = image_config.IMAGE_SIZE
|
||||
|
||||
if form_data.size and 'x' in form_data.size:
|
||||
size = form_data.size
|
||||
@@ -581,41 +613,37 @@ async def image_generations(
|
||||
model = await get_image_model(request)
|
||||
|
||||
try:
|
||||
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
if image_config.IMAGE_GENERATION_ENGINE == 'openai':
|
||||
headers = {
|
||||
'Authorization': f'Bearer {request.app.state.config.IMAGES_OPENAI_API_KEY}',
|
||||
'Authorization': f'Bearer {image_config.IMAGES_OPENAI_API_KEY}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
url = f'{request.app.state.config.IMAGES_OPENAI_API_BASE_URL}/images/generations'
|
||||
if request.app.state.config.IMAGES_OPENAI_API_VERSION:
|
||||
url = f'{url}?api-version={request.app.state.config.IMAGES_OPENAI_API_VERSION}'
|
||||
url = f'{image_config.IMAGES_OPENAI_API_BASE_URL}/images/generations'
|
||||
if image_config.IMAGES_OPENAI_API_VERSION:
|
||||
url = f'{url}?api-version={image_config.IMAGES_OPENAI_API_VERSION}'
|
||||
|
||||
data = {
|
||||
'model': model,
|
||||
'prompt': form_data.prompt,
|
||||
'n': form_data.n,
|
||||
**(
|
||||
{'size': form_data.size or request.app.state.config.IMAGE_SIZE}
|
||||
if (form_data.size or request.app.state.config.IMAGE_SIZE)
|
||||
{'size': form_data.size or image_config.IMAGE_SIZE}
|
||||
if (form_data.size or image_config.IMAGE_SIZE)
|
||||
else {}
|
||||
),
|
||||
**(
|
||||
{}
|
||||
if re.match(
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
image_config.IMAGE_GENERATION_MODEL,
|
||||
)
|
||||
else {'response_format': 'b64_json'}
|
||||
),
|
||||
**(
|
||||
{}
|
||||
if not request.app.state.config.IMAGES_OPENAI_API_PARAMS
|
||||
else request.app.state.config.IMAGES_OPENAI_API_PARAMS
|
||||
),
|
||||
**({} if not image_config.IMAGES_OPENAI_API_PARAMS else image_config.IMAGES_OPENAI_API_PARAMS),
|
||||
}
|
||||
|
||||
session = await get_session()
|
||||
@@ -643,17 +671,17 @@ async def image_generations(
|
||||
images.append({'url': url})
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'gemini':
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'x-goog-api-key': request.app.state.config.IMAGES_GEMINI_API_KEY,
|
||||
'x-goog-api-key': image_config.IMAGES_GEMINI_API_KEY,
|
||||
}
|
||||
|
||||
data = {}
|
||||
|
||||
if (
|
||||
request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD == ''
|
||||
or request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD == 'predict'
|
||||
image_config.IMAGES_GEMINI_ENDPOINT_METHOD == ''
|
||||
or image_config.IMAGES_GEMINI_ENDPOINT_METHOD == 'predict'
|
||||
):
|
||||
model = f'{model}:predict'
|
||||
data = {
|
||||
@@ -664,13 +692,13 @@ async def image_generations(
|
||||
},
|
||||
}
|
||||
|
||||
elif request.app.state.config.IMAGES_GEMINI_ENDPOINT_METHOD == 'generateContent':
|
||||
elif image_config.IMAGES_GEMINI_ENDPOINT_METHOD == 'generateContent':
|
||||
model = f'{model}:generateContent'
|
||||
data = {'contents': [{'parts': [{'text': form_data.prompt}]}]}
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.IMAGES_GEMINI_API_BASE_URL}/models/{model}',
|
||||
url=f'{image_config.IMAGES_GEMINI_API_BASE_URL}/models/{model}',
|
||||
json=data,
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -701,7 +729,7 @@ async def image_generations(
|
||||
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'comfyui':
|
||||
data = {
|
||||
'prompt': form_data.prompt,
|
||||
'width': width,
|
||||
@@ -709,8 +737,8 @@ async def image_generations(
|
||||
'n': form_data.n,
|
||||
}
|
||||
|
||||
if request.app.state.config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else request.app.state.config.IMAGE_STEPS
|
||||
if image_config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else image_config.IMAGE_STEPS
|
||||
|
||||
if form_data.negative_prompt is not None:
|
||||
data['negative_prompt'] = form_data.negative_prompt
|
||||
@@ -719,8 +747,8 @@ async def image_generations(
|
||||
**{
|
||||
'workflow': ComfyUIWorkflow(
|
||||
**{
|
||||
'workflow': request.app.state.config.COMFYUI_WORKFLOW,
|
||||
'nodes': request.app.state.config.COMFYUI_WORKFLOW_NODES,
|
||||
'workflow': image_config.COMFYUI_WORKFLOW,
|
||||
'nodes': image_config.COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
),
|
||||
**data,
|
||||
@@ -730,8 +758,8 @@ async def image_generations(
|
||||
model,
|
||||
form_data,
|
||||
str(uuid.uuid4()),
|
||||
request.app.state.config.COMFYUI_BASE_URL,
|
||||
request.app.state.config.COMFYUI_API_KEY,
|
||||
image_config.COMFYUI_BASE_URL,
|
||||
image_config.COMFYUI_API_KEY,
|
||||
)
|
||||
log.debug(f'res: {res}')
|
||||
|
||||
@@ -739,13 +767,13 @@ async def image_generations(
|
||||
|
||||
for image in res['data']:
|
||||
headers = None
|
||||
if request.app.state.config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.COMFYUI_API_KEY}'}
|
||||
if image_config.COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {image_config.COMFYUI_API_KEY}'}
|
||||
|
||||
image_data, content_type = await get_image_data(
|
||||
image['url'],
|
||||
headers,
|
||||
trusted_base_url=request.app.state.config.COMFYUI_BASE_URL,
|
||||
trusted_base_url=image_config.COMFYUI_BASE_URL,
|
||||
)
|
||||
_, url = await upload_image(
|
||||
request,
|
||||
@@ -756,10 +784,7 @@ async def image_generations(
|
||||
)
|
||||
images.append({'url': url})
|
||||
return images
|
||||
elif (
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE == 'automatic1111'
|
||||
or request.app.state.config.IMAGE_GENERATION_ENGINE == ''
|
||||
):
|
||||
elif image_config.IMAGE_GENERATION_ENGINE == 'automatic1111' or image_config.IMAGE_GENERATION_ENGINE == '':
|
||||
if form_data.model:
|
||||
await set_image_model(request, form_data.model)
|
||||
|
||||
@@ -770,20 +795,20 @@ async def image_generations(
|
||||
'height': height,
|
||||
}
|
||||
|
||||
if request.app.state.config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else request.app.state.config.IMAGE_STEPS
|
||||
if image_config.IMAGE_STEPS is not None or form_data.steps is not None:
|
||||
data['steps'] = form_data.steps if form_data.steps is not None else image_config.IMAGE_STEPS
|
||||
|
||||
if form_data.negative_prompt is not None:
|
||||
data['negative_prompt'] = form_data.negative_prompt
|
||||
|
||||
if request.app.state.config.AUTOMATIC1111_PARAMS:
|
||||
data = {**data, **request.app.state.config.AUTOMATIC1111_PARAMS}
|
||||
if image_config.AUTOMATIC1111_PARAMS:
|
||||
data = {**data, **image_config.AUTOMATIC1111_PARAMS}
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.AUTOMATIC1111_BASE_URL}/sdapi/v1/txt2img',
|
||||
url=f'{image_config.AUTOMATIC1111_BASE_URL}/sdapi/v1/txt2img',
|
||||
json=data,
|
||||
headers={'authorization': get_automatic1111_api_auth(request)},
|
||||
headers={'authorization': get_automatic1111_api_auth(image_config)},
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as r:
|
||||
res = await r.json(content_type=None)
|
||||
@@ -820,23 +845,61 @@ class EditImageForm(BaseModel):
|
||||
|
||||
|
||||
@router.post('/edit')
|
||||
async def edit_images(request: Request, form_data: EditImageForm, user=Depends(get_verified_user)):
|
||||
# Authorize the direct route like /generations and the edit_image tool: enforce the
|
||||
# global image-edit switch and the per-user image-generation permission. The internal
|
||||
# callers (edit_image tool, chat middleware) gate themselves and call image_edits()
|
||||
# directly, so they are unaffected by this wrapper.
|
||||
image_config = await get_image_config()
|
||||
if not image_config.ENABLE_IMAGE_EDIT:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.image_generation', image_config.USER_PERMISSIONS
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
result = await image_edits(request, form_data, user=user)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.IMAGE_EDITED,
|
||||
actor=user,
|
||||
subject_id=None,
|
||||
subject_type='image',
|
||||
data={
|
||||
'model': form_data.model,
|
||||
'size': form_data.size,
|
||||
'n': form_data.n,
|
||||
'prompt_preview': form_data.prompt[:300],
|
||||
},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
async def image_edits(
|
||||
request: Request,
|
||||
form_data: EditImageForm,
|
||||
metadata: dict | None = None,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
image_config = await get_image_config()
|
||||
size = None
|
||||
width, height = None, None
|
||||
metadata = metadata or {}
|
||||
|
||||
if (request.app.state.config.IMAGE_EDIT_SIZE and 'x' in request.app.state.config.IMAGE_EDIT_SIZE) or (
|
||||
if (image_config.IMAGE_EDIT_SIZE and 'x' in image_config.IMAGE_EDIT_SIZE) or (
|
||||
form_data.size and 'x' in form_data.size
|
||||
):
|
||||
size = form_data.size if form_data.size else request.app.state.config.IMAGE_EDIT_SIZE
|
||||
size = form_data.size if form_data.size else image_config.IMAGE_EDIT_SIZE
|
||||
width, height = tuple(map(int, size.split('x')))
|
||||
|
||||
model = request.app.state.config.IMAGE_EDIT_MODEL if form_data.model is None else form_data.model
|
||||
model = image_config.IMAGE_EDIT_MODEL if form_data.model is None else form_data.model
|
||||
|
||||
try:
|
||||
|
||||
@@ -850,15 +913,17 @@ async def image_edits(
|
||||
# called only on the originally-submitted URL; following 3xx redirects
|
||||
# without re-validation would let an attacker reach private IPs via a
|
||||
# public host that redirects internally (e.g. cloud-metadata exfil).
|
||||
validate_url(data)
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
data, ssl=AIOHTTP_CLIENT_SESSION_SSL, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
await asyncio.to_thread(validate_url, data)
|
||||
# SSRF-safe session: re-checks the connect-time IP so a rebinding DNS answer
|
||||
# that passed validate_url cannot reach an internal address.
|
||||
async with get_ssrf_safe_session() as session:
|
||||
async with session.get(
|
||||
data, ssl=AIOHTTP_CLIENT_SESSION_SSL, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS
|
||||
) as r:
|
||||
r.raise_for_status()
|
||||
|
||||
image_data = base64.b64encode(await r.read()).decode('utf-8')
|
||||
return f'data:{r.headers["content-type"]};base64,{image_data}'
|
||||
image_data = base64.b64encode(await r.read()).decode('utf-8')
|
||||
return f'data:{r.headers["content-type"]};base64,{image_data}'
|
||||
|
||||
else:
|
||||
file_id = None
|
||||
@@ -885,27 +950,18 @@ async def image_edits(
|
||||
elif isinstance(form_data.image, list):
|
||||
# Load all images in parallel for better performance
|
||||
form_data.image = list(await asyncio.gather(*[load_url_image(img) for img in form_data.image]))
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e))
|
||||
|
||||
def get_image_file_item(base64_string, param_name='image'):
|
||||
data = base64_string
|
||||
header, encoded = data.split(',', 1)
|
||||
mime_type = header.split(';')[0].lstrip('data:')
|
||||
image_data = base64.b64decode(encoded)
|
||||
return (
|
||||
param_name,
|
||||
(
|
||||
f'{uuid.uuid4()}.png',
|
||||
io.BytesIO(image_data),
|
||||
mime_type if mime_type else 'image/png',
|
||||
),
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error loading image'),
|
||||
)
|
||||
|
||||
try:
|
||||
if request.app.state.config.IMAGE_EDIT_ENGINE == 'openai':
|
||||
if image_config.IMAGE_EDIT_ENGINE == 'openai':
|
||||
headers = {
|
||||
'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_OPENAI_API_KEY}',
|
||||
'Authorization': f'Bearer {image_config.IMAGES_EDIT_OPENAI_API_KEY}',
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS:
|
||||
@@ -921,7 +977,7 @@ async def image_edits(
|
||||
{}
|
||||
if re.match(
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
image_config.IMAGE_EDIT_MODEL,
|
||||
)
|
||||
else {'response_format': 'b64_json'}
|
||||
),
|
||||
@@ -929,14 +985,19 @@ async def image_edits(
|
||||
|
||||
files = []
|
||||
if isinstance(form_data.image, str):
|
||||
files = [get_image_file_item(form_data.image)]
|
||||
image = form_data.image
|
||||
if ENABLE_OPENAI_IMAGE_EDIT_NORMALIZATION:
|
||||
image = normalize_openai_edit_image_data_url(image)
|
||||
files = [get_image_file_item(image)]
|
||||
elif isinstance(form_data.image, list):
|
||||
for img in form_data.image:
|
||||
if ENABLE_OPENAI_IMAGE_EDIT_NORMALIZATION:
|
||||
img = normalize_openai_edit_image_data_url(img)
|
||||
files.append(get_image_file_item(img, 'image[]'))
|
||||
|
||||
url_search_params = ''
|
||||
if request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION:
|
||||
url_search_params += f'?api-version={request.app.state.config.IMAGES_EDIT_OPENAI_API_VERSION}'
|
||||
if image_config.IMAGES_EDIT_OPENAI_API_VERSION:
|
||||
url_search_params += f'?api-version={image_config.IMAGES_EDIT_OPENAI_API_VERSION}'
|
||||
|
||||
# Build multipart form data for aiohttp
|
||||
form = aiohttp.FormData()
|
||||
@@ -955,7 +1016,7 @@ async def image_edits(
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.IMAGES_EDIT_OPENAI_API_BASE_URL}/images/edits{url_search_params}',
|
||||
url=f'{image_config.IMAGES_EDIT_OPENAI_API_BASE_URL}/images/edits{url_search_params}',
|
||||
headers=headers,
|
||||
data=form,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -977,10 +1038,10 @@ async def image_edits(
|
||||
images.append({'url': url})
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_EDIT_ENGINE == 'gemini':
|
||||
elif image_config.IMAGE_EDIT_ENGINE == 'gemini':
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'x-goog-api-key': request.app.state.config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
'x-goog-api-key': image_config.IMAGES_EDIT_GEMINI_API_KEY,
|
||||
}
|
||||
|
||||
model = f'{model}:generateContent'
|
||||
@@ -1010,7 +1071,7 @@ async def image_edits(
|
||||
|
||||
session = await get_session()
|
||||
async with session.post(
|
||||
url=f'{request.app.state.config.IMAGES_EDIT_GEMINI_API_BASE_URL}/models/{model}',
|
||||
url=f'{image_config.IMAGES_EDIT_GEMINI_API_BASE_URL}/models/{model}',
|
||||
json=data,
|
||||
headers=headers,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
@@ -1034,7 +1095,7 @@ async def image_edits(
|
||||
|
||||
return images
|
||||
|
||||
elif request.app.state.config.IMAGE_EDIT_ENGINE == 'comfyui':
|
||||
elif image_config.IMAGE_EDIT_ENGINE == 'comfyui':
|
||||
try:
|
||||
files = []
|
||||
if isinstance(form_data.image, str):
|
||||
@@ -1048,8 +1109,8 @@ async def image_edits(
|
||||
for file_item in files:
|
||||
res = await comfyui_upload_image(
|
||||
file_item,
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
image_config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
image_config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
)
|
||||
comfyui_images.append(res.get('name', file_item[1][0]))
|
||||
except Exception as e:
|
||||
@@ -1068,8 +1129,8 @@ async def image_edits(
|
||||
**{
|
||||
'workflow': ComfyUIWorkflow(
|
||||
**{
|
||||
'workflow': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'nodes': request.app.state.config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
'workflow': image_config.IMAGES_EDIT_COMFYUI_WORKFLOW,
|
||||
'nodes': image_config.IMAGES_EDIT_COMFYUI_WORKFLOW_NODES,
|
||||
}
|
||||
),
|
||||
**data,
|
||||
@@ -1079,8 +1140,8 @@ async def image_edits(
|
||||
model,
|
||||
form_data,
|
||||
str(uuid.uuid4()),
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
image_config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
image_config.IMAGES_EDIT_COMFYUI_API_KEY,
|
||||
)
|
||||
log.debug(f'res: {res}')
|
||||
|
||||
@@ -1099,13 +1160,13 @@ async def image_edits(
|
||||
|
||||
for image_url in image_urls:
|
||||
headers = None
|
||||
if request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {request.app.state.config.IMAGES_EDIT_COMFYUI_API_KEY}'}
|
||||
if image_config.IMAGES_EDIT_COMFYUI_API_KEY:
|
||||
headers = {'Authorization': f'Bearer {image_config.IMAGES_EDIT_COMFYUI_API_KEY}'}
|
||||
|
||||
image_data, content_type = await get_image_data(
|
||||
image_url,
|
||||
headers,
|
||||
trusted_base_url=request.app.state.config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
trusted_base_url=image_config.IMAGES_EDIT_COMFYUI_BASE_URL,
|
||||
)
|
||||
_, url = await upload_image(
|
||||
request,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2,16 +2,27 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Optional
|
||||
from typing import Literal
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.memories import Memories, MemoryModel
|
||||
from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
||||
from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_verified_user
|
||||
from open_webui.utils.memory import (
|
||||
clean_memory_content,
|
||||
clean_memory_path,
|
||||
list_memory_path_groups,
|
||||
memory_vector_text,
|
||||
read_memory_path_rows,
|
||||
search_memory_rows,
|
||||
validate_memory_operations,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
@@ -20,6 +31,21 @@ log = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
async def check_memories_permission(user):
|
||||
config = await Config.get_many('memories.enable', 'user.permissions')
|
||||
if not config.get('memories.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'features.memories', config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
# GetMemories
|
||||
# Let what is remembered here spare someone the cost
|
||||
@@ -33,17 +59,7 @@ async def get_memories(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
return await Memories.get_memories_by_user_id(user.id, db=db)
|
||||
|
||||
@@ -55,10 +71,57 @@ async def get_memories(
|
||||
|
||||
class AddMemoryForm(BaseModel):
|
||||
content: str
|
||||
type: Literal['user', 'context'] = 'context'
|
||||
path: str | None = None
|
||||
|
||||
|
||||
class MemoryUpdateModel(BaseModel):
|
||||
content: str | None = None
|
||||
type: Literal['user', 'context'] | None = None
|
||||
path: str | None = None
|
||||
|
||||
|
||||
class MemoryOperationModel(BaseModel):
|
||||
action: Literal['add', 'replace', 'remove', 'move']
|
||||
id: str | None = None
|
||||
content: str | None = None
|
||||
type: Literal['user', 'context'] | None = None
|
||||
path: str | None = None
|
||||
|
||||
|
||||
class UpdateMemoriesForm(BaseModel):
|
||||
operations: list[MemoryOperationModel]
|
||||
source: Literal['tool', 'background_review'] | None = None
|
||||
|
||||
|
||||
class SearchMemoriesForm(BaseModel):
|
||||
query: str | None = None
|
||||
type: Literal['user', 'context', 'all'] = 'all'
|
||||
path: str | None = None
|
||||
memory_id: str | None = None
|
||||
limit: int = 20
|
||||
|
||||
|
||||
class ListMemoryPathsForm(BaseModel):
|
||||
query: str | None = None
|
||||
type: Literal['user', 'context', 'all'] = 'all'
|
||||
limit: int = 100
|
||||
|
||||
|
||||
class ReadMemoryPathForm(BaseModel):
|
||||
path: str
|
||||
type: Literal['user', 'context', 'all'] = 'all'
|
||||
include_children: bool = True
|
||||
limit: int = 50
|
||||
|
||||
|
||||
def _memory_metadata(memory: MemoryModel) -> dict:
|
||||
return {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
'type': memory.type,
|
||||
'path': memory.path,
|
||||
}
|
||||
|
||||
|
||||
@router.post('/add', response_model=MemoryModel | None)
|
||||
@@ -73,37 +136,128 @@ async def add_memory(
|
||||
own short-lived sessions so a connection is not held during the external
|
||||
embedding API call (``EMBEDDING_FUNCTION``), which can take 1-5+ seconds.
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
content = clean_memory_content(form_data.content)
|
||||
path = clean_memory_path(form_data.path)
|
||||
memory = await Memories.insert_new_memory(
|
||||
user.id,
|
||||
content,
|
||||
memory_type=form_data.type,
|
||||
path=path,
|
||||
meta={'created_by': 'manual'},
|
||||
)
|
||||
|
||||
memory = await Memories.insert_new_memory(user.id, form_data.content)
|
||||
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user)
|
||||
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'text': memory_vector_text(memory.content, memory.path),
|
||||
'vector': vector,
|
||||
'metadata': {'created_at': memory.created_at},
|
||||
'metadata': _memory_metadata(memory),
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MEMORY_CREATED,
|
||||
actor=user,
|
||||
subject_id=memory.id,
|
||||
data={'content_preview': memory.content[:300], 'type': memory.type, 'path': memory.path},
|
||||
)
|
||||
return memory
|
||||
|
||||
|
||||
@router.post('/update', response_model=list[dict])
|
||||
async def update_memories(
|
||||
request: Request,
|
||||
form_data: UpdateMemoriesForm,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
await check_memories_permission(user)
|
||||
|
||||
operations = validate_memory_operations(form_data)
|
||||
metadata = getattr(request.state, 'metadata', {}) or {}
|
||||
source = form_data.source or 'tool'
|
||||
for operation in operations:
|
||||
if operation.get('action') in {'add', 'replace', 'move'}:
|
||||
operation['meta'] = {
|
||||
'created_by': source,
|
||||
'chat_id': metadata.get('chat_id'),
|
||||
'message_id': metadata.get('message_id'),
|
||||
'model': metadata.get('model'),
|
||||
}
|
||||
|
||||
try:
|
||||
results = await Memories.apply_memory_operations(user.id, operations)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
|
||||
upsert_items = []
|
||||
delete_ids = []
|
||||
response = []
|
||||
|
||||
for result in results:
|
||||
memory = result.get('memory')
|
||||
if isinstance(memory, MemoryModel):
|
||||
result = {**result, 'memory': memory.model_dump()}
|
||||
if result.get('status') in {'created', 'updated'}:
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(
|
||||
memory_vector_text(memory.content, memory.path),
|
||||
user=user,
|
||||
)
|
||||
upsert_items.append(
|
||||
{
|
||||
'id': memory.id,
|
||||
'text': memory_vector_text(memory.content, memory.path),
|
||||
'vector': vector,
|
||||
'metadata': _memory_metadata(memory),
|
||||
}
|
||||
)
|
||||
if result.get('status') == 'deleted' and result.get('id'):
|
||||
delete_ids.append(result['id'])
|
||||
response.append(result)
|
||||
|
||||
if upsert_items:
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(collection_name=f'user-memory-{user.id}', items=upsert_items)
|
||||
|
||||
if delete_ids:
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=delete_ids)
|
||||
|
||||
for result in response:
|
||||
status_value = result.get('status')
|
||||
memory = result.get('memory') or {}
|
||||
memory_id = memory.get('id') or result.get('id')
|
||||
|
||||
if status_value == 'created':
|
||||
event = EVENTS.MEMORY_CREATED
|
||||
elif status_value == 'updated':
|
||||
event = EVENTS.MEMORY_UPDATED
|
||||
elif status_value == 'deleted':
|
||||
event = EVENTS.MEMORY_DELETED
|
||||
else:
|
||||
continue
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
event,
|
||||
actor=user,
|
||||
subject_id=memory_id,
|
||||
data={
|
||||
'content_preview': (memory.get('content') or '')[:300],
|
||||
'type': memory.get('type'),
|
||||
'path': memory.get('path'),
|
||||
'operation': result.get('action'),
|
||||
},
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
############################
|
||||
# QueryMemory
|
||||
############################
|
||||
@@ -124,17 +278,7 @@ async def query_memory(
|
||||
# Database operations (get_memories_by_user_id) manage their own short-lived sessions.
|
||||
# This prevents holding a connection during EMBEDDING_FUNCTION()
|
||||
# which makes external embedding API calls (1-5+ seconds).
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
if not memories:
|
||||
@@ -154,7 +298,7 @@ async def query_memory(
|
||||
# same RELEVANCE_THRESHOLD used by RAG ensures only genuinely matching
|
||||
# memories are surfaced (distances are normalised to 0→1, higher is
|
||||
# better).
|
||||
relevance_threshold = getattr(request.app.state.config, 'RELEVANCE_THRESHOLD', 0.0)
|
||||
relevance_threshold = await Config.get('rag.relevance_threshold', 0.0)
|
||||
if results and relevance_threshold > 0.0 and results.distances and results.distances[0]:
|
||||
from open_webui.retrieval.vector.main import SearchResult
|
||||
|
||||
@@ -183,6 +327,61 @@ async def query_memory(
|
||||
return results
|
||||
|
||||
|
||||
@router.post('/search', response_model=list[MemoryModel])
|
||||
async def search_memories(
|
||||
form_data: SearchMemoriesForm,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
await check_memories_permission(user)
|
||||
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
return search_memory_rows(
|
||||
memories,
|
||||
query=form_data.query,
|
||||
path=form_data.path,
|
||||
memory_id=form_data.memory_id,
|
||||
memory_type=form_data.type,
|
||||
limit=form_data.limit,
|
||||
)
|
||||
|
||||
|
||||
@router.post('/paths')
|
||||
async def list_memory_paths(
|
||||
form_data: ListMemoryPathsForm,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
await check_memories_permission(user)
|
||||
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
return list_memory_path_groups(
|
||||
memories,
|
||||
query=form_data.query or '',
|
||||
memory_type=form_data.type,
|
||||
limit=form_data.limit,
|
||||
)
|
||||
|
||||
|
||||
@router.post('/path')
|
||||
async def read_memory_path(
|
||||
form_data: ReadMemoryPathForm,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
await check_memories_permission(user)
|
||||
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
result = read_memory_path_rows(
|
||||
memories,
|
||||
path=form_data.path,
|
||||
memory_type=form_data.type,
|
||||
include_children=form_data.include_children,
|
||||
limit=form_data.limit,
|
||||
)
|
||||
return {
|
||||
**result,
|
||||
'memories': [memory.model_dump() for memory in result['memories']],
|
||||
}
|
||||
|
||||
|
||||
############################
|
||||
# ResetMemoryFromVectorDB
|
||||
############################
|
||||
@@ -199,17 +398,7 @@ async def reset_memory_from_vector_db(
|
||||
calls simultaneously. With a session held, this could block a connection
|
||||
for MINUTES, completely exhausting the connection pool.
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
|
||||
@@ -217,7 +406,10 @@ async def reset_memory_from_vector_db(
|
||||
|
||||
# Generate vectors in parallel
|
||||
vectors = await asyncio.gather(
|
||||
*[request.app.state.EMBEDDING_FUNCTION(memory.content, user=user) for memory in memories]
|
||||
*[
|
||||
request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user)
|
||||
for memory in memories
|
||||
]
|
||||
)
|
||||
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
@@ -225,17 +417,22 @@ async def reset_memory_from_vector_db(
|
||||
items=[
|
||||
{
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'text': memory_vector_text(memory.content, memory.path),
|
||||
'vector': vectors[idx],
|
||||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
},
|
||||
'metadata': _memory_metadata(memory),
|
||||
}
|
||||
for idx, memory in enumerate(memories)
|
||||
],
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MEMORY_RESET,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
subject_type='user',
|
||||
data={'count': len(memories)},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
@@ -250,17 +447,7 @@ async def delete_memory_by_user_id(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
result = await Memories.delete_memories_by_user_id(user.id, db=db)
|
||||
|
||||
@@ -269,6 +456,13 @@ async def delete_memory_by_user_id(
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
except Exception as e:
|
||||
log.error(e)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MEMORY_DELETED,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
subject_type='user',
|
||||
)
|
||||
return True
|
||||
|
||||
return False
|
||||
@@ -290,40 +484,46 @@ async def update_memory_by_id(
|
||||
# Database operations (update_memory_by_id_and_user_id) manage their own
|
||||
# short-lived sessions. This prevents holding a connection during
|
||||
# EMBEDDING_FUNCTION() which makes external API calls (1-5+ seconds).
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content)
|
||||
content = clean_memory_content(form_data.content) if form_data.content is not None else None
|
||||
path = clean_memory_path(form_data.path)
|
||||
if content is None and form_data.type is None and form_data.path is None:
|
||||
raise HTTPException(status_code=400, detail='No memory update provided')
|
||||
memory = await Memories.update_memory_by_id_and_user_id(
|
||||
memory_id,
|
||||
user.id,
|
||||
content,
|
||||
memory_type=form_data.type,
|
||||
path=path,
|
||||
update_path=form_data.path is not None,
|
||||
meta={'created_by': 'manual'},
|
||||
)
|
||||
if memory is None:
|
||||
raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
if form_data.content is not None:
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
if form_data.content is not None or form_data.path is not None:
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user)
|
||||
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'text': memory_vector_text(memory.content, memory.path),
|
||||
'vector': vector,
|
||||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
},
|
||||
'metadata': _memory_metadata(memory),
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MEMORY_UPDATED,
|
||||
actor=user,
|
||||
subject_id=memory.id,
|
||||
data={'content_preview': memory.content[:300], 'type': memory.type, 'path': memory.path},
|
||||
)
|
||||
return memory
|
||||
|
||||
|
||||
@@ -339,22 +539,18 @@ async def delete_memory_by_id(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_MEMORIES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not await has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
await check_memories_permission(user)
|
||||
|
||||
result = await Memories.delete_memory_by_id_and_user_id(memory_id, user.id, db=db)
|
||||
|
||||
if result:
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MEMORY_DELETED,
|
||||
actor=user,
|
||||
subject_id=memory_id,
|
||||
)
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@@ -20,9 +20,11 @@ from fastapi import (
|
||||
from fastapi.responses import RedirectResponse, StreamingResponse
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING, PROFILE_IMAGE_ALLOWED_MIME_TYPES
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import (
|
||||
ModelAccessListResponse,
|
||||
@@ -60,6 +62,9 @@ def _safe_static_redirect_path(url: str) -> str | None:
|
||||
if decoded == path:
|
||||
break
|
||||
path = decoded
|
||||
# Fail closed: a value still encoded after the cap would be decoded further downstream.
|
||||
if unquote(path) != path:
|
||||
return None
|
||||
if '\x00' in path or '\\' in path:
|
||||
return None
|
||||
if not path.startswith('/'):
|
||||
@@ -193,9 +198,19 @@ async def get_models(
|
||||
###########################
|
||||
|
||||
|
||||
@router.get('/base/tags', response_model=list[str])
|
||||
async def get_base_model_tags(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
tags = await Models.get_all_tags(user_id=user.id, is_admin=True, is_base_model=True, db=db)
|
||||
return sorted(tags)
|
||||
|
||||
|
||||
@router.get('/base', response_model=list[ModelResponse])
|
||||
async def get_base_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
return await Models.get_base_models(db=db)
|
||||
async def get_base_models(
|
||||
tag: str | None = None,
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
return await Models.get_base_models(tag=tag, db=db)
|
||||
|
||||
|
||||
###########################
|
||||
@@ -227,7 +242,7 @@ async def create_new_model(
|
||||
):
|
||||
"""Create a new workspace model entry."""
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'workspace.models', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'workspace.models', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -255,7 +270,7 @@ async def create_new_model(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -264,6 +279,13 @@ async def create_new_model(
|
||||
|
||||
model = await Models.insert_new_model(form_data, user.id, db=db)
|
||||
if model:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_CREATED,
|
||||
actor=user,
|
||||
subject_id=model.id,
|
||||
data={'name': model.name},
|
||||
)
|
||||
return model
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -286,7 +308,7 @@ async def export_models(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.models_export',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -319,7 +341,7 @@ async def import_models(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.models_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -356,10 +378,12 @@ async def import_models(
|
||||
else:
|
||||
writable_model_ids = set(existing_model_ids)
|
||||
|
||||
imported_ids = []
|
||||
for model_data in data:
|
||||
model_id = model_data.get('id')
|
||||
|
||||
if model_id and is_valid_model_id(model_id):
|
||||
imported_ids.append(model_id)
|
||||
# Defense-in-depth: skip models referencing inaccessible files
|
||||
try:
|
||||
await _verify_knowledge_file_access(
|
||||
@@ -400,7 +424,7 @@ async def import_models(
|
||||
# metadata-only imports.
|
||||
if 'access_grants' in model_data:
|
||||
updated_model.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
updated_model.access_grants,
|
||||
@@ -413,13 +437,20 @@ async def import_models(
|
||||
model_data['params'] = model_data.get('params', {})
|
||||
new_model = ModelForm(**model_data)
|
||||
new_model.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
new_model.access_grants,
|
||||
'sharing.public_models',
|
||||
)
|
||||
await Models.insert_new_model(user_id=user.id, form_data=new_model, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_IMPORTED,
|
||||
actor=user,
|
||||
subject_type='model',
|
||||
data={'count': len(imported_ids), 'model_ids': imported_ids},
|
||||
)
|
||||
return True
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail='Invalid JSON format')
|
||||
@@ -444,7 +475,15 @@ async def sync_models(
|
||||
user=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
return await Models.sync_models(user.id, form_data.models, db=db)
|
||||
models = await Models.sync_models(user.id, form_data.models, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_SYNCED,
|
||||
actor=user,
|
||||
subject_type='model',
|
||||
data={'count': len(models), 'model_ids': [model.id for model in models]},
|
||||
)
|
||||
return models
|
||||
|
||||
|
||||
###########################
|
||||
@@ -528,11 +567,7 @@ async def get_model_profile_image(
|
||||
|
||||
# Fallback: check arena models stored in config (not in the DB)
|
||||
if not profile_image_url:
|
||||
arena_models = getattr(
|
||||
getattr(request.app.state, 'config', None),
|
||||
'EVALUATION_ARENA_MODELS',
|
||||
[],
|
||||
)
|
||||
arena_models = await Config.get('evaluation.arena.models', []) or []
|
||||
for arena_model in arena_models:
|
||||
if arena_model.get('id') == id:
|
||||
profile_image_url = arena_model.get('meta', {}).get('profile_image_url')
|
||||
@@ -596,7 +631,9 @@ async def get_model_profile_image(
|
||||
|
||||
|
||||
@router.post('/model/toggle', response_model=ModelResponse | None)
|
||||
async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def toggle_model_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
model = await Models.get_model_by_id(id, db=db)
|
||||
if model:
|
||||
if (
|
||||
@@ -613,6 +650,14 @@ async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Async
|
||||
model = await Models.toggle_model_by_id(id, db=db)
|
||||
|
||||
if model:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_ENABLED if model.is_active else EVENTS.MODEL_DISABLED,
|
||||
actor=user,
|
||||
subject_id=model.id,
|
||||
subject_type='model',
|
||||
data={'name': model.name},
|
||||
)
|
||||
return model
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -674,7 +719,7 @@ async def update_model_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -682,6 +727,14 @@ async def update_model_by_id(
|
||||
)
|
||||
|
||||
model = await Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db)
|
||||
if model:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_UPDATED,
|
||||
actor=user,
|
||||
subject_id=model.id,
|
||||
data={'name': model.name},
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@@ -746,7 +799,7 @@ async def update_model_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -757,7 +810,14 @@ async def update_model_access_by_id(
|
||||
|
||||
await Models.update_model_updated_at_by_id(form_data.id, db=db)
|
||||
|
||||
return await Models.get_model_by_id(form_data.id, db=db)
|
||||
model = await Models.get_model_by_id(form_data.id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_ACCESS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=form_data.id,
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
############################
|
||||
@@ -767,6 +827,7 @@ async def update_model_access_by_id(
|
||||
|
||||
@router.post('/model/delete', response_model=bool)
|
||||
async def delete_model_by_id(
|
||||
request: Request,
|
||||
form_data: ModelIdForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
@@ -795,10 +856,22 @@ async def delete_model_by_id(
|
||||
)
|
||||
|
||||
result = await Models.delete_model_by_id(form_data.id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_DELETED,
|
||||
actor=user,
|
||||
subject_id=form_data.id,
|
||||
data={'name': model.name},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@router.delete('/delete/all', response_model=bool)
|
||||
async def delete_all_models(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_all_models(
|
||||
request: Request, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
result = await Models.delete_all_models(db=db)
|
||||
if result:
|
||||
await publish_event(request, EVENTS.MODEL_DELETED, actor=user, subject_type='model')
|
||||
return result
|
||||
|
||||
@@ -9,8 +9,10 @@ from open_webui.config import (
|
||||
ENABLE_ADMIN_EXPORT,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.notes import (
|
||||
NoteForm,
|
||||
@@ -66,7 +68,7 @@ async def get_notes(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -114,7 +116,7 @@ async def get_pinned_notes(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -155,7 +157,7 @@ async def search_notes(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -208,7 +210,7 @@ async def create_new_note(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -216,7 +218,7 @@ async def create_new_note(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -226,6 +228,13 @@ async def create_new_note(
|
||||
|
||||
try:
|
||||
note = await Notes.insert_new_note(user.id, form_data, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.NOTE_CREATED,
|
||||
actor=user,
|
||||
subject_id=note.id,
|
||||
data={'title': note.title},
|
||||
)
|
||||
return note
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -249,7 +258,7 @@ async def get_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -308,7 +317,7 @@ async def update_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -332,7 +341,7 @@ async def update_note_by_id(
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -351,6 +360,13 @@ async def update_note_by_id(
|
||||
to=f'note:{note.id}',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.NOTE_UPDATED,
|
||||
actor=user,
|
||||
subject_id=note.id,
|
||||
data={'title': note.title},
|
||||
)
|
||||
return note
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
@@ -375,7 +391,7 @@ async def update_note_access_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -399,7 +415,7 @@ async def update_note_access_by_id(
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT())
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -411,6 +427,12 @@ async def update_note_access_by_id(
|
||||
note = await Notes.get_note_by_id(id, db=db)
|
||||
pinned_note_ids = await Notes.get_pinned_note_ids(user.id, db=db)
|
||||
note.is_pinned = note.id in pinned_note_ids
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.NOTE_ACCESS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=note.id,
|
||||
)
|
||||
return note
|
||||
|
||||
|
||||
@@ -427,7 +449,7 @@ async def pin_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -453,6 +475,13 @@ async def pin_note_by_id(
|
||||
note = await Notes.toggle_note_pinned_by_id(id, user.id, db=db)
|
||||
pinned_note_ids = await Notes.get_pinned_note_ids(user.id, db=db)
|
||||
note.is_pinned = note.id in pinned_note_ids
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.NOTE_PINNED if note.is_pinned else EVENTS.NOTE_UNPINNED,
|
||||
actor=user,
|
||||
subject_id=note.id,
|
||||
subject_type='note',
|
||||
)
|
||||
return note
|
||||
|
||||
|
||||
@@ -469,7 +498,7 @@ async def delete_note_by_id(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'features.notes', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
user.id, 'features.notes', await Config.get('user.permissions'), db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -494,6 +523,12 @@ async def delete_note_by_id(
|
||||
|
||||
try:
|
||||
note = await Notes.delete_note_by_id(id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.NOTE_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
|
||||
@@ -20,6 +20,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from open_webui.config import UPLOAD_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import (
|
||||
AIOHTTP_CLIENT_SESSION_SSL,
|
||||
AIOHTTP_CLIENT_TIMEOUT,
|
||||
@@ -31,12 +32,13 @@ from open_webui.env import (
|
||||
)
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import Models
|
||||
from open_webui.models.users import UserModel
|
||||
from open_webui.utils.access_control import check_model_access
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
from open_webui.utils.headers import get_custom_headers, include_user_info_headers
|
||||
from open_webui.utils.misc import calculate_sha256
|
||||
from open_webui.utils.payload import (
|
||||
apply_model_params_to_body_ollama,
|
||||
@@ -97,6 +99,8 @@ async def send_request(
|
||||
stream: bool = False,
|
||||
content_type: str | None = None,
|
||||
metadata: dict | None = None,
|
||||
api_config: dict | None = None,
|
||||
request: Request | None = None,
|
||||
):
|
||||
r = None
|
||||
streaming = False
|
||||
@@ -113,6 +117,10 @@ async def send_request(
|
||||
if metadata and metadata.get('chat_id'):
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
|
||||
|
||||
# Custom per-connection headers last so admin-set headers take precedence.
|
||||
if api_config and api_config.get('headers'):
|
||||
headers.update(get_custom_headers(api_config['headers'], user, metadata, request=request))
|
||||
|
||||
r = await session.request(
|
||||
method,
|
||||
url,
|
||||
@@ -181,6 +189,32 @@ def get_api_key(idx, url, configs):
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
OLLAMA_CONFIG_KEYS = {
|
||||
'ENABLE_OLLAMA_API': 'ollama.enable',
|
||||
'OLLAMA_BASE_URLS': 'ollama.base_urls',
|
||||
'OLLAMA_API_CONFIGS': 'ollama.api_configs',
|
||||
}
|
||||
|
||||
|
||||
async def get_ollama_config_values() -> dict:
|
||||
values = await Config.get_many(*OLLAMA_CONFIG_KEYS.values())
|
||||
return {field: values[storage_key] for field, storage_key in OLLAMA_CONFIG_KEYS.items() if storage_key in values}
|
||||
|
||||
|
||||
async def get_ollama_runtime_config() -> tuple[bool, list[str], dict]:
|
||||
values = await Config.get_many('ollama.enable', 'ollama.base_urls', 'ollama.api_configs')
|
||||
return (
|
||||
values.get('ollama.enable'),
|
||||
values.get('ollama.base_urls') or [],
|
||||
values.get('ollama.api_configs') or {},
|
||||
)
|
||||
|
||||
|
||||
async def get_ollama_connection(idx: int) -> tuple[str, dict, str | None]:
|
||||
_, base_urls, api_configs = await get_ollama_runtime_config()
|
||||
url = base_urls[idx]
|
||||
return url, resolve_api_config(api_configs, idx, url), get_api_key(idx, url, api_configs)
|
||||
|
||||
|
||||
@router.head('/')
|
||||
@router.get('/')
|
||||
@@ -236,11 +270,7 @@ async def get_config(
|
||||
user=Depends(get_admin_user),
|
||||
) -> dict:
|
||||
"""Return the current Ollama connection configuration."""
|
||||
return {
|
||||
'ENABLE_OLLAMA_API': request.app.state.config.ENABLE_OLLAMA_API,
|
||||
'OLLAMA_BASE_URLS': request.app.state.config.OLLAMA_BASE_URLS,
|
||||
'OLLAMA_API_CONFIGS': request.app.state.config.OLLAMA_API_CONFIGS,
|
||||
}
|
||||
return await get_ollama_config_values()
|
||||
|
||||
|
||||
class OllamaConfigForm(BaseModel):
|
||||
@@ -258,20 +288,32 @@ async def update_config(
|
||||
user=Depends(get_admin_user),
|
||||
) -> dict:
|
||||
"""Persist updated Ollama connection settings."""
|
||||
request.app.state.config.ENABLE_OLLAMA_API = form_data.ENABLE_OLLAMA_API
|
||||
request.app.state.config.OLLAMA_BASE_URLS = form_data.OLLAMA_BASE_URLS
|
||||
request.app.state.config.OLLAMA_API_CONFIGS = form_data.OLLAMA_API_CONFIGS
|
||||
|
||||
# Prune stale config entries that no longer map to a URL index
|
||||
valid_keys = {str(i) for i in range(len(request.app.state.config.OLLAMA_BASE_URLS))}
|
||||
request.app.state.config.OLLAMA_API_CONFIGS = {
|
||||
k: v for k, v in request.app.state.config.OLLAMA_API_CONFIGS.items() if k in valid_keys
|
||||
}
|
||||
valid_keys = {str(i) for i in range(len(form_data.OLLAMA_BASE_URLS))}
|
||||
api_configs = {k: v for k, v in form_data.OLLAMA_API_CONFIGS.items() if k in valid_keys}
|
||||
|
||||
await Config.upsert(
|
||||
{
|
||||
'ollama.enable': form_data.ENABLE_OLLAMA_API,
|
||||
'ollama.base_urls': form_data.OLLAMA_BASE_URLS,
|
||||
'ollama.api_configs': api_configs,
|
||||
}
|
||||
)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_PROVIDER_CONFIG_UPDATED,
|
||||
actor=user,
|
||||
subject_id='ollama',
|
||||
subject_type='model.provider_config',
|
||||
data={
|
||||
'provider': 'ollama',
|
||||
'enabled': form_data.ENABLE_OLLAMA_API,
|
||||
'base_url_count': len(form_data.OLLAMA_BASE_URLS),
|
||||
},
|
||||
)
|
||||
return {
|
||||
'ENABLE_OLLAMA_API': request.app.state.config.ENABLE_OLLAMA_API,
|
||||
'OLLAMA_BASE_URLS': request.app.state.config.OLLAMA_BASE_URLS,
|
||||
'OLLAMA_API_CONFIGS': request.app.state.config.OLLAMA_API_CONFIGS,
|
||||
'ENABLE_OLLAMA_API': form_data.ENABLE_OLLAMA_API,
|
||||
'OLLAMA_BASE_URLS': form_data.OLLAMA_BASE_URLS,
|
||||
'OLLAMA_API_CONFIGS': api_configs,
|
||||
}
|
||||
|
||||
|
||||
@@ -293,29 +335,32 @@ def merge_models_lists(model_lists) -> list[dict]:
|
||||
return list(merged.values())
|
||||
|
||||
|
||||
def _resolve_api_config(request: Request, idx: int, url: str) -> dict:
|
||||
def resolve_api_config(api_configs: dict, idx: int, url: str) -> dict:
|
||||
"""Look up the API config for a backend by numeric index, falling back to URL key (legacy)."""
|
||||
api_configs = request.app.state.config.OLLAMA_API_CONFIGS
|
||||
return api_configs.get(str(idx), api_configs.get(url, {}))
|
||||
|
||||
|
||||
@cached(
|
||||
ttl=MODELS_CACHE_TTL,
|
||||
key=lambda _, user: f'ollama_all_models_{user.id}' if user else 'ollama_all_models',
|
||||
# key_builder (not key) is the per-call hook in aiocache 0.12; `key=` is a
|
||||
# static key, so a `key=lambda` collapsed every caller to one shared entry.
|
||||
key_builder=lambda _func, request, user=None: f'ollama_all_models_{user.id}' if user else 'ollama_all_models',
|
||||
)
|
||||
async def get_all_models(request: Request, user: UserModel | None = None):
|
||||
"""Aggregate model tags from every enabled Ollama backend."""
|
||||
log.info('get_all_models()')
|
||||
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
models_dict: dict = {'models': []}
|
||||
request.app.state.OLLAMA_MODELS = {}
|
||||
return models_dict
|
||||
|
||||
# Fan-out tag requests to every backend
|
||||
tasks = []
|
||||
for idx, url in enumerate(request.app.state.config.OLLAMA_BASE_URLS):
|
||||
api_config = _resolve_api_config(request, idx, url)
|
||||
base_urls = await Config.get('ollama.base_urls', [])
|
||||
api_configs = await Config.get('ollama.api_configs', {})
|
||||
for idx, url in enumerate(base_urls):
|
||||
api_config = resolve_api_config(api_configs, idx, url)
|
||||
if not api_config:
|
||||
tasks.append(send_get_request(f'{url}/api/tags', user=user))
|
||||
elif api_config.get('enable', True):
|
||||
@@ -325,12 +370,16 @@ async def get_all_models(request: Request, user: UserModel | None = None):
|
||||
|
||||
responses = await asyncio.gather(*tasks)
|
||||
|
||||
# Track which backends failed so we can skip them for /api/ps
|
||||
failed_idxs: set[int] = set()
|
||||
|
||||
# Post-process each response: apply prefix_id, tags, model filtering
|
||||
for idx, response in enumerate(responses):
|
||||
if not response:
|
||||
failed_idxs.add(idx)
|
||||
continue
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[idx]
|
||||
api_config = _resolve_api_config(request, idx, url)
|
||||
url = base_urls[idx]
|
||||
api_config = resolve_api_config(api_configs, idx, url)
|
||||
|
||||
connection_type = api_config.get('connection_type', 'local')
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
@@ -352,7 +401,7 @@ async def get_all_models(request: Request, user: UserModel | None = None):
|
||||
|
||||
# Annotate with expiry info from loaded-model state
|
||||
try:
|
||||
loaded = await get_ollama_loaded_models(request, user=user)
|
||||
loaded = await get_ollama_loaded_models(request, user=user, skip_idxs=failed_idxs)
|
||||
expires_map = {m['model']: m['expires_at'] for m in loaded['models'] if 'expires_at' in m}
|
||||
for m in models_dict['models']:
|
||||
if m['model'] in expires_map:
|
||||
@@ -394,14 +443,14 @@ async def get_ollama_tags(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""List Ollama model tags, optionally from a specific backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
if url_idx is None:
|
||||
result = await get_all_models(request, user=user)
|
||||
else:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
result = await send_request(f'{url}/api/tags', 'GET', key=key, user=user)
|
||||
|
||||
if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL:
|
||||
@@ -414,14 +463,20 @@ async def get_ollama_tags(
|
||||
async def get_ollama_loaded_models(
|
||||
request: Request,
|
||||
user=Depends(get_admin_user),
|
||||
skip_idxs: set[int] | None = None,
|
||||
) -> dict:
|
||||
"""List models currently loaded in Ollama memory across all backends."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
return {'models': []}
|
||||
|
||||
tasks = []
|
||||
for idx, url in enumerate(request.app.state.config.OLLAMA_BASE_URLS):
|
||||
api_config = _resolve_api_config(request, idx, url)
|
||||
base_urls = await Config.get('ollama.base_urls', [])
|
||||
api_configs = await Config.get('ollama.api_configs', {})
|
||||
for idx, url in enumerate(base_urls):
|
||||
if skip_idxs and idx in skip_idxs:
|
||||
tasks.append(asyncio.ensure_future(asyncio.sleep(0, None)))
|
||||
continue
|
||||
api_config = resolve_api_config(api_configs, idx, url)
|
||||
if not api_config:
|
||||
tasks.append(send_get_request(f'{url}/api/ps', user=user))
|
||||
elif api_config.get('enable', True):
|
||||
@@ -434,7 +489,7 @@ async def get_ollama_loaded_models(
|
||||
for idx, response in enumerate(responses):
|
||||
if not response:
|
||||
continue
|
||||
api_config = _resolve_api_config(request.app.state.config, idx, request.app.state.config.OLLAMA_BASE_URLS[idx])
|
||||
api_config = resolve_api_config(api_configs, idx, base_urls[idx])
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
for m in response.get('models', []):
|
||||
@@ -450,19 +505,19 @@ async def get_ollama_versions(
|
||||
url_idx: int | None = None,
|
||||
):
|
||||
"""Return the lowest Ollama version across all configured backends."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
return {'version': False}
|
||||
|
||||
if url_idx is not None:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
return await send_request(f'{url}/api/version', 'GET')
|
||||
|
||||
# Fan-out to every enabled backend
|
||||
tasks = []
|
||||
for idx, url in enumerate(request.app.state.config.OLLAMA_BASE_URLS):
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
for idx, url in enumerate(await Config.get('ollama.base_urls', [])):
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
if api_config.get('enable', True):
|
||||
tasks.append(send_get_request(f'{url}/api/version', api_config.get('key')))
|
||||
@@ -511,11 +566,11 @@ async def unload_model(
|
||||
results = []
|
||||
errors = []
|
||||
for idx in url_indices:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
str(idx), request.app.state.config.OLLAMA_API_CONFIGS.get(url, {})
|
||||
url = (await Config.get('ollama.base_urls', []))[idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(idx), (await Config.get('ollama.api_configs', {})).get(url, {})
|
||||
)
|
||||
key = get_api_key(idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
key = get_api_key(idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
if prefix_id and model.startswith(f'{prefix_id}.'):
|
||||
@@ -552,20 +607,20 @@ async def pull_model(
|
||||
url_idx: int = 0,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
form_data = form_data.model_dump(exclude_none=True)
|
||||
form_data['model'] = form_data.get('model', form_data.get('name'))
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
log.info(f'url: {url}')
|
||||
|
||||
# Admins may pull from any registry
|
||||
return await send_request(
|
||||
f'{url}/api/pull',
|
||||
payload=json.dumps({**form_data, 'insecure': True}),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -588,7 +643,7 @@ async def push_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Push a local model to a remote registry."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
if url_idx is None:
|
||||
@@ -598,13 +653,13 @@ async def push_model(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = models[form_data.model]['urls'][0]
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
log.debug(f'url: {url}')
|
||||
|
||||
return await send_request(
|
||||
f'{url}/api/push',
|
||||
payload=form_data.model_dump_json(exclude_none=True).encode(),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -627,16 +682,16 @@ async def create_model(
|
||||
url_idx: int = 0,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
log.debug(f'form_data: {form_data}')
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
|
||||
return await send_request(
|
||||
f'{url}/api/create',
|
||||
payload=form_data.model_dump_json(exclude_none=True).encode(),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -658,7 +713,7 @@ async def copy_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Duplicate an existing model under a new name."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
if url_idx is None:
|
||||
@@ -668,8 +723,8 @@ async def copy_model(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.source))
|
||||
url_idx = models[form_data.source]['urls'][0]
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
await send_request(
|
||||
f'{url}/api/copy',
|
||||
@@ -677,6 +732,13 @@ async def copy_model(
|
||||
key=key,
|
||||
user=user,
|
||||
)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_PROVIDER_MODEL_CREATED,
|
||||
actor=user,
|
||||
subject_id=form_data.destination,
|
||||
data={'provider': 'ollama', 'source': form_data.source, 'url_idx': url_idx},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
@@ -689,7 +751,7 @@ async def delete_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Remove a model from an Ollama backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = form_data.model_dump(exclude_none=True)
|
||||
@@ -703,8 +765,8 @@ async def delete_model(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(model))
|
||||
url_idx = models[model]['urls'][0]
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
await send_request(
|
||||
f'{url}/api/delete',
|
||||
@@ -713,6 +775,13 @@ async def delete_model(
|
||||
key=key,
|
||||
user=user,
|
||||
)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_PROVIDER_MODEL_DELETED,
|
||||
actor=user,
|
||||
subject_id=model,
|
||||
data={'provider': 'ollama', 'url_idx': url_idx},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
@@ -723,7 +792,7 @@ async def show_model_info(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Retrieve model metadata from the Ollama backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = form_data.model_dump(exclude_none=True)
|
||||
@@ -739,8 +808,8 @@ async def show_model_info(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(model))
|
||||
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
return await send_request(
|
||||
f'{url}/api/show',
|
||||
@@ -770,7 +839,7 @@ async def embed(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Generate embeddings via the Ollama /api/embed endpoint."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
log.info(f'generate_ollama_batch_embeddings {form_data}')
|
||||
@@ -787,12 +856,12 @@ async def embed(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -824,7 +893,7 @@ async def embeddings(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Generate embeddings via the legacy Ollama /api/embeddings endpoint."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
log.info(f'generate_ollama_embeddings {form_data}')
|
||||
@@ -841,12 +910,12 @@ async def embeddings(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS)
|
||||
key = get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {})))
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -886,7 +955,7 @@ async def generate_completion(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Run text completion via Ollama /api/generate."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
await check_model_access(user, await Models.get_model_by_id(form_data.model), BYPASS_MODEL_ACCESS_CONTROL)
|
||||
@@ -900,10 +969,10 @@ async def generate_completion(
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.MODEL_NOT_FOUND(form_data.model))
|
||||
url_idx = random.choice(models[model]['urls'])
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}),
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
@@ -913,7 +982,7 @@ async def generate_completion(
|
||||
return await send_request(
|
||||
f'{url}/api/generate',
|
||||
payload=form_data.model_dump_json(exclude_none=True).encode(),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=True,
|
||||
)
|
||||
@@ -973,7 +1042,7 @@ async def get_ollama_url(request: Request, model: str, url_idx: int | None = Non
|
||||
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(model),
|
||||
)
|
||||
url_idx = random.choice(models[model].get('urls', []))
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
return url, url_idx
|
||||
|
||||
|
||||
@@ -986,7 +1055,7 @@ async def generate_chat_completion(
|
||||
user=Depends(get_verified_user), # noqa: B008
|
||||
):
|
||||
"""Forward a chat completion request to an Ollama backend."""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
# NOTE: We intentionally do NOT use Depends(get_async_session) here.
|
||||
@@ -1035,7 +1104,7 @@ async def generate_chat_completion(
|
||||
await check_model_access(user, None, bypass_filter)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = _resolve_api_config(request, url_idx, url)
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -1044,11 +1113,13 @@ async def generate_chat_completion(
|
||||
return await send_request(
|
||||
f'{url}/api/chat',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=form_data.stream,
|
||||
content_type='application/x-ndjson',
|
||||
metadata=metadata,
|
||||
api_config=api_config,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
@@ -1121,7 +1192,7 @@ async def generate_openai_completion(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = _resolve_api_config(request, url_idx, url)
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -1130,10 +1201,12 @@ async def generate_openai_completion(
|
||||
return await send_request(
|
||||
f'{url}/v1/completions',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
metadata=metadata,
|
||||
api_config=api_config,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
@@ -1178,7 +1251,7 @@ async def generate_openai_chat_completion(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = _resolve_api_config(request, url_idx, url)
|
||||
api_config = resolve_api_config((await Config.get('ollama.api_configs', {})), url_idx, url)
|
||||
|
||||
prefix_id = api_config.get('prefix_id')
|
||||
if prefix_id:
|
||||
@@ -1187,10 +1260,12 @@ async def generate_openai_chat_completion(
|
||||
return await send_request(
|
||||
f'{url}/v1/chat/completions',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
metadata=metadata,
|
||||
api_config=api_config,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
@@ -1211,7 +1286,7 @@ async def generate_anthropic_messages(
|
||||
|
||||
See https://docs.ollama.com/api/anthropic-compatibility
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = {**form_data}
|
||||
@@ -1227,9 +1302,9 @@ async def generate_anthropic_messages(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}), # Legacy support
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}), # Legacy support
|
||||
)
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
@@ -1239,10 +1314,12 @@ async def generate_anthropic_messages(
|
||||
return await send_request(
|
||||
f'{url}/v1/messages',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
content_type='text/event-stream' if payload.get('stream', False) else None,
|
||||
api_config=api_config,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
@@ -1269,7 +1346,7 @@ async def generate_responses(
|
||||
|
||||
See https://ollama.com/blog/responses-api
|
||||
"""
|
||||
if not request.app.state.config.ENABLE_OLLAMA_API:
|
||||
if not await Config.get('ollama.enable'):
|
||||
raise HTTPException(status_code=503, detail=ERROR_MESSAGES.OLLAMA_API_DISABLED)
|
||||
|
||||
payload = form_data.model_dump()
|
||||
@@ -1285,9 +1362,9 @@ async def generate_responses(
|
||||
await check_model_access(user, None)
|
||||
|
||||
url, url_idx = await get_ollama_url(request, payload['model'], url_idx, user)
|
||||
api_config = request.app.state.config.OLLAMA_API_CONFIGS.get(
|
||||
api_config = (await Config.get('ollama.api_configs', {})).get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OLLAMA_API_CONFIGS.get(url, {}), # Legacy support
|
||||
(await Config.get('ollama.api_configs', {})).get(url, {}), # Legacy support
|
||||
)
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
@@ -1297,10 +1374,12 @@ async def generate_responses(
|
||||
return await send_request(
|
||||
f'{url}/v1/responses',
|
||||
payload=json.dumps(payload),
|
||||
key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS),
|
||||
key=get_api_key(url_idx, url, (await Config.get('ollama.api_configs', {}))),
|
||||
user=user,
|
||||
stream=payload.get('stream', False),
|
||||
content_type='text/event-stream' if payload.get('stream', False) else None,
|
||||
api_config=api_config,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
@@ -1317,7 +1396,7 @@ async def get_openai_models(
|
||||
model_list = await get_all_models(request, user=user)
|
||||
raw_models = model_list['models']
|
||||
else:
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx]
|
||||
model_list = await send_request(f'{url}/api/tags', 'GET')
|
||||
raw_models = model_list.get('models', [])
|
||||
|
||||
@@ -1394,10 +1473,13 @@ async def download_file_stream(
|
||||
|
||||
if done:
|
||||
f.close()
|
||||
hashed = calculate_sha256(file_path, chunk_size)
|
||||
hashed = await asyncio.to_thread(calculate_sha256, file_path, chunk_size)
|
||||
|
||||
with open(file_path, 'rb') as blob_f:
|
||||
blob_data = blob_f.read()
|
||||
def _read_blob():
|
||||
with open(file_path, 'rb') as blob_f:
|
||||
return blob_f.read()
|
||||
|
||||
blob_data = await asyncio.to_thread(_read_blob)
|
||||
|
||||
blob_url = f'{ollama_url}/api/blobs/sha256:{hashed}'
|
||||
async with session.post(
|
||||
@@ -1429,7 +1511,7 @@ async def download_model(
|
||||
detail='Invalid file_url. Only URLs from allowed hosts are permitted.',
|
||||
)
|
||||
|
||||
url = request.app.state.config.OLLAMA_BASE_URLS[url_idx if url_idx is not None else 0]
|
||||
url = (await Config.get('ollama.base_urls', []))[url_idx if url_idx is not None else 0]
|
||||
file_name = parse_huggingface_url(form_data.url)
|
||||
|
||||
if not file_name:
|
||||
@@ -1450,7 +1532,7 @@ async def upload_model(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Upload a local model file, push it as a blob, and create the model in Ollama."""
|
||||
ollama_url = request.app.state.config.OLLAMA_BASE_URLS[url_idx if url_idx is not None else 0]
|
||||
ollama_url = (await Config.get('ollama.base_urls', []))[url_idx if url_idx is not None else 0]
|
||||
|
||||
filename = os.path.basename(file.filename)
|
||||
file_path = os.path.join(UPLOAD_DIR, filename)
|
||||
@@ -1458,12 +1540,16 @@ async def upload_model(
|
||||
|
||||
# Stage 1: persist the uploaded file to disk
|
||||
chunk_size = 1024 * 1024 * 2 # 2 MiB
|
||||
with open(file_path, 'wb') as out_f:
|
||||
while True:
|
||||
chunk = file.file.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
out_f.write(chunk)
|
||||
|
||||
def _persist_upload():
|
||||
with open(file_path, 'wb') as out_f:
|
||||
while True:
|
||||
chunk = file.file.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
out_f.write(chunk)
|
||||
|
||||
await asyncio.to_thread(_persist_upload)
|
||||
|
||||
async def file_process_stream():
|
||||
nonlocal ollama_url
|
||||
@@ -1471,7 +1557,7 @@ async def upload_model(
|
||||
log.info(f'Total Model Size: {total_size}')
|
||||
|
||||
# Stage 2: hash the file and emit SSE progress
|
||||
file_hash = calculate_sha256(file_path, chunk_size)
|
||||
file_hash = await asyncio.to_thread(calculate_sha256, file_path, chunk_size)
|
||||
log.info(f'Model Hash: {file_hash}')
|
||||
|
||||
try:
|
||||
@@ -1483,8 +1569,11 @@ async def upload_model(
|
||||
yield f'data: {json.dumps({"progress": progress, "total": total_size, "completed": bytes_read})}\n\n'
|
||||
|
||||
# Stage 3: push blob to Ollama
|
||||
with open(file_path, 'rb') as f:
|
||||
blob_data = f.read()
|
||||
def _read_blob():
|
||||
with open(file_path, 'rb') as f:
|
||||
return f.read()
|
||||
|
||||
blob_data = await asyncio.to_thread(_read_blob)
|
||||
|
||||
session = await get_session()
|
||||
blob_url = f'{ollama_url}/api/blobs/sha256:{file_hash}'
|
||||
|
||||
@@ -22,6 +22,7 @@ from open_webui.config import (
|
||||
CACHE_DIR,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import (
|
||||
AIOHTTP_CLIENT_SESSION_SSL,
|
||||
AIOHTTP_CLIENT_TIMEOUT,
|
||||
@@ -34,10 +35,11 @@ from open_webui.env import (
|
||||
)
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.models import Models
|
||||
from open_webui.models.users import UserModel
|
||||
from open_webui.utils.access_control import check_model_access, has_connection_access
|
||||
from open_webui.utils.access_control import check_model_access, has_connection_access, has_permission
|
||||
from open_webui.utils.anthropic import get_anthropic_models, is_anthropic_url
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import get_custom_headers, include_user_info_headers
|
||||
@@ -207,7 +209,7 @@ async def get_headers_and_cookies(
|
||||
headers['Authorization'] = f'Bearer {token}'
|
||||
|
||||
if config.get('headers') and isinstance(config.get('headers'), dict):
|
||||
custom_headers = get_custom_headers(config.get('headers'), user, metadata)
|
||||
custom_headers = get_custom_headers(config.get('headers'), user, metadata, request=request)
|
||||
headers.update(custom_headers)
|
||||
|
||||
return headers, cookies
|
||||
@@ -236,15 +238,72 @@ def get_microsoft_entra_id_access_token():
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
LLAMACPP_LOADED_STATES = {'loaded', 'sleeping'}
|
||||
LLAMACPP_UNLOADED_STATES = {'loading', 'unloaded'}
|
||||
|
||||
|
||||
def get_llamacpp_model_loaded_state(model: dict, provider: str, manual_model_ids: bool = False) -> bool | None:
|
||||
if provider != 'llama.cpp':
|
||||
return None
|
||||
|
||||
status = model.get('status')
|
||||
if isinstance(status, dict):
|
||||
value = status.get('value')
|
||||
if value in LLAMACPP_LOADED_STATES:
|
||||
return True
|
||||
if value in LLAMACPP_UNLOADED_STATES:
|
||||
return False
|
||||
|
||||
if not manual_model_ids and 'status' not in model:
|
||||
return True
|
||||
|
||||
return None
|
||||
|
||||
|
||||
OPENAI_CONFIG_KEYS = {
|
||||
'ENABLE_OPENAI_API': 'openai.enable',
|
||||
'OPENAI_API_BASE_URLS': 'openai.api_base_urls',
|
||||
'OPENAI_API_KEYS': 'openai.api_keys',
|
||||
'OPENAI_API_CONFIGS': 'openai.api_configs',
|
||||
}
|
||||
|
||||
|
||||
async def get_openai_config() -> dict:
|
||||
values = await Config.get_many(*OPENAI_CONFIG_KEYS.values())
|
||||
return {field: values[storage_key] for field, storage_key in OPENAI_CONFIG_KEYS.items() if storage_key in values}
|
||||
|
||||
|
||||
async def get_openai_runtime_config() -> tuple[bool, list[str], list[str], dict]:
|
||||
values = await Config.get_many('openai.enable', 'openai.api_base_urls', 'openai.api_keys', 'openai.api_configs')
|
||||
return (
|
||||
values.get('openai.enable'),
|
||||
values.get('openai.api_base_urls') or [],
|
||||
values.get('openai.api_keys') or [],
|
||||
values.get('openai.api_configs') or {},
|
||||
)
|
||||
|
||||
|
||||
async def normalize_openai_api_keys(api_base_urls: list[str], api_keys: list[str]) -> list[str]:
|
||||
if len(api_keys) > len(api_base_urls):
|
||||
api_keys = api_keys[: len(api_base_urls)]
|
||||
elif len(api_keys) < len(api_base_urls):
|
||||
api_keys = [*api_keys, *([''] * (len(api_base_urls) - len(api_keys)))]
|
||||
|
||||
await Config.upsert({'openai.api_keys': api_keys})
|
||||
return api_keys
|
||||
|
||||
|
||||
async def get_openai_connection(idx: int) -> tuple[str, str, dict]:
|
||||
_, api_base_urls, api_keys, api_configs = await get_openai_runtime_config()
|
||||
url = api_base_urls[idx]
|
||||
key = api_keys[idx]
|
||||
api_config = api_configs.get(str(idx), api_configs.get(url, {}))
|
||||
return url, key, api_config
|
||||
|
||||
|
||||
@router.get('/config')
|
||||
async def get_config(request: Request, user=Depends(get_admin_user)):
|
||||
return {
|
||||
'ENABLE_OPENAI_API': request.app.state.config.ENABLE_OPENAI_API,
|
||||
'OPENAI_API_BASE_URLS': request.app.state.config.OPENAI_API_BASE_URLS,
|
||||
'OPENAI_API_KEYS': request.app.state.config.OPENAI_API_KEYS,
|
||||
'OPENAI_API_CONFIGS': request.app.state.config.OPENAI_API_CONFIGS,
|
||||
}
|
||||
return await get_openai_config()
|
||||
|
||||
|
||||
class OpenAIConfigForm(BaseModel):
|
||||
@@ -256,42 +315,57 @@ class OpenAIConfigForm(BaseModel):
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_config(request: Request, form_data: OpenAIConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.ENABLE_OPENAI_API = form_data.ENABLE_OPENAI_API
|
||||
request.app.state.config.OPENAI_API_BASE_URLS = form_data.OPENAI_API_BASE_URLS
|
||||
request.app.state.config.OPENAI_API_KEYS = form_data.OPENAI_API_KEYS
|
||||
api_keys = form_data.OPENAI_API_KEYS
|
||||
|
||||
# Check if API KEYS length is same than API URLS length
|
||||
if len(request.app.state.config.OPENAI_API_KEYS) != len(request.app.state.config.OPENAI_API_BASE_URLS):
|
||||
if len(request.app.state.config.OPENAI_API_KEYS) > len(request.app.state.config.OPENAI_API_BASE_URLS):
|
||||
request.app.state.config.OPENAI_API_KEYS = request.app.state.config.OPENAI_API_KEYS[
|
||||
: len(request.app.state.config.OPENAI_API_BASE_URLS)
|
||||
]
|
||||
else:
|
||||
request.app.state.config.OPENAI_API_KEYS += [''] * (
|
||||
len(request.app.state.config.OPENAI_API_BASE_URLS) - len(request.app.state.config.OPENAI_API_KEYS)
|
||||
)
|
||||
if len(api_keys) > len(form_data.OPENAI_API_BASE_URLS):
|
||||
api_keys = api_keys[: len(form_data.OPENAI_API_BASE_URLS)]
|
||||
elif len(api_keys) < len(form_data.OPENAI_API_BASE_URLS):
|
||||
api_keys = [*api_keys, *([''] * (len(form_data.OPENAI_API_BASE_URLS) - len(api_keys)))]
|
||||
|
||||
request.app.state.config.OPENAI_API_CONFIGS = form_data.OPENAI_API_CONFIGS
|
||||
valid_keys = set(map(str, range(len(form_data.OPENAI_API_BASE_URLS))))
|
||||
api_configs = {key: value for key, value in form_data.OPENAI_API_CONFIGS.items() if key in valid_keys}
|
||||
|
||||
# Remove the API configs that are not in the API URLS
|
||||
keys = list(map(str, range(len(request.app.state.config.OPENAI_API_BASE_URLS))))
|
||||
request.app.state.config.OPENAI_API_CONFIGS = {
|
||||
key: value for key, value in request.app.state.config.OPENAI_API_CONFIGS.items() if key in keys
|
||||
}
|
||||
await Config.upsert(
|
||||
{
|
||||
'openai.enable': form_data.ENABLE_OPENAI_API,
|
||||
'openai.api_base_urls': form_data.OPENAI_API_BASE_URLS,
|
||||
'openai.api_keys': api_keys,
|
||||
'openai.api_configs': api_configs,
|
||||
}
|
||||
)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.MODEL_PROVIDER_CONFIG_UPDATED,
|
||||
actor=user,
|
||||
subject_id='openai',
|
||||
subject_type='model.provider_config',
|
||||
data={
|
||||
'provider': 'openai',
|
||||
'enabled': form_data.ENABLE_OPENAI_API,
|
||||
'base_url_count': len(form_data.OPENAI_API_BASE_URLS),
|
||||
},
|
||||
)
|
||||
|
||||
return {
|
||||
'ENABLE_OPENAI_API': request.app.state.config.ENABLE_OPENAI_API,
|
||||
'OPENAI_API_BASE_URLS': request.app.state.config.OPENAI_API_BASE_URLS,
|
||||
'OPENAI_API_KEYS': request.app.state.config.OPENAI_API_KEYS,
|
||||
'OPENAI_API_CONFIGS': request.app.state.config.OPENAI_API_CONFIGS,
|
||||
'ENABLE_OPENAI_API': form_data.ENABLE_OPENAI_API,
|
||||
'OPENAI_API_BASE_URLS': form_data.OPENAI_API_BASE_URLS,
|
||||
'OPENAI_API_KEYS': api_keys,
|
||||
'OPENAI_API_CONFIGS': api_configs,
|
||||
}
|
||||
|
||||
|
||||
@router.post('/audio/speech')
|
||||
async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
if user.role != 'admin' and not await has_permission(user.id, 'chat.tts', await Config.get('user.permissions')):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
idx = None
|
||||
try:
|
||||
idx = request.app.state.config.OPENAI_API_BASE_URLS.index('https://api.openai.com/v1')
|
||||
_, api_base_urls, _, _ = await get_openai_runtime_config()
|
||||
idx = api_base_urls.index('https://api.openai.com/v1')
|
||||
|
||||
body = await request.body()
|
||||
name = hashlib.sha256(body).hexdigest()
|
||||
@@ -305,12 +379,7 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
if file_path.is_file():
|
||||
return FileResponse(file_path)
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, user=user)
|
||||
|
||||
@@ -360,29 +429,15 @@ async def speech(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
async def get_all_models_responses(request: Request, user: UserModel) -> list:
|
||||
if not request.app.state.config.ENABLE_OPENAI_API:
|
||||
enable_openai_api, api_base_urls, api_keys, api_configs = await get_openai_runtime_config()
|
||||
if not enable_openai_api:
|
||||
return []
|
||||
|
||||
# Cache config values locally to avoid repeated Redis lookups.
|
||||
# Each access to request.app.state.config.<KEY> triggers a Redis GET;
|
||||
# caching here avoids hundreds of redundant round-trips.
|
||||
api_base_urls = request.app.state.config.OPENAI_API_BASE_URLS
|
||||
api_keys = list(request.app.state.config.OPENAI_API_KEYS)
|
||||
api_configs = request.app.state.config.OPENAI_API_CONFIGS
|
||||
|
||||
# Check if API KEYS length is same than API URLS length
|
||||
num_urls = len(api_base_urls)
|
||||
num_keys = len(api_keys)
|
||||
|
||||
if num_keys != num_urls:
|
||||
# if there are more keys than urls, remove the extra keys
|
||||
if num_keys > num_urls:
|
||||
api_keys = api_keys[:num_urls]
|
||||
request.app.state.config.OPENAI_API_KEYS = api_keys
|
||||
# if there are more urls than keys, add empty keys
|
||||
else:
|
||||
api_keys += [''] * (num_urls - num_keys)
|
||||
request.app.state.config.OPENAI_API_KEYS = api_keys
|
||||
api_keys = await normalize_openai_api_keys(api_base_urls, api_keys)
|
||||
|
||||
request_tasks = []
|
||||
for idx, url in enumerate(api_base_urls):
|
||||
@@ -487,18 +542,17 @@ async def get_filtered_models(models, user, db=None):
|
||||
|
||||
@cached(
|
||||
ttl=MODELS_CACHE_TTL,
|
||||
key=lambda _, user: f'openai_all_models_{user.id}' if user else 'openai_all_models',
|
||||
# key_builder (not key) is the per-call hook in aiocache 0.12; `key=` is a
|
||||
# static key, so a `key=lambda` collapsed every caller to one shared entry.
|
||||
key_builder=lambda _func, request, user=None: f'openai_all_models_{user.id}' if user else 'openai_all_models',
|
||||
)
|
||||
async def get_all_models(request: Request, user: UserModel) -> dict[str, list]:
|
||||
log.info('get_all_models()')
|
||||
|
||||
if not request.app.state.config.ENABLE_OPENAI_API:
|
||||
enable_openai_api, api_base_urls, _, api_configs = await get_openai_runtime_config()
|
||||
if not enable_openai_api:
|
||||
return {'data': []}
|
||||
|
||||
# Cache config value locally to avoid repeated Redis lookups inside
|
||||
# the nested loop in get_merged_models (one GET per model otherwise).
|
||||
api_base_urls = request.app.state.config.OPENAI_API_BASE_URLS
|
||||
|
||||
responses = await get_all_models_responses(request, user=user)
|
||||
|
||||
def extract_data(response):
|
||||
@@ -539,21 +593,25 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]:
|
||||
continue
|
||||
|
||||
if model_id and model_id not in models:
|
||||
api_config = api_configs.get(str(idx), api_configs.get(base_url, {}))
|
||||
provider = model.get('provider', '')
|
||||
merged = {
|
||||
**model,
|
||||
'name': model.get('name', model_id),
|
||||
'owned_by': 'openai',
|
||||
'openai': model,
|
||||
'connection_type': model.get('connection_type', 'external'),
|
||||
'provider': model.get('provider', ''),
|
||||
'provider': provider,
|
||||
'urlIdx': idx,
|
||||
}
|
||||
|
||||
# llama.cpp router mode: derive loaded state from
|
||||
# the status object returned by GET /v1/models.
|
||||
status = model.get('status')
|
||||
if isinstance(status, dict) and 'value' in status:
|
||||
merged['loaded'] = status['value'] in ('loaded', 'sleeping')
|
||||
loaded = get_llamacpp_model_loaded_state(
|
||||
model,
|
||||
provider,
|
||||
manual_model_ids=bool(api_config.get('model_ids')),
|
||||
)
|
||||
if loaded is not None:
|
||||
merged['loaded'] = loaded
|
||||
|
||||
models[model_id] = merged
|
||||
|
||||
@@ -569,7 +627,7 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]:
|
||||
@router.get('/models')
|
||||
@router.get('/models/{url_idx}')
|
||||
async def get_models(request: Request, url_idx: int | None = None, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_OPENAI_API:
|
||||
if not await Config.get('openai.enable'):
|
||||
raise HTTPException(status_code=503, detail='OpenAI API is disabled')
|
||||
|
||||
models = {
|
||||
@@ -579,13 +637,7 @@ async def get_models(request: Request, url_idx: int | None = None, user=Depends(
|
||||
if url_idx is None:
|
||||
models = await get_all_models(request, user=user)
|
||||
else:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[url_idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[url_idx]
|
||||
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(url_idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(url_idx)
|
||||
|
||||
r = None
|
||||
async with aiohttp.ClientSession(
|
||||
@@ -1114,13 +1166,7 @@ async def generate_chat_completion(
|
||||
detail=ERROR_MESSAGES.MODEL_NOT_FOUND(),
|
||||
)
|
||||
|
||||
# Get the API config for the model
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
request.app.state.config.OPENAI_API_BASE_URLS[idx], {}
|
||||
), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
prefix_id = api_config.get('prefix_id', None)
|
||||
if prefix_id:
|
||||
@@ -1135,9 +1181,6 @@ async def generate_chat_completion(
|
||||
'role': user.role,
|
||||
}
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
|
||||
# Check if model is a reasoning model that needs special handling
|
||||
if is_openai_new_model(payload['model']):
|
||||
payload = openai_reasoning_model_handler(payload)
|
||||
@@ -1303,12 +1346,7 @@ async def embeddings(request: Request, form_data: dict, user):
|
||||
if model_id in models:
|
||||
idx = models[model_id]['urlIdx']
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
r = None
|
||||
streaming = False
|
||||
@@ -1426,12 +1464,7 @@ async def responses(
|
||||
if model_id in models:
|
||||
idx = models[model_id]['urlIdx']
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(url, {}), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
r = None
|
||||
streaming = False
|
||||
@@ -1535,14 +1568,7 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
|
||||
if model_id in models:
|
||||
idx = models[model_id]['urlIdx']
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[idx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[idx]
|
||||
api_config = request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
str(idx),
|
||||
request.app.state.config.OPENAI_API_CONFIGS.get(
|
||||
request.app.state.config.OPENAI_API_BASE_URLS[idx], {}
|
||||
), # Legacy support
|
||||
)
|
||||
url, key, api_config = await get_openai_connection(idx)
|
||||
|
||||
r = None
|
||||
streaming = False
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
@@ -18,6 +19,8 @@ from fastapi import (
|
||||
from open_webui.config import CACHE_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.routers.openai import get_all_models_responses
|
||||
from open_webui.utils.auth import get_admin_user
|
||||
from pydantic import BaseModel
|
||||
@@ -51,6 +54,12 @@ def get_sorted_filters(model_id, models):
|
||||
return sorted_filters
|
||||
|
||||
|
||||
async def get_openai_connection(url_idx: int) -> tuple[str, str]:
|
||||
base_urls = await Config.get('openai.api_base_urls', [])
|
||||
api_keys = await Config.get('openai.api_keys', [])
|
||||
return base_urls[url_idx], api_keys[url_idx]
|
||||
|
||||
|
||||
async def process_pipeline_inlet_filter(request, payload, user, models):
|
||||
user = {'id': user.id, 'email': user.email, 'name': user.name, 'role': user.role}
|
||||
model_id = payload['model']
|
||||
@@ -69,8 +78,7 @@ async def process_pipeline_inlet_filter(request, payload, user, models):
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
if not key:
|
||||
continue
|
||||
@@ -133,8 +141,7 @@ async def process_pipeline_outlet_filter(request, payload, user, models):
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
if not key:
|
||||
continue
|
||||
@@ -194,11 +201,12 @@ async def get_pipelines_list(request: Request, user=Depends(get_admin_user)):
|
||||
log.debug(f'get_pipelines_list: get_openai_models_responses returned {responses}')
|
||||
|
||||
urlIdxs = [idx for idx, response in enumerate(responses) if response is not None and 'pipelines' in response]
|
||||
base_urls = await Config.get('openai.api_base_urls', [])
|
||||
|
||||
return {
|
||||
'data': [
|
||||
{
|
||||
'url': request.app.state.config.OPENAI_API_BASE_URLS[urlIdx],
|
||||
'url': base_urls[urlIdx],
|
||||
'idx': urlIdx,
|
||||
}
|
||||
for urlIdx in urlIdxs
|
||||
@@ -229,12 +237,14 @@ async def upload_pipeline(
|
||||
|
||||
response = None
|
||||
try:
|
||||
# Save the uploaded file
|
||||
with open(file_path, 'wb') as buffer:
|
||||
shutil.copyfileobj(file.file, buffer)
|
||||
# Save the uploaded file off the event loop (uploads can be large).
|
||||
def _save_upload():
|
||||
with open(file_path, 'wb') as buffer:
|
||||
shutil.copyfileobj(file.file, buffer)
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
await asyncio.to_thread(_save_upload)
|
||||
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
headers = {'Authorization': f'Bearer {key}'}
|
||||
|
||||
@@ -257,6 +267,13 @@ async def upload_pipeline(
|
||||
response.raise_for_status()
|
||||
data = await response.json()
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PIPELINE_UPLOADED,
|
||||
actor=user,
|
||||
subject_id=data.get('id') or filename,
|
||||
data={'url_idx': urlIdx, 'filename': filename},
|
||||
)
|
||||
return {**data}
|
||||
except Exception as e:
|
||||
# Handle connection error here
|
||||
@@ -294,8 +311,7 @@ async def add_pipeline(request: Request, form_data: AddPipelineForm, user=Depend
|
||||
try:
|
||||
urlIdx = form_data.urlIdx
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.post(
|
||||
@@ -307,6 +323,13 @@ async def add_pipeline(request: Request, form_data: AddPipelineForm, user=Depend
|
||||
response.raise_for_status()
|
||||
data = await response.json()
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PIPELINE_ADDED,
|
||||
actor=user,
|
||||
subject_id=data.get('id') or form_data.url,
|
||||
data={'url_idx': urlIdx, 'url': form_data.url},
|
||||
)
|
||||
return {**data}
|
||||
except Exception as e:
|
||||
# Handle connection error here
|
||||
@@ -338,8 +361,7 @@ async def delete_pipeline(request: Request, form_data: DeletePipelineForm, user=
|
||||
try:
|
||||
urlIdx = form_data.urlIdx
|
||||
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.delete(
|
||||
@@ -351,6 +373,13 @@ async def delete_pipeline(request: Request, form_data: DeletePipelineForm, user=
|
||||
response.raise_for_status()
|
||||
data = await response.json()
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PIPELINE_DELETED,
|
||||
actor=user,
|
||||
subject_id=form_data.id,
|
||||
data={'url_idx': urlIdx},
|
||||
)
|
||||
return {**data}
|
||||
except Exception as e:
|
||||
# Handle connection error here
|
||||
@@ -375,8 +404,7 @@ async def delete_pipeline(request: Request, form_data: DeletePipelineForm, user=
|
||||
async def get_pipelines(request: Request, urlIdx: Optional[int] = None, user=Depends(get_admin_user)):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(
|
||||
@@ -416,8 +444,7 @@ async def get_pipeline_valves(
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(
|
||||
@@ -428,6 +455,13 @@ async def get_pipeline_valves(
|
||||
response.raise_for_status()
|
||||
data = await response.json()
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PIPELINE_VALVES_UPDATED,
|
||||
actor=user,
|
||||
subject_id=pipeline_id,
|
||||
data={'url_idx': urlIdx},
|
||||
)
|
||||
return {**data}
|
||||
except Exception as e:
|
||||
# Handle connection error here
|
||||
@@ -457,8 +491,7 @@ async def get_pipeline_valves_spec(
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.get(
|
||||
@@ -499,8 +532,7 @@ async def update_pipeline_valves(
|
||||
):
|
||||
response = None
|
||||
try:
|
||||
url = request.app.state.config.OPENAI_API_BASE_URLS[urlIdx]
|
||||
key = request.app.state.config.OPENAI_API_KEYS[urlIdx]
|
||||
url, key = await get_openai_connection(urlIdx)
|
||||
|
||||
async with aiohttp.ClientSession(trust_env=True) as session:
|
||||
async with session.post(
|
||||
|
||||
@@ -5,8 +5,10 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.prompt_history import (
|
||||
PromptHistories,
|
||||
@@ -149,13 +151,13 @@ async def create_new_prompt(
|
||||
await has_permission(
|
||||
user.id,
|
||||
'workspace.prompts',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
or await has_permission(
|
||||
user.id,
|
||||
'workspace.prompts_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -165,7 +167,7 @@ async def create_new_prompt(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -177,6 +179,13 @@ async def create_new_prompt(
|
||||
prompt = await Prompts.insert_new_prompt(user.id, form_data, db=db)
|
||||
|
||||
if prompt:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_CREATED,
|
||||
actor=user,
|
||||
subject_id=prompt.id,
|
||||
data={'name': prompt.name, 'command': prompt.command},
|
||||
)
|
||||
return prompt
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
@@ -281,7 +290,7 @@ async def update_prompt_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -291,6 +300,13 @@ async def update_prompt_by_id(
|
||||
# Use the ID from the found prompt
|
||||
updated_prompt = await Prompts.update_prompt_by_id(prompt.id, form_data, user.id, db=db)
|
||||
if updated_prompt:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=updated_prompt.id,
|
||||
data={'name': updated_prompt.name, 'command': updated_prompt.command},
|
||||
)
|
||||
return updated_prompt
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -306,6 +322,7 @@ async def update_prompt_by_id(
|
||||
|
||||
@router.post('/id/{prompt_id}/update/meta', response_model=PromptModel | None)
|
||||
async def update_prompt_metadata(
|
||||
request: Request,
|
||||
prompt_id: str,
|
||||
form_data: PromptMetadataForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -349,6 +366,13 @@ async def update_prompt_metadata(
|
||||
prompt.id, form_data.name, form_data.command, form_data.tags, db=db
|
||||
)
|
||||
if updated_prompt:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_UPDATED,
|
||||
actor=user,
|
||||
subject_id=updated_prompt.id,
|
||||
data={'name': updated_prompt.name, 'command': updated_prompt.command},
|
||||
)
|
||||
return updated_prompt
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -359,6 +383,7 @@ async def update_prompt_metadata(
|
||||
|
||||
@router.post('/id/{prompt_id}/update/version', response_model=PromptModel | None)
|
||||
async def set_prompt_version(
|
||||
request: Request,
|
||||
prompt_id: str,
|
||||
form_data: PromptVersionUpdateForm,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -389,6 +414,13 @@ async def set_prompt_version(
|
||||
|
||||
updated_prompt = await Prompts.update_prompt_version(prompt.id, form_data.version_id, db=db)
|
||||
if updated_prompt:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_VERSION_UPDATED,
|
||||
actor=user,
|
||||
subject_id=updated_prompt.id,
|
||||
data={'version_id': updated_prompt.version_id},
|
||||
)
|
||||
return updated_prompt
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -438,7 +470,7 @@ async def update_prompt_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -447,7 +479,15 @@ async def update_prompt_access_by_id(
|
||||
|
||||
await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
|
||||
|
||||
return await Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
updated_prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_ACCESS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=prompt_id,
|
||||
data={'name': updated_prompt.name if updated_prompt else None},
|
||||
)
|
||||
return updated_prompt
|
||||
|
||||
|
||||
############################
|
||||
@@ -457,7 +497,10 @@ async def update_prompt_access_by_id(
|
||||
|
||||
@router.post('/id/{prompt_id}/toggle', response_model=PromptModel | None)
|
||||
async def toggle_prompt_active(
|
||||
prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
request: Request,
|
||||
prompt_id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
@@ -485,6 +528,14 @@ async def toggle_prompt_active(
|
||||
|
||||
result = await Prompts.toggle_prompt_active(prompt.id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_ENABLED if result.is_active else EVENTS.PROMPT_DISABLED,
|
||||
actor=user,
|
||||
subject_id=result.id,
|
||||
subject_type='prompt',
|
||||
data={'name': result.name, 'command': result.command},
|
||||
)
|
||||
return result
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
@@ -499,7 +550,10 @@ async def toggle_prompt_active(
|
||||
|
||||
@router.delete('/id/{prompt_id}/delete', response_model=bool)
|
||||
async def delete_prompt_by_id(
|
||||
prompt_id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)
|
||||
request: Request,
|
||||
prompt_id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
prompt = await Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
@@ -526,6 +580,14 @@ async def delete_prompt_by_id(
|
||||
)
|
||||
|
||||
result = await Prompts.delete_prompt_by_id(prompt.id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.PROMPT_DELETED,
|
||||
actor=user,
|
||||
subject_id=prompt.id,
|
||||
data={'name': prompt.name, 'command': prompt.command},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
|
||||
+1104
-757
File diff suppressed because it is too large
Load Diff
@@ -16,6 +16,7 @@ from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, s
|
||||
from fastapi.responses import JSONResponse
|
||||
from open_webui.config import OAUTH_PROVIDERS
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import SCIM_AUTH_PROVIDER
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.groups import GroupModel, Groups
|
||||
@@ -259,10 +260,6 @@ def get_scim_auth(request: Request, authorization: Optional[str] = Header(None))
|
||||
enable_scim = getattr(request.app.state, 'ENABLE_SCIM', False)
|
||||
log.info(f'SCIM auth check - raw ENABLE_SCIM: {enable_scim}, type: {type(enable_scim)}')
|
||||
|
||||
# Handle both ConfigVar and direct value
|
||||
if hasattr(enable_scim, 'value'):
|
||||
enable_scim = enable_scim.value
|
||||
|
||||
if not enable_scim:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -271,9 +268,6 @@ def get_scim_auth(request: Request, authorization: Optional[str] = Header(None))
|
||||
|
||||
# Verify the SCIM token
|
||||
scim_token = getattr(request.app.state, 'SCIM_TOKEN', None)
|
||||
# Handle both ConfigVar and direct value
|
||||
if hasattr(scim_token, 'value'):
|
||||
scim_token = scim_token.value
|
||||
log.debug(f'SCIM token configured: {bool(scim_token)}')
|
||||
if not scim_token or not hmac.compare_digest(token, scim_token):
|
||||
raise HTTPException(
|
||||
@@ -636,6 +630,18 @@ async def create_user(
|
||||
await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db)
|
||||
new_user = await Users.get_user_by_id(user_id, db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_CREATED,
|
||||
subject_id=new_user.id,
|
||||
source='scim',
|
||||
data={
|
||||
'email': new_user.email,
|
||||
'role': new_user.role,
|
||||
'external_id': user_data.externalId,
|
||||
},
|
||||
)
|
||||
|
||||
return await user_to_scim(new_user, request, db=db)
|
||||
|
||||
|
||||
@@ -672,7 +678,10 @@ async def update_user(
|
||||
if user_data.emails and len(user_data.emails) > 0:
|
||||
update_data['email'] = user_data.emails[0].value
|
||||
|
||||
if user_data.active is not None:
|
||||
# Do not let SCIM's active flag demote an existing admin: a routine IdP sync or misconfiguration
|
||||
# must not silently strip a locally-provisioned admin's role and lock the instance out. Admin
|
||||
# role changes go through the dedicated admin endpoints, not SCIM provisioning.
|
||||
if user_data.active is not None and user.role != 'admin':
|
||||
update_data['role'] = 'user' if user_data.active else 'pending'
|
||||
|
||||
if user_data.photos and len(user_data.photos) > 0:
|
||||
@@ -691,6 +700,16 @@ async def update_user(
|
||||
await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db)
|
||||
updated_user = await Users.get_user_by_id(user_id, db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_UPDATED,
|
||||
subject_id=user_id,
|
||||
source='scim',
|
||||
data={
|
||||
'updated_fields': list(update_data.keys()) + (['externalId'] if user_data.externalId else []),
|
||||
},
|
||||
)
|
||||
|
||||
return await user_to_scim(updated_user, request, db=db)
|
||||
|
||||
|
||||
@@ -719,7 +738,9 @@ async def patch_user(
|
||||
|
||||
if op == 'replace':
|
||||
if path == 'active':
|
||||
update_data['role'] = 'user' if value else 'pending'
|
||||
# Same guard as update_user: never demote an existing admin via SCIM.
|
||||
if user.role != 'admin':
|
||||
update_data['role'] = 'user' if value else 'pending'
|
||||
elif path == 'userName':
|
||||
update_data['email'] = value
|
||||
elif path == 'displayName':
|
||||
@@ -743,6 +764,14 @@ async def patch_user(
|
||||
else:
|
||||
updated_user = user
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_UPDATED,
|
||||
subject_id=user_id,
|
||||
source='scim',
|
||||
data={'updated_fields': list(update_data.keys())},
|
||||
)
|
||||
|
||||
return await user_to_scim(updated_user, request, db=db)
|
||||
|
||||
|
||||
@@ -768,6 +797,14 @@ async def delete_user(
|
||||
detail='Failed to delete user',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_DELETED,
|
||||
subject_id=user_id,
|
||||
source='scim',
|
||||
data={'email': user.email},
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@@ -885,6 +922,22 @@ async def create_group(
|
||||
|
||||
new_group = await Groups.get_group_by_id(new_group.id, db=db)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_CREATED,
|
||||
subject_id=new_group.id,
|
||||
source='scim',
|
||||
data={'name': new_group.name, 'member_ids': member_ids, 'member_count': len(member_ids)},
|
||||
)
|
||||
if member_ids:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_ADDED,
|
||||
subject_id=new_group.id,
|
||||
source='scim',
|
||||
data={'member_ids': member_ids, 'count': len(member_ids)},
|
||||
)
|
||||
|
||||
return await group_to_scim(new_group, request, db=db)
|
||||
|
||||
|
||||
@@ -913,9 +966,15 @@ async def update_group(
|
||||
)
|
||||
|
||||
# Handle members if provided
|
||||
added_member_ids = []
|
||||
removed_member_ids = []
|
||||
if group_data.members is not None:
|
||||
old_member_ids = set(await Groups.get_group_user_ids_by_id(group_id, db) or [])
|
||||
member_ids = [member.value for member in group_data.members]
|
||||
await Groups.set_group_user_ids_by_id(group_id, member_ids, db=db)
|
||||
new_member_ids = set(member_ids)
|
||||
added_member_ids = sorted(new_member_ids - old_member_ids)
|
||||
removed_member_ids = sorted(old_member_ids - new_member_ids)
|
||||
|
||||
# Update group
|
||||
updated_group = await Groups.update_group_by_id(group_id, update_form, db=db)
|
||||
@@ -925,6 +984,30 @@ async def update_group(
|
||||
detail='Failed to update group',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_UPDATED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'updated_fields': ['name', 'members'] if group_data.members is not None else ['name']},
|
||||
)
|
||||
if added_member_ids:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_ADDED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'member_ids': added_member_ids, 'count': len(added_member_ids)},
|
||||
)
|
||||
if removed_member_ids:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_REMOVED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'member_ids': removed_member_ids, 'count': len(removed_member_ids)},
|
||||
)
|
||||
|
||||
return await group_to_scim(updated_group, request, db=db)
|
||||
|
||||
|
||||
@@ -950,6 +1033,8 @@ async def patch_group(
|
||||
name=group.name,
|
||||
description=group.description,
|
||||
)
|
||||
added_member_ids = []
|
||||
removed_member_ids = []
|
||||
|
||||
for operation in patch_data.Operations:
|
||||
op = operation.op.lower()
|
||||
@@ -961,7 +1046,12 @@ async def patch_group(
|
||||
update_form.name = value
|
||||
elif path == 'members':
|
||||
# Replace all members
|
||||
await Groups.set_group_user_ids_by_id(group_id, [member['value'] for member in value], db=db)
|
||||
old_member_ids = set(await Groups.get_group_user_ids_by_id(group_id, db) or [])
|
||||
new_member_ids = [member['value'] for member in value]
|
||||
await Groups.set_group_user_ids_by_id(group_id, new_member_ids, db=db)
|
||||
new_member_ids_set = set(new_member_ids)
|
||||
added_member_ids.extend(sorted(new_member_ids_set - old_member_ids))
|
||||
removed_member_ids.extend(sorted(old_member_ids - new_member_ids_set))
|
||||
|
||||
elif op == 'add':
|
||||
if path == 'members':
|
||||
@@ -970,11 +1060,13 @@ async def patch_group(
|
||||
for member in value:
|
||||
if isinstance(member, dict) and 'value' in member:
|
||||
await Groups.add_users_to_group(group_id, [member['value']], db=db)
|
||||
added_member_ids.append(member['value'])
|
||||
elif op == 'remove':
|
||||
if path and path.startswith('members[value eq'):
|
||||
# Remove specific member
|
||||
member_id = path.split('"')[1]
|
||||
await Groups.remove_users_from_group(group_id, [member_id], db=db)
|
||||
removed_member_ids.append(member_id)
|
||||
|
||||
# Update group
|
||||
updated_group = await Groups.update_group_by_id(group_id, update_form, db=db)
|
||||
@@ -984,6 +1076,30 @@ async def patch_group(
|
||||
detail='Failed to update group',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_UPDATED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'operation_count': len(patch_data.Operations)},
|
||||
)
|
||||
if added_member_ids:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_ADDED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'member_ids': sorted(set(added_member_ids)), 'count': len(set(added_member_ids))},
|
||||
)
|
||||
if removed_member_ids:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_MEMBER_REMOVED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'member_ids': sorted(set(removed_member_ids)), 'count': len(set(removed_member_ids))},
|
||||
)
|
||||
|
||||
return await group_to_scim(updated_group, request, db=db)
|
||||
|
||||
|
||||
@@ -1009,4 +1125,12 @@ async def delete_group(
|
||||
detail='Failed to delete group',
|
||||
)
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.GROUP_DELETED,
|
||||
subject_id=group_id,
|
||||
source='scim',
|
||||
data={'name': group.name},
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
@@ -4,8 +4,10 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.skills import (
|
||||
SkillAccessListResponse,
|
||||
@@ -129,8 +131,8 @@ async def export_skills(
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.skills',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
'workspace.skills_export',
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -156,8 +158,9 @@ async def create_new_skill(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'workspace.skills', request.app.state.config.USER_PERMISSIONS, db=db
|
||||
if user.role != 'admin' and not (
|
||||
await has_permission(user.id, 'workspace.skills', await Config.get('user.permissions'), db=db)
|
||||
or await has_permission(user.id, 'workspace.skills_import', await Config.get('user.permissions'), db=db)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -179,7 +182,7 @@ async def create_new_skill(
|
||||
# grants in the create payload, bypassing the sharing.public_skills gate
|
||||
# that the dedicated /access/update endpoint already enforces.
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -189,17 +192,26 @@ async def create_new_skill(
|
||||
try:
|
||||
skill = await Skills.insert_new_skill(user.id, form_data, db=db)
|
||||
if skill:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.SKILL_CREATED,
|
||||
actor=user,
|
||||
subject_id=skill.id,
|
||||
data={'name': skill.name},
|
||||
)
|
||||
return skill
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error creating skill'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to create skill: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error creating skill'),
|
||||
)
|
||||
|
||||
|
||||
@@ -292,7 +304,7 @@ async def update_skill_by_id(
|
||||
# they may set, so a non-admin owner cannot make their own skill publicly
|
||||
# readable/writable without sharing.public_skills permission.
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -307,16 +319,25 @@ async def update_skill_by_id(
|
||||
skill = await Skills.update_skill_by_id(id, updated, db=db)
|
||||
|
||||
if skill:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.SKILL_UPDATED,
|
||||
actor=user,
|
||||
subject_id=skill.id,
|
||||
data={'name': skill.name},
|
||||
)
|
||||
return skill
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error updating skill'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating skill'),
|
||||
)
|
||||
|
||||
|
||||
@@ -361,7 +382,7 @@ async def update_skill_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -370,7 +391,15 @@ async def update_skill_access_by_id(
|
||||
|
||||
await AccessGrants.set_access_grants('skill', id, form_data.access_grants, db=db)
|
||||
|
||||
return await Skills.get_skill_by_id(id, db=db)
|
||||
skill = await Skills.get_skill_by_id(id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.SKILL_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'access_updated': True, 'name': skill.name if skill else None},
|
||||
)
|
||||
return skill
|
||||
|
||||
|
||||
############################
|
||||
@@ -379,7 +408,12 @@ async def update_skill_access_by_id(
|
||||
|
||||
|
||||
@router.post('/id/{id}/toggle', response_model=Optional[SkillModel])
|
||||
async def toggle_skill_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def toggle_skill_by_id(
|
||||
request: Request,
|
||||
id: str,
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
skill = await Skills.get_skill_by_id(id, db=db)
|
||||
if skill:
|
||||
if (
|
||||
@@ -396,6 +430,13 @@ async def toggle_skill_by_id(id: str, user=Depends(get_verified_user), db: Async
|
||||
skill = await Skills.toggle_skill_by_id(id, db=db)
|
||||
|
||||
if skill:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.SKILL_ENABLED if skill.is_active else EVENTS.SKILL_DISABLED,
|
||||
actor=user,
|
||||
subject_id=skill.id,
|
||||
data={'name': skill.name},
|
||||
)
|
||||
return skill
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -450,4 +491,12 @@ async def delete_skill_by_id(
|
||||
)
|
||||
|
||||
result = await Skills.delete_skill_by_id(id, db=db)
|
||||
if result:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.SKILL_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': skill.name},
|
||||
)
|
||||
return result
|
||||
|
||||
@@ -16,6 +16,7 @@ from open_webui.config import (
|
||||
DEFAULT_VOICE_MODE_PROMPT_TEMPLATE,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES, TASKS
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.routers.pipelines import process_pipeline_inlet_filter
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.chat import generate_chat_completion
|
||||
@@ -36,6 +37,36 @@ log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
TASK_CONFIG_KEYS = {
|
||||
'TASK_MODEL': 'task.model.default',
|
||||
'TASK_MODEL_EXTERNAL': 'task.model.external',
|
||||
'TITLE_GENERATION_PROMPT_TEMPLATE': 'task.title.prompt_template',
|
||||
'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': 'task.image.prompt_template',
|
||||
'ENABLE_AUTOCOMPLETE_GENERATION': 'task.autocomplete.enable',
|
||||
'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': 'task.autocomplete.input_max_length',
|
||||
'AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE': 'task.autocomplete.prompt_template',
|
||||
'TAGS_GENERATION_PROMPT_TEMPLATE': 'task.tags.prompt_template',
|
||||
'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': 'task.follow_up.prompt_template',
|
||||
'ENABLE_FOLLOW_UP_GENERATION': 'task.follow_up.enable',
|
||||
'ENABLE_TAGS_GENERATION': 'task.tags.enable',
|
||||
'ENABLE_TITLE_GENERATION': 'task.title.enable',
|
||||
'ENABLE_SEARCH_QUERY_GENERATION': 'task.query.search.enable',
|
||||
'ENABLE_RETRIEVAL_QUERY_GENERATION': 'task.query.retrieval.enable',
|
||||
'QUERY_GENERATION_PROMPT_TEMPLATE': 'task.query.prompt_template',
|
||||
'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': 'task.tools.prompt_template',
|
||||
'ENABLE_VOICE_MODE_PROMPT': 'task.voice.prompt.enable',
|
||||
'VOICE_MODE_PROMPT_TEMPLATE': 'task.voice.prompt_template',
|
||||
}
|
||||
|
||||
|
||||
async def get_config_values(key_map: dict[str, str]) -> dict:
|
||||
values = await Config.get_many(*key_map.values())
|
||||
return {field: values[storage_key] for field, storage_key in key_map.items() if storage_key in values}
|
||||
|
||||
|
||||
def config_updates(data: dict, key_map: dict[str, str]) -> dict:
|
||||
return {key_map[field]: value for field, value in data.items() if field in key_map}
|
||||
|
||||
|
||||
##################################
|
||||
#
|
||||
@@ -59,25 +90,7 @@ async def check_active_chats(request: Request, form_data: ActiveChatsForm, user=
|
||||
|
||||
@router.get('/config')
|
||||
async def get_task_config(request: Request, user=Depends(get_verified_user)):
|
||||
return {
|
||||
'TASK_MODEL': request.app.state.config.TASK_MODEL,
|
||||
'TASK_MODEL_EXTERNAL': request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
'TITLE_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE,
|
||||
'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_AUTOCOMPLETE_GENERATION': request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION,
|
||||
'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH,
|
||||
'TAGS_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE,
|
||||
'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_FOLLOW_UP_GENERATION': request.app.state.config.ENABLE_FOLLOW_UP_GENERATION,
|
||||
'ENABLE_TAGS_GENERATION': request.app.state.config.ENABLE_TAGS_GENERATION,
|
||||
'ENABLE_TITLE_GENERATION': request.app.state.config.ENABLE_TITLE_GENERATION,
|
||||
'ENABLE_SEARCH_QUERY_GENERATION': request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION,
|
||||
'ENABLE_RETRIEVAL_QUERY_GENERATION': request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION,
|
||||
'QUERY_GENERATION_PROMPT_TEMPLATE': request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE,
|
||||
'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE,
|
||||
'ENABLE_VOICE_MODE_PROMPT': request.app.state.config.ENABLE_VOICE_MODE_PROMPT,
|
||||
'VOICE_MODE_PROMPT_TEMPLATE': request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE,
|
||||
}
|
||||
return await get_config_values(TASK_CONFIG_KEYS)
|
||||
|
||||
|
||||
class TaskConfigForm(BaseModel):
|
||||
@@ -88,6 +101,7 @@ class TaskConfigForm(BaseModel):
|
||||
IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE: str
|
||||
ENABLE_AUTOCOMPLETE_GENERATION: bool
|
||||
AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH: int
|
||||
AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE: str
|
||||
TAGS_GENERATION_PROMPT_TEMPLATE: str
|
||||
FOLLOW_UP_GENERATION_PROMPT_TEMPLATE: str
|
||||
ENABLE_FOLLOW_UP_GENERATION: bool
|
||||
@@ -102,56 +116,13 @@ class TaskConfigForm(BaseModel):
|
||||
|
||||
@router.post('/config/update')
|
||||
async def update_task_config(request: Request, form_data: TaskConfigForm, user=Depends(get_admin_user)):
|
||||
request.app.state.config.TASK_MODEL = form_data.TASK_MODEL
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL = form_data.TASK_MODEL_EXTERNAL
|
||||
request.app.state.config.ENABLE_TITLE_GENERATION = form_data.ENABLE_TITLE_GENERATION
|
||||
request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE = form_data.TITLE_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.ENABLE_FOLLOW_UP_GENERATION = form_data.ENABLE_FOLLOW_UP_GENERATION
|
||||
request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE = form_data.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE = form_data.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION = form_data.ENABLE_AUTOCOMPLETE_GENERATION
|
||||
request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH = (
|
||||
form_data.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH
|
||||
)
|
||||
|
||||
request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE = form_data.TAGS_GENERATION_PROMPT_TEMPLATE
|
||||
request.app.state.config.ENABLE_TAGS_GENERATION = form_data.ENABLE_TAGS_GENERATION
|
||||
request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION = form_data.ENABLE_SEARCH_QUERY_GENERATION
|
||||
request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION = form_data.ENABLE_RETRIEVAL_QUERY_GENERATION
|
||||
|
||||
request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE = form_data.QUERY_GENERATION_PROMPT_TEMPLATE
|
||||
request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE = form_data.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE
|
||||
|
||||
request.app.state.config.ENABLE_VOICE_MODE_PROMPT = form_data.ENABLE_VOICE_MODE_PROMPT
|
||||
request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE = form_data.VOICE_MODE_PROMPT_TEMPLATE
|
||||
|
||||
return {
|
||||
'TASK_MODEL': request.app.state.config.TASK_MODEL,
|
||||
'TASK_MODEL_EXTERNAL': request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
'ENABLE_TITLE_GENERATION': request.app.state.config.ENABLE_TITLE_GENERATION,
|
||||
'TITLE_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE,
|
||||
'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_AUTOCOMPLETE_GENERATION': request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION,
|
||||
'AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH': request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH,
|
||||
'TAGS_GENERATION_PROMPT_TEMPLATE': request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_TAGS_GENERATION': request.app.state.config.ENABLE_TAGS_GENERATION,
|
||||
'ENABLE_FOLLOW_UP_GENERATION': request.app.state.config.ENABLE_FOLLOW_UP_GENERATION,
|
||||
'FOLLOW_UP_GENERATION_PROMPT_TEMPLATE': request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE,
|
||||
'ENABLE_SEARCH_QUERY_GENERATION': request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION,
|
||||
'ENABLE_RETRIEVAL_QUERY_GENERATION': request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION,
|
||||
'QUERY_GENERATION_PROMPT_TEMPLATE': request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE,
|
||||
'TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE': request.app.state.config.TOOLS_FUNCTION_CALLING_PROMPT_TEMPLATE,
|
||||
'ENABLE_VOICE_MODE_PROMPT': request.app.state.config.ENABLE_VOICE_MODE_PROMPT,
|
||||
'VOICE_MODE_PROMPT_TEMPLATE': request.app.state.config.VOICE_MODE_PROMPT_TEMPLATE,
|
||||
}
|
||||
await Config.upsert(config_updates(form_data.model_dump(), TASK_CONFIG_KEYS))
|
||||
return await get_config_values(TASK_CONFIG_KEYS)
|
||||
|
||||
|
||||
@router.post('/title/completions')
|
||||
async def generate_title(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_TITLE_GENERATION:
|
||||
if not await Config.get('task.title.enable'):
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={'detail': 'Title generation is disabled'},
|
||||
@@ -181,15 +152,16 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating chat title using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.TITLE_GENERATION_PROMPT_TEMPLATE
|
||||
title_template = await Config.get('task.title.prompt_template')
|
||||
if title_template != '':
|
||||
template = title_template
|
||||
else:
|
||||
template = DEFAULT_TITLE_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -234,7 +206,7 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver
|
||||
|
||||
@router.post('/follow_up/completions')
|
||||
async def generate_follow_ups(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_FOLLOW_UP_GENERATION:
|
||||
if not await Config.get('task.follow_up.enable'):
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={'detail': 'Follow-up generation is disabled'},
|
||||
@@ -259,15 +231,16 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating chat title using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.FOLLOW_UP_GENERATION_PROMPT_TEMPLATE
|
||||
follow_up_template = await Config.get('task.follow_up.prompt_template')
|
||||
if follow_up_template != '':
|
||||
template = follow_up_template
|
||||
else:
|
||||
template = DEFAULT_FOLLOW_UP_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -303,7 +276,7 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge
|
||||
|
||||
@router.post('/tags/completions')
|
||||
async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_TAGS_GENERATION:
|
||||
if not await Config.get('task.tags.enable'):
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_200_OK,
|
||||
content={'detail': 'Tags generation is disabled'},
|
||||
@@ -328,15 +301,16 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating chat tags using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.TAGS_GENERATION_PROMPT_TEMPLATE
|
||||
tags_template = await Config.get('task.tags.prompt_template')
|
||||
if tags_template != '':
|
||||
template = tags_template
|
||||
else:
|
||||
template = DEFAULT_TAGS_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -391,15 +365,16 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends(
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating image prompt using model {task_model_id} for user {user.email} ')
|
||||
|
||||
if request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE != '':
|
||||
template = request.app.state.config.IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE
|
||||
image_prompt_template = await Config.get('task.image.prompt_template')
|
||||
if image_prompt_template != '':
|
||||
template = image_prompt_template
|
||||
else:
|
||||
template = DEFAULT_IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -437,13 +412,13 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends(
|
||||
async def generate_queries(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
type = form_data.get('type')
|
||||
if type == 'web_search':
|
||||
if not request.app.state.config.ENABLE_SEARCH_QUERY_GENERATION:
|
||||
if not await Config.get('task.query.search.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Search query generation'),
|
||||
)
|
||||
elif type == 'retrieval':
|
||||
if not request.app.state.config.ENABLE_RETRIEVAL_QUERY_GENERATION:
|
||||
if not await Config.get('task.query.retrieval.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Query generation'),
|
||||
@@ -472,15 +447,16 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating {type} queries using model {task_model_id} for user {user.email}')
|
||||
|
||||
if (request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE).strip() != '':
|
||||
template = request.app.state.config.QUERY_GENERATION_PROMPT_TEMPLATE
|
||||
query_template = await Config.get('task.query.prompt_template')
|
||||
if query_template.strip() != '':
|
||||
template = query_template
|
||||
else:
|
||||
template = DEFAULT_QUERY_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -515,7 +491,7 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v
|
||||
|
||||
@router.post('/auto/completions')
|
||||
async def generate_autocompletion(request: Request, form_data: dict, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_AUTOCOMPLETE_GENERATION:
|
||||
if not await Config.get('task.autocomplete.enable'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Autocompletion generation'),
|
||||
@@ -525,11 +501,12 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
|
||||
prompt = form_data.get('prompt')
|
||||
messages = form_data.get('messages')
|
||||
|
||||
if request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH > 0:
|
||||
if len(prompt) > request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH:
|
||||
autocomplete_input_max_length = await Config.get('task.autocomplete.input_max_length')
|
||||
if autocomplete_input_max_length > 0:
|
||||
if len(prompt) > autocomplete_input_max_length:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.INPUT_TOO_LONG(request.app.state.config.AUTOCOMPLETE_GENERATION_INPUT_MAX_LENGTH),
|
||||
detail=ERROR_MESSAGES.INPUT_TOO_LONG(autocomplete_input_max_length),
|
||||
)
|
||||
|
||||
if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'):
|
||||
@@ -551,15 +528,16 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
log.debug(f'generating autocompletion using model {task_model_id} for user {user.email}')
|
||||
|
||||
if (request.app.state.config.AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE).strip() != '':
|
||||
template = request.app.state.config.AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE
|
||||
autocomplete_template = await Config.get('task.autocomplete.prompt_template')
|
||||
if autocomplete_template.strip() != '':
|
||||
template = autocomplete_template
|
||||
else:
|
||||
template = DEFAULT_AUTOCOMPLETE_GENERATION_PROMPT_TEMPLATE
|
||||
|
||||
@@ -614,8 +592,8 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver
|
||||
# If the user has a custom task model, use that model
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
request.app.state.config.TASK_MODEL,
|
||||
request.app.state.config.TASK_MODEL_EXTERNAL,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
|
||||
|
||||
@@ -13,11 +13,14 @@ import aiohttp
|
||||
from fastapi import APIRouter, Depends, Request, Response, WebSocket
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from open_webui.config import TERMINAL_PROXY_HEADERS
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.access_control import has_connection_access
|
||||
from open_webui.utils.auth import get_verified_user
|
||||
from open_webui.utils.tools import bearer_auth_header, normalize_bearer_token
|
||||
from starlette.background import BackgroundTask
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -43,6 +46,9 @@ def _sanitize_proxy_path(path: str) -> str | None:
|
||||
if once == decoded:
|
||||
break
|
||||
decoded = once
|
||||
# Fail closed: still encoded after the cap means the upstream would decode further into traversal.
|
||||
if unquote(decoded) != decoded:
|
||||
return None
|
||||
had_trailing_slash = decoded.endswith('/')
|
||||
normalized = posixpath.normpath(decoded)
|
||||
# Remove any leading slashes that would reset the base
|
||||
@@ -59,7 +65,7 @@ def _sanitize_proxy_path(path: str) -> str | None:
|
||||
@router.get('/')
|
||||
async def list_terminal_servers(request: Request, user=Depends(get_verified_user)):
|
||||
"""Return terminal servers the authenticated user has access to."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
|
||||
|
||||
return [
|
||||
@@ -84,7 +90,7 @@ async def proxy_terminal(
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
"""Proxy a request to the admin terminal server identified by *server_id*."""
|
||||
connections = request.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
connection = next((c for c in connections if c.get('id') == server_id), None)
|
||||
|
||||
if connection is None:
|
||||
@@ -121,15 +127,15 @@ async def proxy_terminal(
|
||||
auth_type = connection.get('auth_type', 'bearer')
|
||||
|
||||
if auth_type == 'bearer':
|
||||
headers['Authorization'] = f'Bearer {connection.get("key", "")}'
|
||||
headers.update(bearer_auth_header(connection.get('key', '')))
|
||||
elif auth_type == 'session':
|
||||
cookies = request.cookies
|
||||
headers['Authorization'] = f'Bearer {request.state.token.credentials}'
|
||||
headers.update(bearer_auth_header(request.state.token.credentials))
|
||||
elif auth_type == 'system_oauth':
|
||||
cookies = request.cookies
|
||||
oauth_token = request.headers.get('x-oauth-access-token', '')
|
||||
if oauth_token:
|
||||
headers['Authorization'] = f'Bearer {oauth_token}'
|
||||
headers.update(bearer_auth_header(oauth_token))
|
||||
# auth_type == "none": no Authorization header
|
||||
|
||||
content_type = request.headers.get('content-type')
|
||||
@@ -206,7 +212,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str):
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
from open_webui.utils.auth import decode_token
|
||||
from open_webui.utils.auth import decode_token, is_valid_token
|
||||
|
||||
# First-message authentication
|
||||
try:
|
||||
@@ -217,7 +223,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str):
|
||||
return None
|
||||
token = payload.get('token', '')
|
||||
data = decode_token(token)
|
||||
if data is None or 'id' not in data:
|
||||
if data is None or 'id' not in data or not await is_valid_token(data, getattr(ws.app.state, 'redis', None)):
|
||||
await ws.close(code=4001, reason='Invalid token')
|
||||
return None
|
||||
user = await Users.get_user_by_id(data['id'])
|
||||
@@ -232,7 +238,7 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str):
|
||||
return None
|
||||
|
||||
# Resolve terminal server
|
||||
connections = ws.app.state.config.TERMINAL_SERVER_CONNECTIONS or []
|
||||
connections = await Config.get('terminal_server.connections', []) or []
|
||||
connection = next((c for c in connections if c.get('id') == server_id), None)
|
||||
|
||||
if connection is None:
|
||||
@@ -282,13 +288,19 @@ async def ws_terminal(
|
||||
|
||||
import urllib.parse
|
||||
|
||||
# Encode session_id as an opaque path segment so it cannot smuggle '?'/'#'/'&' (at any
|
||||
# decode depth) and inject an attacker-chosen user_id ahead of the one appended below.
|
||||
safe_session_id = urllib.parse.quote(session_id, safe='')
|
||||
|
||||
if policy_id:
|
||||
upstream_url = f'{ws_base}/p/{policy_id}/api/terminals/{session_id}'
|
||||
upstream_url = f'{ws_base}/p/{policy_id}/api/terminals/{safe_session_id}'
|
||||
else:
|
||||
upstream_url = f'{ws_base}/api/terminals/{session_id}'
|
||||
upstream_url = f'{ws_base}/api/terminals/{safe_session_id}'
|
||||
if upstream_params:
|
||||
upstream_url += f'?{urllib.parse.urlencode(upstream_params)}'
|
||||
|
||||
app = ws.scope.get('app')
|
||||
opened = False
|
||||
session = aiohttp.ClientSession()
|
||||
try:
|
||||
async with session.ws_connect(upstream_url, ssl=AIOHTTP_CLIENT_SESSION_SSL) as upstream:
|
||||
@@ -298,9 +310,19 @@ async def ws_terminal(
|
||||
# First-message auth to upstream terminal server
|
||||
auth_type = connection.get('auth_type', 'bearer')
|
||||
if auth_type == 'bearer':
|
||||
key = connection.get('key', '')
|
||||
key = normalize_bearer_token(connection.get('key', ''))
|
||||
await upstream.send_str(_json.dumps({'type': 'auth', 'token': key}))
|
||||
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.TERMINAL_SESSION_OPENED,
|
||||
actor=user,
|
||||
subject_id=session_id,
|
||||
subject_type='terminal.session',
|
||||
data={'server_id': server_id},
|
||||
)
|
||||
opened = True
|
||||
|
||||
async def _client_to_upstream():
|
||||
"""Forward client → upstream."""
|
||||
try:
|
||||
@@ -349,6 +371,15 @@ async def ws_terminal(
|
||||
log.exception('Terminal WebSocket proxy error: %s', e)
|
||||
finally:
|
||||
await session.close()
|
||||
if opened:
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.TERMINAL_SESSION_CLOSED,
|
||||
actor=user,
|
||||
subject_id=session_id,
|
||||
subject_type='terminal.session',
|
||||
data={'server_id': server_id},
|
||||
)
|
||||
try:
|
||||
await ws.close()
|
||||
except Exception:
|
||||
|
||||
@@ -11,8 +11,10 @@ from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, CACHE_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.models.tools import (
|
||||
@@ -30,6 +32,7 @@ from open_webui.utils.access_control import (
|
||||
)
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.plugin import (
|
||||
get_tools_cache,
|
||||
get_tool_module_from_cache,
|
||||
load_tool_module_by_id,
|
||||
replace_imports,
|
||||
@@ -69,13 +72,17 @@ async def get_tools(
|
||||
tools = []
|
||||
|
||||
# Local Tools
|
||||
tools_cache = get_tools_cache(request)
|
||||
for tool in await Tools.get_tools(defer_content=True, db=db):
|
||||
tool_module = request.app.state.TOOLS.get(tool.id) if hasattr(request.app.state, 'TOOLS') else None
|
||||
tool_module = tools_cache.get(tool.id)
|
||||
has_user_valves = (
|
||||
hasattr(tool_module, 'UserValves') if tool_module else (tool.meta.has_user_valves if tool.meta else False)
|
||||
)
|
||||
tools.append(
|
||||
ToolUserResponse(
|
||||
**{
|
||||
**tool.model_dump(),
|
||||
'has_user_valves': (hasattr(tool_module, 'UserValves') if tool_module else False),
|
||||
'has_user_valves': has_user_valves,
|
||||
}
|
||||
)
|
||||
)
|
||||
@@ -84,7 +91,7 @@ async def get_tools(
|
||||
server_access_grants = {}
|
||||
for server in await get_tool_servers(request):
|
||||
server_idx = server.get('idx', 0)
|
||||
connections = request.app.state.config.TOOL_SERVER_CONNECTIONS
|
||||
connections = await Config.get('tool_server.connections', [])
|
||||
if server_idx >= len(connections):
|
||||
log.warning(
|
||||
f'Tool server index {server_idx} out of range '
|
||||
@@ -113,13 +120,14 @@ async def get_tools(
|
||||
)
|
||||
|
||||
# MCP Tool Servers
|
||||
for server in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
if server.get('type', 'openapi') == 'mcp' and server.get('config', {}).get('enable'):
|
||||
server_id = server.get('info', {}).get('id')
|
||||
for server in await Config.get('tool_server.connections', []):
|
||||
if server.get('type', 'openapi') == 'mcp' and (server.get('config') or {}).get('enable'):
|
||||
info = server.get('info') or {}
|
||||
server_id = info.get('id')
|
||||
auth_type = server.get('auth_type', 'none')
|
||||
|
||||
session_token = None
|
||||
if auth_type in ('oauth_2.1', 'oauth_2.1_static'):
|
||||
if auth_type in ('oauth_2.1', 'oauth_2.1_static') and server_id:
|
||||
splits = server_id.split(':')
|
||||
server_id = splits[-1] if len(splits) > 1 else server_id
|
||||
|
||||
@@ -127,9 +135,9 @@ async def get_tools(
|
||||
user.id, f'mcp:{server_id}'
|
||||
)
|
||||
|
||||
server_config = server.get('config', {})
|
||||
server_config = server.get('config') or {}
|
||||
|
||||
tool_id = f'server:mcp:{server.get("info", {}).get("id")}'
|
||||
tool_id = f'server:mcp:{info.get("id")}'
|
||||
server_access_grants[tool_id] = server_config.get('access_grants', [])
|
||||
|
||||
tools.append(
|
||||
@@ -137,9 +145,9 @@ async def get_tools(
|
||||
**{
|
||||
'id': tool_id,
|
||||
'user_id': tool_id,
|
||||
'name': server.get('info', {}).get('name', 'MCP Tool Server'),
|
||||
'name': info.get('name', 'MCP Tool Server'),
|
||||
'meta': {
|
||||
'description': server.get('info', {}).get('description', ''),
|
||||
'description': info.get('description', ''),
|
||||
},
|
||||
'updated_at': int(time.time()),
|
||||
'created_at': int(time.time()),
|
||||
@@ -285,8 +293,13 @@ async def load_tool_from_url(request: Request, form_data: LoadUrlForm, user=Depe
|
||||
'name': tool_name,
|
||||
'content': data,
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=ERROR_MESSAGES.DEFAULT(e))
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error fetching tool'),
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
@@ -303,7 +316,7 @@ async def export_tools(
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id,
|
||||
'workspace.tools_export',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -331,11 +344,11 @@ async def create_new_tools(
|
||||
):
|
||||
"""Create a new tool from user-supplied Python source code."""
|
||||
if user.role != 'admin' and not (
|
||||
await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
await has_permission(user.id, 'workspace.tools', await Config.get('user.permissions'), db=db)
|
||||
or await has_permission(
|
||||
user.id,
|
||||
'workspace.tools_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -356,7 +369,7 @@ async def create_new_tools(
|
||||
if tools is None:
|
||||
try:
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -366,8 +379,9 @@ async def create_new_tools(
|
||||
form_data.content = replace_imports(form_data.content)
|
||||
tool_module, frontmatter = await load_tool_module_by_id(form_data.id, content=form_data.content)
|
||||
form_data.meta.manifest = frontmatter
|
||||
form_data.meta.has_user_valves = hasattr(tool_module, 'UserValves')
|
||||
|
||||
TOOLS = request.app.state.TOOLS
|
||||
TOOLS = get_tools_cache(request)
|
||||
TOOLS[form_data.id] = tool_module
|
||||
|
||||
specs = get_tool_specs(TOOLS[form_data.id])
|
||||
@@ -377,17 +391,26 @@ async def create_new_tools(
|
||||
tool_cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if tools:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.TOOL_CREATED,
|
||||
actor=user,
|
||||
subject_id=tools.id,
|
||||
data={'name': tools.name},
|
||||
)
|
||||
return tools
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error creating tools'),
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to load the tool by id {form_data.id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error creating tool'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -484,8 +507,8 @@ async def update_tools_by_id(
|
||||
# Content edits trigger exec on load — gate them behind workspace.tools (matches /create).
|
||||
if form_data.content != tools.content:
|
||||
if user.role != 'admin' and not (
|
||||
await has_permission(user.id, 'workspace.tools', request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
or await has_permission(user.id, 'workspace.tools_import', request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
await has_permission(user.id, 'workspace.tools', await Config.get('user.permissions'), db=db)
|
||||
or await has_permission(user.id, 'workspace.tools_import', await Config.get('user.permissions'), db=db)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -496,14 +519,15 @@ async def update_tools_by_id(
|
||||
form_data.content = replace_imports(form_data.content)
|
||||
tool_module, frontmatter = await load_tool_module_by_id(id, content=form_data.content)
|
||||
form_data.meta.manifest = frontmatter
|
||||
form_data.meta.has_user_valves = hasattr(tool_module, 'UserValves')
|
||||
|
||||
TOOLS = request.app.state.TOOLS
|
||||
TOOLS = get_tools_cache(request)
|
||||
TOOLS[id] = tool_module
|
||||
|
||||
specs = get_tool_specs(TOOLS[id])
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -519,6 +543,13 @@ async def update_tools_by_id(
|
||||
tools = await Tools.update_tool_by_id(id, updated, db=db)
|
||||
|
||||
if tools:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.TOOL_UPDATED,
|
||||
actor=user,
|
||||
subject_id=tools.id,
|
||||
data={'name': tools.name},
|
||||
)
|
||||
return tools
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -526,10 +557,12 @@ async def update_tools_by_id(
|
||||
detail=ERROR_MESSAGES.DEFAULT('Error updating tools'),
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating tool'),
|
||||
)
|
||||
|
||||
|
||||
@@ -574,7 +607,7 @@ async def update_tool_access_by_id(
|
||||
)
|
||||
|
||||
form_data.access_grants = await filter_allowed_access_grants(
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
@@ -583,7 +616,15 @@ async def update_tool_access_by_id(
|
||||
|
||||
await AccessGrants.set_access_grants('tool', id, form_data.access_grants, db=db)
|
||||
|
||||
return await Tools.get_tool_by_id(id, db=db)
|
||||
tools = await Tools.get_tool_by_id(id, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.TOOL_ACCESS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': tools.name if tools else None},
|
||||
)
|
||||
return tools
|
||||
|
||||
|
||||
############################
|
||||
@@ -623,9 +664,15 @@ async def delete_tools_by_id(
|
||||
|
||||
result = await Tools.delete_tool_by_id(id, db=db)
|
||||
if result:
|
||||
TOOLS = request.app.state.TOOLS
|
||||
if id in TOOLS:
|
||||
del TOOLS[id]
|
||||
TOOLS = get_tools_cache(request)
|
||||
TOOLS.pop(id, None)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.TOOL_DELETED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'name': tools.name},
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
@@ -668,7 +715,7 @@ async def get_tools_valves_by_id(
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error getting tool valves'),
|
||||
)
|
||||
|
||||
|
||||
@@ -707,11 +754,7 @@ async def get_tools_valves_spec_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if id in request.app.state.TOOLS:
|
||||
tools_module = request.app.state.TOOLS[id]
|
||||
else:
|
||||
tools_module, _ = await load_tool_module_by_id(id)
|
||||
request.app.state.TOOLS[id] = tools_module
|
||||
tools_module, _ = await get_tool_module_from_cache(request, id)
|
||||
|
||||
if hasattr(tools_module, 'Valves'):
|
||||
Valves = tools_module.Valves
|
||||
@@ -758,11 +801,7 @@ async def update_tools_valves_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if id in request.app.state.TOOLS:
|
||||
tools_module = request.app.state.TOOLS[id]
|
||||
else:
|
||||
tools_module, _ = await load_tool_module_by_id(id)
|
||||
request.app.state.TOOLS[id] = tools_module
|
||||
tools_module, _ = await get_tool_module_from_cache(request, id)
|
||||
|
||||
if not hasattr(tools_module, 'Valves'):
|
||||
raise HTTPException(
|
||||
@@ -776,12 +815,18 @@ async def update_tools_valves_by_id(
|
||||
valves = Valves(**form_data)
|
||||
valves_dict = valves.model_dump(exclude_unset=True)
|
||||
await Tools.update_tool_valves_by_id(id, valves_dict, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.TOOL_VALVES_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
)
|
||||
return valves_dict
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to update tool valves by id {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating tool valves'),
|
||||
)
|
||||
|
||||
|
||||
@@ -823,7 +868,7 @@ async def get_tools_user_valves_by_id(
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error getting tool user valves'),
|
||||
)
|
||||
|
||||
|
||||
@@ -857,11 +902,7 @@ async def get_tools_user_valves_spec_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if id in request.app.state.TOOLS:
|
||||
tools_module = request.app.state.TOOLS[id]
|
||||
else:
|
||||
tools_module, _ = await load_tool_module_by_id(id)
|
||||
request.app.state.TOOLS[id] = tools_module
|
||||
tools_module, _ = await get_tool_module_from_cache(request, id)
|
||||
|
||||
if hasattr(tools_module, 'UserValves'):
|
||||
UserValves = tools_module.UserValves
|
||||
@@ -903,11 +944,7 @@ async def update_tools_user_valves_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
if id in request.app.state.TOOLS:
|
||||
tools_module = request.app.state.TOOLS[id]
|
||||
else:
|
||||
tools_module, _ = await load_tool_module_by_id(id)
|
||||
request.app.state.TOOLS[id] = tools_module
|
||||
tools_module, _ = await get_tool_module_from_cache(request, id)
|
||||
|
||||
if hasattr(tools_module, 'UserValves'):
|
||||
UserValves = tools_module.UserValves
|
||||
@@ -917,12 +954,19 @@ async def update_tools_user_valves_by_id(
|
||||
user_valves = UserValves(**form_data)
|
||||
user_valves_dict = user_valves.model_dump(exclude_unset=True)
|
||||
await Tools.update_user_valves_by_id_and_user_id(id, user.id, user_valves_dict, db=db)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.TOOL_VALVES_UPDATED,
|
||||
actor=user,
|
||||
subject_id=id,
|
||||
data={'scope': 'user'},
|
||||
)
|
||||
return user_valves_dict
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to update user valves by id {id}: {e}')
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=ERROR_MESSAGES.DEFAULT(str(e)),
|
||||
detail=ERROR_MESSAGES.DEFAULT(e, 'Error updating tool user valves'),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
|
||||
@@ -9,9 +9,11 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import FileResponse, Response, StreamingResponse
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.env import ENABLE_PROFILE_IMAGE_URL_FORWARDING, PROFILE_IMAGE_ALLOWED_MIME_TYPES, STATIC_DIR
|
||||
from open_webui.internal.db import get_async_session
|
||||
from open_webui.models.auths import Auths
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
from open_webui.models.users import (
|
||||
@@ -38,7 +40,7 @@ from open_webui.utils.auth import (
|
||||
get_verified_user,
|
||||
validate_password,
|
||||
)
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -157,7 +159,7 @@ async def get_user_permissisions(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
user_permissions = await get_permissions(user.id, request.app.state.config.USER_PERMISSIONS, db=db)
|
||||
user_permissions = await get_permissions(user.id, await Config.get('user.permissions'), db=db)
|
||||
|
||||
return user_permissions
|
||||
|
||||
@@ -177,6 +179,8 @@ class WorkspacePermissions(BaseModel):
|
||||
prompts_export: bool = False
|
||||
tools_import: bool = False
|
||||
tools_export: bool = False
|
||||
skills_import: bool = False
|
||||
skills_export: bool = False
|
||||
|
||||
|
||||
class SharingPermissions(BaseModel):
|
||||
@@ -201,6 +205,8 @@ class AccessGrantsPermissions(BaseModel):
|
||||
|
||||
|
||||
class ChatPermissions(BaseModel):
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
controls: bool = True
|
||||
valves: bool = True
|
||||
system_prompt: bool = True
|
||||
@@ -215,6 +221,7 @@ class ChatPermissions(BaseModel):
|
||||
edit: bool = True
|
||||
share: bool = True
|
||||
export: bool = True
|
||||
import_: bool = Field(default=True, alias='import')
|
||||
stt: bool = True
|
||||
tts: bool = True
|
||||
call: bool = True
|
||||
@@ -236,6 +243,7 @@ class FeaturesPermissions(BaseModel):
|
||||
memories: bool = True
|
||||
automations: bool = False
|
||||
calendar: bool = True
|
||||
webhooks: bool = False
|
||||
|
||||
|
||||
class SettingsPermissions(BaseModel):
|
||||
@@ -253,20 +261,43 @@ class UserPermissions(BaseModel):
|
||||
|
||||
@router.get('/default/permissions', response_model=UserPermissions)
|
||||
async def get_default_user_permissions(request: Request, user=Depends(get_admin_user)):
|
||||
user_permissions = await Config.get('user.permissions')
|
||||
return {
|
||||
'workspace': WorkspacePermissions(**request.app.state.config.USER_PERMISSIONS.get('workspace', {})),
|
||||
'sharing': SharingPermissions(**request.app.state.config.USER_PERMISSIONS.get('sharing', {})),
|
||||
'access_grants': AccessGrantsPermissions(**request.app.state.config.USER_PERMISSIONS.get('access_grants', {})),
|
||||
'chat': ChatPermissions(**request.app.state.config.USER_PERMISSIONS.get('chat', {})),
|
||||
'features': FeaturesPermissions(**request.app.state.config.USER_PERMISSIONS.get('features', {})),
|
||||
'settings': SettingsPermissions(**request.app.state.config.USER_PERMISSIONS.get('settings', {})),
|
||||
'workspace': WorkspacePermissions(**user_permissions.get('workspace', {})),
|
||||
'sharing': SharingPermissions(**user_permissions.get('sharing', {})),
|
||||
'access_grants': AccessGrantsPermissions(**user_permissions.get('access_grants', {})),
|
||||
'chat': ChatPermissions(**user_permissions.get('chat', {})),
|
||||
'features': FeaturesPermissions(**user_permissions.get('features', {})),
|
||||
'settings': SettingsPermissions(**user_permissions.get('settings', {})),
|
||||
}
|
||||
|
||||
|
||||
@router.post('/default/permissions')
|
||||
async def update_default_user_permissions(request: Request, form_data: UserPermissions, user=Depends(get_admin_user)):
|
||||
request.app.state.config.USER_PERMISSIONS = form_data.model_dump()
|
||||
return request.app.state.config.USER_PERMISSIONS
|
||||
user_permissions = form_data.model_dump(by_alias=True)
|
||||
await Config.upsert({'user.permissions': user_permissions})
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_PERMISSIONS_UPDATED,
|
||||
actor=user,
|
||||
subject_id='user.permissions',
|
||||
subject_type='config',
|
||||
)
|
||||
return user_permissions
|
||||
|
||||
|
||||
@router.get('/default/permissions/defaults', response_model=UserPermissions)
|
||||
async def get_default_user_permissions_defaults(user=Depends(get_admin_user)):
|
||||
from open_webui.config import DEFAULT_USER_PERMISSIONS
|
||||
|
||||
return {
|
||||
'workspace': WorkspacePermissions(**DEFAULT_USER_PERMISSIONS.get('workspace', {})),
|
||||
'sharing': SharingPermissions(**DEFAULT_USER_PERMISSIONS.get('sharing', {})),
|
||||
'access_grants': AccessGrantsPermissions(**DEFAULT_USER_PERMISSIONS.get('access_grants', {})),
|
||||
'chat': ChatPermissions(**DEFAULT_USER_PERMISSIONS.get('chat', {})),
|
||||
'features': FeaturesPermissions(**DEFAULT_USER_PERMISSIONS.get('features', {})),
|
||||
'settings': SettingsPermissions(**DEFAULT_USER_PERMISSIONS.get('settings', {})),
|
||||
}
|
||||
|
||||
|
||||
############################
|
||||
@@ -294,6 +325,14 @@ async def update_user_settings_by_session_user(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if user.role != 'admin' and not await has_permission(
|
||||
user.id, 'settings.interface', request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
updated_user_settings = form_data.model_dump()
|
||||
ui_settings = updated_user_settings.get('ui')
|
||||
if (
|
||||
@@ -303,7 +342,7 @@ async def update_user_settings_by_session_user(
|
||||
and not await has_permission(
|
||||
user.id,
|
||||
'features.direct_tool_servers',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
await Config.get('user.permissions'),
|
||||
)
|
||||
):
|
||||
# If the user is not an admin and does not have permission to use tool servers, remove the key
|
||||
@@ -311,6 +350,12 @@ async def update_user_settings_by_session_user(
|
||||
|
||||
user = await Users.update_user_settings_by_id(user.id, updated_user_settings, db=db)
|
||||
if user:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_SETTINGS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
)
|
||||
return user.settings
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -330,7 +375,7 @@ async def get_user_status_by_session_user(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_USER_STATUS:
|
||||
if not await Config.get('users.enable_status'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
|
||||
@@ -351,7 +396,7 @@ async def update_user_status_by_session_user(
|
||||
user=Depends(get_verified_user),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
if not request.app.state.config.ENABLE_USER_STATUS:
|
||||
if not await Config.get('users.enable_status'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACTION_PROHIBITED,
|
||||
@@ -359,6 +404,12 @@ async def update_user_status_by_session_user(
|
||||
# user already fetched by get_verified_user — no need to refetch
|
||||
updated = await Users.update_user_status_by_id(user.id, form_data, db=db)
|
||||
if updated:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_STATUS_UPDATED,
|
||||
actor=user,
|
||||
subject_id=user.id,
|
||||
)
|
||||
return updated
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
@@ -541,6 +592,7 @@ async def get_user_active_status_by_id(
|
||||
|
||||
@router.post('/{user_id}/update', response_model=UserModel | None)
|
||||
async def update_user_by_id(
|
||||
request: Request,
|
||||
user_id: str,
|
||||
form_data: UserUpdateForm,
|
||||
session_user: UserModel = Depends(get_admin_user),
|
||||
@@ -591,7 +643,7 @@ async def update_user_by_id(
|
||||
except Exception as e:
|
||||
raise HTTPException(400, detail=str(e))
|
||||
|
||||
hashed = get_password_hash(form_data.password)
|
||||
hashed = await get_password_hash(form_data.password)
|
||||
await Auths.update_user_password_by_id(user_id, hashed, db=db)
|
||||
|
||||
# Build update dict from only the provided fields
|
||||
@@ -620,6 +672,30 @@ async def update_user_by_id(
|
||||
# privileges cached in SESSION_POOL are invalidated.
|
||||
if updated_user.role != user.role:
|
||||
await disconnect_user_sessions(user_id)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_ROLE_UPDATED,
|
||||
actor=session_user,
|
||||
subject_id=user_id,
|
||||
data={'role': updated_user.role},
|
||||
)
|
||||
else:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_UPDATED,
|
||||
actor=session_user,
|
||||
subject_id=user_id,
|
||||
data={'updated_fields': list(update_data.keys())},
|
||||
)
|
||||
if form_data.password:
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.AUTH_PASSWORD_CHANGED,
|
||||
actor=session_user,
|
||||
subject_id=user_id,
|
||||
subject_type='user',
|
||||
source='admin',
|
||||
)
|
||||
return updated_user
|
||||
|
||||
raise HTTPException(
|
||||
@@ -639,7 +715,9 @@ async def update_user_by_id(
|
||||
|
||||
|
||||
@router.delete('/{user_id}', response_model=bool)
|
||||
async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)):
|
||||
async def delete_user_by_id(
|
||||
request: Request, user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)
|
||||
):
|
||||
# Prevent deletion of the primary admin user
|
||||
try:
|
||||
first_user = await Users.get_first_user(db=db)
|
||||
@@ -662,6 +740,12 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user), db: Asyn
|
||||
|
||||
if result:
|
||||
await disconnect_user_sessions(user_id)
|
||||
await publish_event(
|
||||
request,
|
||||
EVENTS.USER_DELETED,
|
||||
actor=user,
|
||||
subject_id=user_id,
|
||||
)
|
||||
return True
|
||||
|
||||
raise HTTPException(
|
||||
|
||||
@@ -7,6 +7,7 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from open_webui.config import DATA_DIR, ENABLE_ADMIN_EXPORT
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.models.chats import ChatTitleMessagesForm
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.code_interpreter import execute_code_jupyter
|
||||
from open_webui.utils.misc import get_gravatar_url
|
||||
@@ -41,27 +42,27 @@ async def format_code(form_data: CodeForm, user=Depends(get_admin_user)):
|
||||
|
||||
@router.post('/code/execute')
|
||||
async def execute_code(request: Request, form_data: CodeForm, user=Depends(get_verified_user)):
|
||||
if not request.app.state.config.ENABLE_CODE_EXECUTION:
|
||||
if not await Config.get('code_execution.enable'):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=ERROR_MESSAGES.FEATURE_DISABLED('Code execution'),
|
||||
)
|
||||
|
||||
if request.app.state.config.CODE_EXECUTION_ENGINE == 'jupyter':
|
||||
if await Config.get('code_execution.engine') == 'jupyter':
|
||||
output = await execute_code_jupyter(
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_URL,
|
||||
await Config.get('code_execution.jupyter.url'),
|
||||
form_data.code,
|
||||
(
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_TOKEN
|
||||
if request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH == 'token'
|
||||
await Config.get('code_execution.jupyter.auth_token')
|
||||
if await Config.get('code_execution.jupyter.auth') == 'token'
|
||||
else None
|
||||
),
|
||||
(
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH_PASSWORD
|
||||
if request.app.state.config.CODE_EXECUTION_JUPYTER_AUTH == 'password'
|
||||
await Config.get('code_execution.jupyter.auth_password')
|
||||
if await Config.get('code_execution.jupyter.auth') == 'password'
|
||||
else None
|
||||
),
|
||||
request.app.state.config.CODE_EXECUTION_JUPYTER_TIMEOUT,
|
||||
await Config.get('code_execution.jupyter.timeout'),
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@@ -38,7 +38,7 @@ from open_webui.models.users import UserNameResponse, Users
|
||||
from open_webui.socket.utils import RedisDict, RedisLock, YdocManager
|
||||
from open_webui.tasks import create_task, stop_item_tasks
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import decode_token
|
||||
from open_webui.utils.auth import decode_token, is_valid_token
|
||||
from open_webui.utils.redis import (
|
||||
build_sentinel_url,
|
||||
get_redis_connection,
|
||||
@@ -342,9 +342,12 @@ async def usage(sid, data):
|
||||
async def connect(sid, environ, auth):
|
||||
user = None
|
||||
if auth and 'token' in auth:
|
||||
scope = (environ or {}).get('asgi.scope') or {}
|
||||
fastapi_app = scope.get('app')
|
||||
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
|
||||
data = decode_token(auth['token'])
|
||||
|
||||
if data is not None and 'id' in data:
|
||||
if data is not None and 'id' in data and await is_valid_token(data, redis):
|
||||
user = await Users.get_user_by_id(data['id'])
|
||||
|
||||
if user:
|
||||
@@ -369,8 +372,12 @@ async def user_join(sid, data):
|
||||
if not auth or 'token' not in auth:
|
||||
return
|
||||
|
||||
environ = sio.get_environ(sid) or {}
|
||||
scope = environ.get('asgi.scope') or {}
|
||||
fastapi_app = scope.get('app')
|
||||
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
|
||||
token_data = decode_token(auth['token'])
|
||||
if token_data is None or 'id' not in token_data:
|
||||
if token_data is None or 'id' not in token_data or not await is_valid_token(token_data, redis):
|
||||
return
|
||||
|
||||
user = await Users.get_user_by_id(token_data['id'])
|
||||
@@ -416,8 +423,12 @@ async def join_channel(sid, data):
|
||||
if not auth or 'token' not in auth:
|
||||
return
|
||||
|
||||
environ = sio.get_environ(sid) or {}
|
||||
scope = environ.get('asgi.scope') or {}
|
||||
fastapi_app = scope.get('app')
|
||||
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
|
||||
data = decode_token(auth['token'])
|
||||
if data is None or 'id' not in data:
|
||||
if data is None or 'id' not in data or not await is_valid_token(data, redis):
|
||||
return
|
||||
|
||||
user = await Users.get_user_by_id(data['id'])
|
||||
@@ -438,8 +449,12 @@ async def join_note(sid, data):
|
||||
if not auth or 'token' not in auth:
|
||||
return
|
||||
|
||||
environ = sio.get_environ(sid) or {}
|
||||
scope = environ.get('asgi.scope') or {}
|
||||
fastapi_app = scope.get('app')
|
||||
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
|
||||
token_data = decode_token(auth['token'])
|
||||
if token_data is None or 'id' not in token_data:
|
||||
if token_data is None or 'id' not in token_data or not await is_valid_token(token_data, redis):
|
||||
return
|
||||
|
||||
user = await Users.get_user_by_id(token_data['id'])
|
||||
@@ -757,11 +772,13 @@ async def yjs_document_update(sid, data):
|
||||
@sio.on('ydoc:document:leave')
|
||||
async def yjs_document_leave(sid, data):
|
||||
"""Handle user leaving a document"""
|
||||
user = SESSION_POOL.get(sid)
|
||||
if not user: # authenticated session required (parity with sibling handlers)
|
||||
return
|
||||
try:
|
||||
document_id = normalize_document_id(data['document_id'])
|
||||
user_id = data.get('user_id', sid)
|
||||
|
||||
log.info(f'User {user_id} leaving document {document_id}')
|
||||
log.info(f'User {user["id"]} leaving document {document_id}')
|
||||
|
||||
# Remove user from the document
|
||||
await YDOC_MANAGER.remove_user(document_id=document_id, user_id=sid)
|
||||
@@ -769,10 +786,10 @@ async def yjs_document_leave(sid, data):
|
||||
# Leave Socket.IO room
|
||||
await sio.leave_room(sid, f'doc_{document_id}')
|
||||
|
||||
# Notify other users
|
||||
# Notify other users; user_id is the authenticated identity, not client-supplied
|
||||
await sio.emit(
|
||||
'ydoc:user:left',
|
||||
{'document_id': document_id, 'user_id': user_id},
|
||||
{'document_id': document_id, 'user_id': user['id']},
|
||||
room=f'doc_{document_id}',
|
||||
)
|
||||
|
||||
@@ -787,16 +804,21 @@ async def yjs_document_leave(sid, data):
|
||||
@sio.on('ydoc:awareness:update')
|
||||
async def yjs_awareness_update(sid, data):
|
||||
"""Handle awareness updates (cursors, selections, etc.)"""
|
||||
user = SESSION_POOL.get(sid)
|
||||
if not user: # authenticated session required (parity with sibling handlers)
|
||||
return
|
||||
try:
|
||||
document_id = data['document_id']
|
||||
user_id = data.get('user_id', sid)
|
||||
document_id = normalize_document_id(data['document_id'])
|
||||
room = f'doc_{document_id}'
|
||||
if room not in sio.rooms(sid): # must have joined the document first
|
||||
return
|
||||
update = data['update']
|
||||
|
||||
# Broadcast awareness update to all other users in the document
|
||||
# Broadcast to the room; user_id is the authenticated identity, not client-supplied
|
||||
await sio.emit(
|
||||
'ydoc:awareness:update',
|
||||
{'document_id': document_id, 'user_id': user_id, 'update': update},
|
||||
room=f'doc_{document_id}',
|
||||
{'document_id': document_id, 'user_id': user['id'], 'update': update},
|
||||
room=room,
|
||||
skip_sid=sid,
|
||||
)
|
||||
|
||||
@@ -841,11 +863,14 @@ async def _make_channel_emitter(request_info):
|
||||
async def _emit_channel_update(content: str, done: bool = False):
|
||||
from open_webui.models.messages import MessageForm, Messages
|
||||
|
||||
msg = await Messages.get_message_by_id(message_id)
|
||||
if not msg or msg.channel_id != channel_id:
|
||||
return
|
||||
|
||||
update_form = MessageForm(content=content)
|
||||
if done:
|
||||
# Merge done flag into existing meta (preserve model_id etc.)
|
||||
msg = await Messages.get_message_by_id(message_id)
|
||||
existing_meta = (msg.meta or {}) if msg else {}
|
||||
existing_meta = msg.meta or {}
|
||||
update_form = MessageForm(
|
||||
content=content,
|
||||
meta={**existing_meta, 'done': True},
|
||||
@@ -1015,9 +1040,10 @@ async def get_event_call(request_info):
|
||||
async def __event_caller__(event_data):
|
||||
session_id = request_info['session_id']
|
||||
|
||||
# Fast-fail if the client has disconnected.
|
||||
if session_id not in SESSION_POOL:
|
||||
log.warning(f'Event caller: session {session_id} no longer connected')
|
||||
# session_id is client-supplied; only the requesting user's own live session may be targeted.
|
||||
session = SESSION_POOL.get(session_id)
|
||||
if session is None or session.get('id') != request_info.get('user_id'):
|
||||
log.warning(f'Event caller: session {session_id} not owned by requesting user or disconnected')
|
||||
return {'error': 'Client session disconnected.'}
|
||||
|
||||
try:
|
||||
|
||||
@@ -18,6 +18,7 @@ from fastapi import Request
|
||||
|
||||
from open_webui.models.channels import Channel, ChannelMember, Channels
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.memories import Memories
|
||||
from open_webui.models.messages import Message, Messages
|
||||
@@ -33,9 +34,15 @@ from open_webui.routers.images import (
|
||||
)
|
||||
from open_webui.routers.memories import (
|
||||
AddMemoryForm,
|
||||
ListMemoryPathsForm,
|
||||
MemoryUpdateModel,
|
||||
QueryMemoryForm,
|
||||
query_memory,
|
||||
ReadMemoryPathForm,
|
||||
SearchMemoriesForm,
|
||||
UpdateMemoriesForm,
|
||||
list_memory_paths as _list_memory_paths,
|
||||
read_memory_path as _read_memory_path,
|
||||
search_memories as _search_memories,
|
||||
update_memories as _update_memories,
|
||||
update_memory_by_id,
|
||||
)
|
||||
from open_webui.routers.memories import (
|
||||
@@ -225,10 +232,10 @@ async def search_web(
|
||||
return json.dumps({'error': 'Request context not available'})
|
||||
|
||||
try:
|
||||
engine = __request__.app.state.config.WEB_SEARCH_ENGINE
|
||||
engine = await Config.get('web.search.engine')
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
|
||||
configured = __request__.app.state.config.WEB_SEARCH_RESULT_COUNT
|
||||
configured = await Config.get('web.search.result_count')
|
||||
max_count = 5 if configured is None else configured
|
||||
count = max(1, min(count, max_count)) if count is not None else max_count
|
||||
|
||||
@@ -261,12 +268,12 @@ async def fetch_url(
|
||||
return json.dumps({'error': 'Request context not available'})
|
||||
|
||||
try:
|
||||
content, _ = await asyncio.to_thread(get_content_from_url, __request__, url)
|
||||
content, _ = await get_content_from_url(__request__, url)
|
||||
|
||||
# Truncate if configured (WEB_FETCH_MAX_CONTENT_LENGTH)
|
||||
# Guard: content may be None if the web loader silently failed
|
||||
if content is not None:
|
||||
max_length = getattr(__request__.app.state.config, 'WEB_FETCH_MAX_CONTENT_LENGTH', None)
|
||||
max_length = await Config.get('web.fetch.max_content_length')
|
||||
if max_length and max_length > 0 and len(content) > max_length:
|
||||
content = content[:max_length] + '\n\n[Content truncated...]'
|
||||
else:
|
||||
@@ -358,10 +365,11 @@ async def edit_image(
|
||||
__message_id__: str = None,
|
||||
) -> str:
|
||||
"""
|
||||
Edit existing images based on a text prompt.
|
||||
Transform one or more existing images according to a text prompt.
|
||||
Supports targeted edits such as adding, removing, replacing, inpainting, extending, or compositing image content.
|
||||
|
||||
:param prompt: A description of the changes to make to the images
|
||||
:param image_urls: A list of URLs of the images to edit
|
||||
:param prompt: A description of the transformation to apply to the provided images
|
||||
:param image_urls: Source image URLs to modify or use as composition inputs
|
||||
:return: Confirmation that the images were edited, or an error message
|
||||
"""
|
||||
if __request__ is None:
|
||||
@@ -475,7 +483,7 @@ async def execute_code(
|
||||
)
|
||||
code = blocking_code + '\n' + code
|
||||
|
||||
engine = getattr(__request__.app.state.config, 'CODE_INTERPRETER_ENGINE', 'pyodide')
|
||||
engine = await Config.get('code_interpreter.engine', 'pyodide')
|
||||
if engine == 'pyodide':
|
||||
# Execute via frontend pyodide using bidirectional event call
|
||||
if __event_call__ is None:
|
||||
@@ -514,20 +522,14 @@ async def execute_code(
|
||||
elif engine == 'jupyter':
|
||||
from open_webui.utils.code_interpreter import execute_code_jupyter
|
||||
|
||||
jupyter_auth = await Config.get('code_interpreter.jupyter.auth')
|
||||
|
||||
output = await execute_code_jupyter(
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_URL,
|
||||
await Config.get('code_interpreter.jupyter.url'),
|
||||
code,
|
||||
(
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_TOKEN
|
||||
if __request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH == 'token'
|
||||
else None
|
||||
),
|
||||
(
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH_PASSWORD
|
||||
if __request__.app.state.config.CODE_INTERPRETER_JUPYTER_AUTH == 'password'
|
||||
else None
|
||||
),
|
||||
__request__.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT,
|
||||
(await Config.get('code_interpreter.jupyter.auth_token') if jupyter_auth == 'token' else None),
|
||||
(await Config.get('code_interpreter.jupyter.auth_password') if jupyter_auth == 'password' else None),
|
||||
await Config.get('code_interpreter.jupyter.timeout'),
|
||||
)
|
||||
|
||||
stdout = output.get('stdout', '')
|
||||
@@ -593,17 +595,88 @@ async def execute_code(
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def search_memories(
|
||||
query: str,
|
||||
count: int = 5,
|
||||
async def list_memory_paths(
|
||||
query: str = '',
|
||||
count: int = 100,
|
||||
type: str = 'all',
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Search the user's stored memories for relevant information.
|
||||
List saved memory paths to find existing memory groups before writing or moving memories.
|
||||
|
||||
:param query: The search query to find relevant memories
|
||||
:param query: Optional query to filter memory paths or contents
|
||||
:param count: Maximum number of paths to return
|
||||
:param type: "user", "context", or "all"
|
||||
:return: JSON with memory paths, counts, children, and update times
|
||||
"""
|
||||
try:
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
result = await _list_memory_paths(
|
||||
ListMemoryPathsForm(
|
||||
query=query or None,
|
||||
type=type if type in {'user', 'context', 'all'} else 'all',
|
||||
limit=count,
|
||||
),
|
||||
user,
|
||||
)
|
||||
return json.dumps(result, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
log.exception(f'list_memory_paths error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
||||
|
||||
async def read_memory_path(
|
||||
path: str,
|
||||
count: int = 50,
|
||||
type: str = 'all',
|
||||
include_children: bool = True,
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Read saved memories at a memory path, including nearby parent and child paths.
|
||||
|
||||
:param path: Memory path to read
|
||||
:param count: Maximum number of memories to return
|
||||
:param type: "user", "context", or "all"
|
||||
:param include_children: Include memories under child paths
|
||||
:return: JSON with parent paths, child paths, and memories at the path
|
||||
"""
|
||||
try:
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
result = await _read_memory_path(
|
||||
ReadMemoryPathForm(
|
||||
path=path,
|
||||
type=type if type in {'user', 'context', 'all'} else 'all',
|
||||
include_children=include_children,
|
||||
limit=count,
|
||||
),
|
||||
user,
|
||||
)
|
||||
return json.dumps(result, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
log.exception(f'read_memory_path error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
||||
|
||||
async def search_memories(
|
||||
query: str = '',
|
||||
count: int = 5,
|
||||
type: str = 'all',
|
||||
path: Optional[str] = None,
|
||||
memory_id: Optional[str] = None,
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Search or browse saved memories by content, path, type, or memory ID.
|
||||
|
||||
:param query: Optional query to search memory content and path
|
||||
:param count: Number of memories to return (default 5)
|
||||
:param type: "user", "context", or "all"
|
||||
:param path: Optional memory path to search around
|
||||
:param memory_id: Optional exact memory ID to read
|
||||
:return: JSON with matching memories and their dates
|
||||
"""
|
||||
if __request__ is None:
|
||||
@@ -612,28 +685,34 @@ async def search_memories(
|
||||
try:
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
|
||||
results = await query_memory(
|
||||
__request__,
|
||||
QueryMemoryForm(content=query, k=count),
|
||||
memories = await _search_memories(
|
||||
SearchMemoriesForm(
|
||||
query=query or None,
|
||||
type=type if type in {'user', 'context', 'all'} else 'all',
|
||||
path=path,
|
||||
memory_id=memory_id,
|
||||
limit=count,
|
||||
),
|
||||
user,
|
||||
)
|
||||
|
||||
if results and hasattr(results, 'documents') and results.documents:
|
||||
memories = []
|
||||
for doc_idx, doc in enumerate(results.documents[0]):
|
||||
memory_id = None
|
||||
if results.ids and results.ids[0]:
|
||||
memory_id = results.ids[0][doc_idx]
|
||||
created_at = 'Unknown'
|
||||
if results.metadatas and results.metadatas[0][doc_idx].get('created_at'):
|
||||
created_at = time.strftime(
|
||||
'%Y-%m-%d',
|
||||
time.localtime(results.metadatas[0][doc_idx]['created_at']),
|
||||
)
|
||||
memories.append({'id': memory_id, 'date': created_at, 'content': doc})
|
||||
return json.dumps(memories, ensure_ascii=False)
|
||||
else:
|
||||
if not memories:
|
||||
return json.dumps([])
|
||||
|
||||
return json.dumps(
|
||||
[
|
||||
{
|
||||
'id': memory.id,
|
||||
'type': memory.type,
|
||||
'path': memory.path,
|
||||
'content': memory.content,
|
||||
'created_at': time.strftime('%Y-%m-%d', time.localtime(memory.created_at)),
|
||||
'updated_at': time.strftime('%Y-%m-%d', time.localtime(memory.updated_at)),
|
||||
}
|
||||
for memory in memories
|
||||
],
|
||||
ensure_ascii=False,
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception(f'search_memories error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
@@ -641,13 +720,17 @@ async def search_memories(
|
||||
|
||||
async def add_memory(
|
||||
content: str,
|
||||
type: str = 'user',
|
||||
path: Optional[str] = None,
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Store a new memory for the user.
|
||||
Save a user-provided preference, fact, or instruction as memory for future chats.
|
||||
|
||||
:param content: The memory content to store
|
||||
:param type: Use "user" for facts/preferences about the user, or "context" for other durable context
|
||||
:param path: Optional stable memory address for grouping related memories
|
||||
:return: Confirmation that the memory was stored
|
||||
"""
|
||||
if __request__ is None:
|
||||
@@ -658,27 +741,73 @@ async def add_memory(
|
||||
|
||||
memory = await _add_memory(
|
||||
__request__,
|
||||
AddMemoryForm(content=content),
|
||||
AddMemoryForm(content=content, type=Memories.normalize_memory_type(type), path=path),
|
||||
user,
|
||||
)
|
||||
|
||||
return json.dumps({'status': 'success', 'id': memory.id}, ensure_ascii=False)
|
||||
return json.dumps(
|
||||
{'status': 'success', 'id': memory.id, 'type': memory.type, 'path': memory.path},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception(f'add_memory error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
||||
|
||||
async def update_memory(
|
||||
operations: list[dict],
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Apply a batch of memory changes after learning durable information.
|
||||
|
||||
Use type "user" for facts, preferences, or instructions about the user.
|
||||
Use type "context" for other durable context that may help future chats.
|
||||
Path is optional. Use it as a stable memory address to group related memories.
|
||||
Prefer an existing path from list_memory_paths when one fits.
|
||||
Leave path empty when no useful grouping is clear.
|
||||
|
||||
Operation shapes:
|
||||
- {"action": "add", "content": "...", "type": "user"|"context", "path": "..."}
|
||||
- {"action": "replace", "id": "...", "content": "...", "type": "user"|"context", "path": "..."}
|
||||
- {"action": "move", "id": "...", "path": "..."}
|
||||
- {"action": "remove", "id": "..."}
|
||||
|
||||
:param operations: Memory operations to apply in one request
|
||||
:return: JSON with operation results
|
||||
"""
|
||||
if __request__ is None:
|
||||
return json.dumps({'error': 'Request context not available'})
|
||||
|
||||
try:
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
operation_results = await _update_memories(
|
||||
__request__,
|
||||
UpdateMemoriesForm(operations=operations),
|
||||
user,
|
||||
)
|
||||
return json.dumps(operation_results, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
log.exception(f'update_memory error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
||||
|
||||
async def replace_memory_content(
|
||||
memory_id: str,
|
||||
content: str,
|
||||
type: Optional[str] = None,
|
||||
path: Optional[str] = None,
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Update the content of an existing memory by its ID.
|
||||
Update an existing saved memory by its ID when its content needs correction.
|
||||
|
||||
:param memory_id: The ID of the memory to update
|
||||
:param content: The new content for the memory
|
||||
:param type: Optional "user" or "context" type for the updated memory
|
||||
:param path: Optional stable memory address for grouping related memories
|
||||
:return: Confirmation that the memory was updated
|
||||
"""
|
||||
if __request__ is None:
|
||||
@@ -690,12 +819,22 @@ async def replace_memory_content(
|
||||
memory = await update_memory_by_id(
|
||||
memory_id=memory_id,
|
||||
request=__request__,
|
||||
form_data=MemoryUpdateModel(content=content),
|
||||
form_data=MemoryUpdateModel(
|
||||
content=content,
|
||||
type=Memories.normalize_memory_type(type) if type else None,
|
||||
path=path,
|
||||
),
|
||||
user=user,
|
||||
)
|
||||
|
||||
return json.dumps(
|
||||
{'status': 'success', 'id': memory.id, 'content': memory.content},
|
||||
{
|
||||
'status': 'success',
|
||||
'id': memory.id,
|
||||
'type': memory.type,
|
||||
'path': memory.path,
|
||||
'content': memory.content,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
except Exception as e:
|
||||
@@ -709,7 +848,7 @@ async def delete_memory(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Delete a memory by its ID.
|
||||
Delete a saved memory by its ID.
|
||||
|
||||
:param memory_id: The ID of the memory to delete
|
||||
:return: Confirmation that the memory was deleted
|
||||
@@ -740,7 +879,7 @@ async def list_memories(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
List all stored memories for the user.
|
||||
List all stored memories for the user, including IDs and timestamps.
|
||||
|
||||
:return: JSON list of all memories with id, content, and dates
|
||||
"""
|
||||
@@ -753,16 +892,18 @@ async def list_memories(
|
||||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
|
||||
if memories:
|
||||
result = [
|
||||
memory_rows = [
|
||||
{
|
||||
'id': m.id,
|
||||
'type': m.type,
|
||||
'path': m.path,
|
||||
'content': m.content,
|
||||
'created_at': time.strftime('%Y-%m-%d %H:%M', time.localtime(m.created_at)),
|
||||
'updated_at': time.strftime('%Y-%m-%d %H:%M', time.localtime(m.updated_at)),
|
||||
}
|
||||
for m in memories
|
||||
]
|
||||
return json.dumps(result, ensure_ascii=False)
|
||||
return json.dumps(memory_rows, ensure_ascii=False)
|
||||
else:
|
||||
return json.dumps([])
|
||||
except Exception as e:
|
||||
@@ -784,7 +925,7 @@ async def search_notes(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Search the user's notes by title and content.
|
||||
Search the user's saved notes by title and content.
|
||||
|
||||
:param query: The search query to find matching notes
|
||||
:param count: Maximum number of results to return (default: 5)
|
||||
@@ -987,7 +1128,7 @@ async def replace_note_content(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Update the content of a note. Use this to modify task lists, add notes, or update content.
|
||||
Update the markdown content, and optionally the title, of an existing note.
|
||||
|
||||
:param note_id: The ID of the note to update
|
||||
:param content: The new markdown content for the note
|
||||
@@ -1064,6 +1205,7 @@ async def search_chats(
|
||||
) -> str:
|
||||
"""
|
||||
Search the user's previous chat conversations by title and message content.
|
||||
Helpful for finding details from earlier conversations.
|
||||
|
||||
:param query: The search query to find matching chats
|
||||
:param count: Maximum number of results to return (default: 5)
|
||||
@@ -1102,7 +1244,7 @@ async def search_chats(
|
||||
|
||||
# Find a matching message snippet
|
||||
snippet = ''
|
||||
messages = chat.chat.get('history', {}).get('messages', {})
|
||||
messages = (getattr(chat, 'chat', None) or {}).get('history', {}).get('messages', {})
|
||||
lower_query = query.lower()
|
||||
|
||||
for msg_id, msg in messages.items():
|
||||
@@ -1141,7 +1283,8 @@ async def view_chat(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the full conversation history of a chat by its ID.
|
||||
Get the full conversation history of a chat by its ID after a relevant
|
||||
previous chat has been identified.
|
||||
|
||||
:param chat_id: The ID of the chat to retrieve
|
||||
:return: JSON with the chat's id, title, and messages
|
||||
@@ -1211,7 +1354,7 @@ async def search_channels(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Search for channels by name and description that the user has access to.
|
||||
Search channels by name and description to find accessible team spaces.
|
||||
|
||||
:param query: The search query to find matching channels
|
||||
:param count: Maximum number of results to return (default: 5)
|
||||
@@ -1265,7 +1408,8 @@ async def search_channel_messages(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Search for messages in channels the user is a member of, including thread replies.
|
||||
Search messages in channels the user is a member of, including thread replies.
|
||||
Helpful for finding prior team/channel discussion.
|
||||
|
||||
:param query: The search query to find matching messages
|
||||
:param count: Maximum number of results to return (default: 10)
|
||||
@@ -1493,7 +1637,8 @@ async def list_knowledge_bases(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
List the user's accessible knowledge bases.
|
||||
List the user's accessible knowledge bases so a relevant internal source
|
||||
can be chosen.
|
||||
|
||||
:param count: Maximum number of KBs to return (default: 10)
|
||||
:param skip: Number of results to skip for pagination (default: 0)
|
||||
@@ -1551,7 +1696,8 @@ async def search_knowledge_bases(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Search the user's accessible knowledge bases by name and description.
|
||||
Search the user's accessible knowledge bases by name and description to find
|
||||
a relevant internal source.
|
||||
|
||||
:param query: The search query to find matching knowledge bases
|
||||
:param count: Maximum number of results to return (default: 5)
|
||||
@@ -1614,6 +1760,7 @@ async def search_knowledge_files(
|
||||
"""
|
||||
Search files by filename across knowledge bases the user has access to.
|
||||
When the model has attached knowledge, searches only within attached KBs and files.
|
||||
Helpful when looking for a specific document or file name.
|
||||
|
||||
:param query: The search query to find matching files by filename
|
||||
:param knowledge_id: Optional KB id to limit search to a specific knowledge base
|
||||
@@ -1785,6 +1932,7 @@ async def grep_knowledge_files(
|
||||
Search for exact text across knowledge files. Returns matching lines with line numbers.
|
||||
Unlike query_knowledge_files (semantic/vector search), this performs exact string matching.
|
||||
Automatically detects regex patterns (e.g. "error|warn", "version \\d+").
|
||||
Helpful for literal strings, identifiers, error messages, or regex-style searches.
|
||||
|
||||
:param pattern: The text pattern to search for (regex auto-detected)
|
||||
:param file_id: Optional file ID to search within a single file only
|
||||
@@ -2347,6 +2495,7 @@ async def query_knowledge_files(
|
||||
"""
|
||||
Search knowledge base files using semantic/vector search. Searches across collections (KBs),
|
||||
individual files, and notes that the user has access to.
|
||||
Helpful for internal documentation, uploaded knowledge, and attached model knowledge.
|
||||
|
||||
:param query: The search query to find semantically relevant content
|
||||
:param knowledge_ids: Optional list of KB ids to limit search to specific knowledge bases
|
||||
@@ -2383,6 +2532,7 @@ async def query_knowledge_files(
|
||||
from open_webui.models.files import Files
|
||||
from open_webui.models.knowledge import Knowledges
|
||||
from open_webui.models.notes import Notes
|
||||
from open_webui.retrieval.external import retrieve_external_knowledge
|
||||
from open_webui.retrieval.utils import query_collection
|
||||
|
||||
user_id = __user__.get('id')
|
||||
@@ -2394,6 +2544,7 @@ async def query_knowledge_files(
|
||||
return json.dumps({'error': 'Embedding function not configured'})
|
||||
|
||||
collection_names = []
|
||||
external_knowledges = []
|
||||
note_results = [] # Notes aren't vectorized, handle separately
|
||||
|
||||
# If model has attached knowledge, use those
|
||||
@@ -2416,7 +2567,10 @@ async def query_knowledge_files(
|
||||
user_group_ids=set(user_group_ids),
|
||||
)
|
||||
):
|
||||
collection_names.append(item_id)
|
||||
if (knowledge.meta or {}).get('source') == 'external':
|
||||
external_knowledges.append(knowledge)
|
||||
else:
|
||||
collection_names.append(item_id)
|
||||
|
||||
elif item_type == 'file':
|
||||
# Individual file - use file-{id} as collection name
|
||||
@@ -2462,7 +2616,10 @@ async def query_knowledge_files(
|
||||
user_group_ids=set(user_group_ids),
|
||||
)
|
||||
):
|
||||
collection_names.append(knowledge_id)
|
||||
if (knowledge.meta or {}).get('source') == 'external':
|
||||
external_knowledges.append(knowledge)
|
||||
else:
|
||||
collection_names.append(knowledge_id)
|
||||
else:
|
||||
# No model knowledge and no specific IDs - search all accessible KBs
|
||||
result = await Knowledges.search_knowledge_bases(
|
||||
@@ -2475,7 +2632,11 @@ async def query_knowledge_files(
|
||||
skip=0,
|
||||
limit=50,
|
||||
)
|
||||
collection_names = [knowledge_base.id for knowledge_base in result.items]
|
||||
for knowledge_base in result.items:
|
||||
if (knowledge_base.meta or {}).get('source') == 'external':
|
||||
external_knowledges.append(knowledge_base)
|
||||
else:
|
||||
collection_names.append(knowledge_base.id)
|
||||
|
||||
chunks = []
|
||||
|
||||
@@ -2507,6 +2668,31 @@ async def query_knowledge_files(
|
||||
chunk_info['distance'] = distances[idx]
|
||||
chunks.append(chunk_info)
|
||||
|
||||
for knowledge in external_knowledges:
|
||||
query_results = await retrieve_external_knowledge(
|
||||
__request__,
|
||||
knowledge,
|
||||
queries=[query],
|
||||
count=count,
|
||||
user=type('UserContext', (), {'id': user_id, 'role': user_role})(),
|
||||
)
|
||||
documents = query_results.get('documents', [[]])[0]
|
||||
metadatas = query_results.get('metadatas', [[]])[0]
|
||||
distances = query_results.get('distances', [[]])[0]
|
||||
|
||||
for idx, doc in enumerate(documents):
|
||||
metadata = metadatas[idx] if idx < len(metadatas) else {}
|
||||
chunk_info = {
|
||||
'content': doc,
|
||||
'source': metadata.get('source', metadata.get('name', knowledge.name)),
|
||||
'file_id': metadata.get('file_id', f'external-{knowledge.id}'),
|
||||
'type': 'external',
|
||||
'knowledge_id': knowledge.id,
|
||||
}
|
||||
if idx < len(distances):
|
||||
chunk_info['distance'] = distances[idx]
|
||||
chunks.append(chunk_info)
|
||||
|
||||
# Limit to requested count
|
||||
chunks = chunks[:count]
|
||||
|
||||
@@ -2525,7 +2711,7 @@ async def query_knowledge_bases(
|
||||
"""
|
||||
Search knowledge bases by semantic similarity to query.
|
||||
Finds KBs whose name/description match the meaning of your query.
|
||||
Use this to discover relevant knowledge bases before querying their files.
|
||||
Helpful for discovering which knowledge base to query next.
|
||||
|
||||
:param query: Natural language query describing what you're looking for
|
||||
:param count: Maximum results (default: 5)
|
||||
@@ -2731,9 +2917,7 @@ async def create_tasks(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Create a task checklist to track progress on multi-step work.
|
||||
Call this once at the start to define all steps, then use
|
||||
update_task to mark each task as you complete it.
|
||||
Create a visible task checklist for multi-step work so progress can be shown in chat.
|
||||
|
||||
:param tasks: List of task items. Each item: content (string, required), status (pending|in_progress|completed|cancelled, default pending), id (optional, auto-generated).
|
||||
:return: JSON with the full task list and summary counts
|
||||
@@ -2784,9 +2968,7 @@ async def update_task(
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Mark a single task as completed, in_progress, pending, or cancelled.
|
||||
Call this after finishing each step. You MUST call this for every
|
||||
task, including the very last one.
|
||||
Mark a single visible task item as completed, in_progress, pending, or cancelled.
|
||||
|
||||
:param id: The task ID to update
|
||||
:param status: New status: completed, in_progress, pending, or cancelled (default: completed)
|
||||
@@ -3226,8 +3408,7 @@ async def search_calendar_events(
|
||||
) -> str:
|
||||
"""
|
||||
Search calendar events, reminders, and scheduled items by text and/or date range.
|
||||
Use this to check what's coming up, find a specific event or reminder, or list
|
||||
the user's schedule for a time period.
|
||||
Helpful for finding upcoming events, reminders, or schedule items.
|
||||
|
||||
:param query: Search text to match against event title, description, or location (optional)
|
||||
:param start: Only return events starting at or after this datetime, e.g. "2026-04-20 00:00" (optional)
|
||||
|
||||
@@ -482,6 +482,11 @@ async def _kb_ls(args: list[str], flags: set[str], user: dict, model_knowledge:
|
||||
path_arg = args[0] if args else None
|
||||
|
||||
kb_ids = await _get_accessible_kb_ids(user, model_knowledge, knowledge_id=None)
|
||||
direct_files = (
|
||||
[f for f in await _get_accessible_files(user, model_knowledge) if not f.get('knowledge_id')]
|
||||
if model_knowledge
|
||||
else []
|
||||
)
|
||||
|
||||
# If path_arg looks like a KB ID, scope to that KB
|
||||
target_kb_id = None
|
||||
@@ -497,7 +502,7 @@ async def _kb_ls(args: list[str], flags: set[str], user: dict, model_knowledge:
|
||||
if target_kb_id:
|
||||
kb_ids = [(kid, kn, kd) for kid, kn, kd in kb_ids if kid == target_kb_id]
|
||||
|
||||
if not kb_ids:
|
||||
if not kb_ids and not direct_files:
|
||||
return 'No knowledge bases found.'
|
||||
|
||||
lines = []
|
||||
@@ -540,6 +545,12 @@ async def _kb_ls(args: list[str], flags: set[str], user: dict, model_knowledge:
|
||||
lines.append(' (empty)')
|
||||
lines.append('')
|
||||
|
||||
if direct_files and not target_kb_id and not dir_path:
|
||||
lines.append('Attached Files:')
|
||||
for f in direct_files:
|
||||
lines.append(f' {f["id"]} {f["filename"]} {_fmt_size(f)} {_fmt_date(f)}')
|
||||
lines.append('')
|
||||
|
||||
return '\n'.join(lines).rstrip()
|
||||
|
||||
|
||||
@@ -958,7 +969,12 @@ async def _kb_sed(
|
||||
async def _kb_tree(args: list[str], flags: set[str], user: dict, model_knowledge: list[dict] | None) -> str:
|
||||
"""Show directory tree structure."""
|
||||
kb_ids = await _get_accessible_kb_ids(user, model_knowledge)
|
||||
if not kb_ids:
|
||||
direct_files = (
|
||||
[f for f in await _get_accessible_files(user, model_knowledge) if not f.get('knowledge_id')]
|
||||
if model_knowledge
|
||||
else []
|
||||
)
|
||||
if not kb_ids and not direct_files:
|
||||
return 'No knowledge bases found.'
|
||||
|
||||
dir_scope = args[0].strip('/') if args else None
|
||||
@@ -1007,6 +1023,14 @@ async def _kb_tree(args: list[str], flags: set[str], user: dict, model_knowledge
|
||||
output.append(f'\n {total_dirs} directories, {total_files} files')
|
||||
output.append('')
|
||||
|
||||
if direct_files and not dir_scope:
|
||||
output.append('Attached Files:')
|
||||
for idx, f in enumerate(direct_files):
|
||||
connector = '└── ' if idx == len(direct_files) - 1 else '├── '
|
||||
output.append(f' {connector}{f["filename"]}')
|
||||
output.append(f'\n 0 directories, {len(direct_files)} files')
|
||||
output.append('')
|
||||
|
||||
return '\n'.join(output).rstrip()
|
||||
|
||||
|
||||
|
||||
@@ -114,7 +114,7 @@ async def has_access(
|
||||
Check if a user has the specified permission using an in-memory access_grants list.
|
||||
|
||||
Used for config-driven resources (arena models, tool servers) that store
|
||||
access control as JSON in ConfigVar rather than in the access_grant DB table.
|
||||
access control as JSON config rather than in the access_grant DB table.
|
||||
|
||||
Semantics:
|
||||
- None or [] → private (owner-only, deny all)
|
||||
@@ -321,7 +321,8 @@ async def check_model_access(
|
||||
return
|
||||
|
||||
if model_info:
|
||||
if user.role == 'user':
|
||||
# Enforce for every non-admin role (including pending); never fail open.
|
||||
if user.role != 'admin':
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
|
||||
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
|
||||
|
||||
@@ -38,25 +38,33 @@ async def has_access_to_file(
|
||||
if file.user_id == user.id:
|
||||
return True
|
||||
|
||||
# Check if the file is associated with any knowledge bases the user has access to
|
||||
# Check if the file is associated with any knowledge bases the user has access to.
|
||||
# An object (knowledge base or workspace model) confers write/delete on a file only when
|
||||
# the object's OWNER owns that file; otherwise a read-only file laundered into an object
|
||||
# the user controls would gain write/delete on it (CWE-863). Read access is unaffected.
|
||||
knowledge_bases = await Knowledges.get_knowledges_by_file_id(file_id, db=db)
|
||||
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)}
|
||||
for knowledge_base in knowledge_bases:
|
||||
if knowledge_base.user_id == user.id or await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge_base.id,
|
||||
permission=access_type,
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
):
|
||||
if (
|
||||
knowledge_base.user_id == user.id
|
||||
or await AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge_base.id,
|
||||
permission=access_type,
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
)
|
||||
) and (access_type == 'read' or knowledge_base.user_id == file.user_id):
|
||||
return True
|
||||
|
||||
knowledge_base_id = file.meta.get('collection_name') if file.meta else None
|
||||
if knowledge_base_id:
|
||||
knowledge_bases = await Knowledges.get_knowledge_bases_by_user_id(user.id, access_type, db=db)
|
||||
for knowledge_base in knowledge_bases:
|
||||
if knowledge_base.id == knowledge_base_id:
|
||||
if knowledge_base.id == knowledge_base_id and (
|
||||
access_type == 'read' or knowledge_base.user_id == file.user_id
|
||||
):
|
||||
return True
|
||||
|
||||
# Check if the file is associated with any channels the user has access to
|
||||
@@ -78,12 +86,14 @@ async def has_access_to_file(
|
||||
if accessible_ids:
|
||||
return True
|
||||
|
||||
# Check if the file is directly attached to a shared workspace model
|
||||
# Check if the file is directly attached to a shared workspace model (per the ownership
|
||||
# note above, model write is conferred only for files the model owner owns).
|
||||
for model in await Models.get_models_by_user_id(user.id, permission=access_type, db=db):
|
||||
knowledge_items = getattr(model.meta, 'knowledge', None) or []
|
||||
for item in knowledge_items:
|
||||
if isinstance(item, dict) and item.get('type') == 'file' and item.get('id') == file.id:
|
||||
return True
|
||||
if access_type == 'read' or model.user_id == file.user_id:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
from open_webui.models.access_grants import AccessGrants
|
||||
from open_webui.models.folders import FolderModel, Folders
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
|
||||
async def has_folder_access(user_id: str, folder: FolderModel, permission: str, db: AsyncSession) -> bool:
|
||||
"""Check if user has access to folder directly or via ancestor inheritance."""
|
||||
if folder.user_id == user_id:
|
||||
return True
|
||||
|
||||
if await AccessGrants.has_access(
|
||||
user_id=user_id,
|
||||
resource_type='folder',
|
||||
resource_id=folder.id,
|
||||
permission=permission,
|
||||
db=db,
|
||||
):
|
||||
return True
|
||||
# Check ancestor chain for inherited access
|
||||
if folder.parent_id:
|
||||
parent = await Folders.get_folder_by_id(folder.parent_id, db=db)
|
||||
if parent:
|
||||
return await has_folder_access(user_id, parent, permission, db)
|
||||
return False
|
||||
@@ -89,6 +89,26 @@ async def get_anthropic_models(url: str, key: str, user: UserModel = None) -> di
|
||||
##############################
|
||||
|
||||
|
||||
def _copy_cache_control(source: dict, target: dict) -> dict:
|
||||
if isinstance(source, dict) and 'cache_control' in source:
|
||||
target['cache_control'] = source['cache_control']
|
||||
return target
|
||||
|
||||
|
||||
def _has_cache_control(blocks: list) -> bool:
|
||||
return any(isinstance(block, dict) and 'cache_control' in block for block in blocks)
|
||||
|
||||
|
||||
def _finalize_openai_content(blocks: list) -> str | list:
|
||||
if not blocks:
|
||||
return ''
|
||||
|
||||
if len(blocks) == 1 and blocks[0].get('type') == 'text' and not _has_cache_control(blocks):
|
||||
return blocks[0].get('text', '')
|
||||
|
||||
return blocks
|
||||
|
||||
|
||||
def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
"""
|
||||
Convert an Anthropic Messages API request to OpenAI Chat Completions format.
|
||||
@@ -112,14 +132,21 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
if isinstance(system, str):
|
||||
messages.append({'role': 'system', 'content': system})
|
||||
elif isinstance(system, list):
|
||||
# Anthropic supports system as list of content blocks
|
||||
text_parts = []
|
||||
openai_content = []
|
||||
for block in system:
|
||||
if isinstance(block, dict) and block.get('type') == 'text':
|
||||
text_parts.append(block.get('text', ''))
|
||||
openai_content.append(
|
||||
_copy_cache_control(
|
||||
block,
|
||||
{
|
||||
'type': 'text',
|
||||
'text': block.get('text', ''),
|
||||
},
|
||||
)
|
||||
)
|
||||
elif isinstance(block, str):
|
||||
text_parts.append(block)
|
||||
messages.append({'role': 'system', 'content': '\n'.join(text_parts)})
|
||||
openai_content.append({'type': 'text', 'text': block})
|
||||
messages.append({'role': 'system', 'content': _finalize_openai_content(openai_content)})
|
||||
|
||||
# Convert messages
|
||||
for msg in anthropic_payload.get('messages', []):
|
||||
@@ -138,10 +165,13 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
|
||||
if block_type == 'text':
|
||||
openai_content.append(
|
||||
{
|
||||
'type': 'text',
|
||||
'text': block.get('text', ''),
|
||||
}
|
||||
_copy_cache_control(
|
||||
block,
|
||||
{
|
||||
'type': 'text',
|
||||
'text': block.get('text', ''),
|
||||
},
|
||||
)
|
||||
)
|
||||
elif block_type == 'image':
|
||||
source = block.get('source', {})
|
||||
@@ -149,19 +179,25 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
media_type = source.get('media_type', 'image/png')
|
||||
data = source.get('data', '')
|
||||
openai_content.append(
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {
|
||||
'url': f'data:{media_type};base64,{data}',
|
||||
_copy_cache_control(
|
||||
block,
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {
|
||||
'url': f'data:{media_type};base64,{data}',
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
elif source.get('type') == 'url':
|
||||
openai_content.append(
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {'url': source.get('url', '')},
|
||||
}
|
||||
_copy_cache_control(
|
||||
block,
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {'url': source.get('url', '')},
|
||||
},
|
||||
)
|
||||
)
|
||||
elif block_type == 'tool_use':
|
||||
tool_calls.append(
|
||||
@@ -196,10 +232,13 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
|
||||
if content_type == 'text':
|
||||
converted_parts.append(
|
||||
{
|
||||
'type': 'text',
|
||||
'text': content_block.get('text', ''),
|
||||
}
|
||||
_copy_cache_control(
|
||||
content_block,
|
||||
{
|
||||
'type': 'text',
|
||||
'text': content_block.get('text', ''),
|
||||
},
|
||||
)
|
||||
)
|
||||
elif content_type == 'image':
|
||||
source = content_block.get('source', {})
|
||||
@@ -207,21 +246,27 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
media_type = source.get('media_type', 'image/png')
|
||||
data = source.get('data', '')
|
||||
converted_parts.append(
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {
|
||||
'url': f'data:{media_type};base64,{data}',
|
||||
_copy_cache_control(
|
||||
content_block,
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {
|
||||
'url': f'data:{media_type};base64,{data}',
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
elif source.get('type') == 'url':
|
||||
converted_parts.append(
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {
|
||||
'url': source.get('url', ''),
|
||||
_copy_cache_control(
|
||||
content_block,
|
||||
{
|
||||
'type': 'image_url',
|
||||
'image_url': {
|
||||
'url': source.get('url', ''),
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
elif content_type == 'document':
|
||||
# Documents have no direct OpenAI equivalent;
|
||||
@@ -254,7 +299,9 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
converted_parts.append({'type': 'text', 'text': search_text})
|
||||
|
||||
# Flatten to string when only text parts are present
|
||||
if all(part.get('type') == 'text' for part in converted_parts):
|
||||
if all(part.get('type') == 'text' for part in converted_parts) and not _has_cache_control(
|
||||
converted_parts
|
||||
):
|
||||
tool_content = '\n'.join(part.get('text', '') for part in converted_parts)
|
||||
elif converted_parts:
|
||||
tool_content = converted_parts
|
||||
@@ -287,21 +334,13 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
# Assistant message with tool calls
|
||||
msg_dict = {'role': role}
|
||||
if openai_content:
|
||||
# If there's only text, flatten it
|
||||
if len(openai_content) == 1 and openai_content[0]['type'] == 'text':
|
||||
msg_dict['content'] = openai_content[0]['text']
|
||||
else:
|
||||
msg_dict['content'] = openai_content
|
||||
msg_dict['content'] = _finalize_openai_content(openai_content)
|
||||
else:
|
||||
msg_dict['content'] = ''
|
||||
msg_dict['tool_calls'] = tool_calls
|
||||
messages.append(msg_dict)
|
||||
elif openai_content:
|
||||
# If there's only a single text block, flatten it to a string
|
||||
if len(openai_content) == 1 and openai_content[0]['type'] == 'text':
|
||||
messages.append({'role': role, 'content': openai_content[0]['text']})
|
||||
else:
|
||||
messages.append({'role': role, 'content': openai_content})
|
||||
messages.append({'role': role, 'content': _finalize_openai_content(openai_content)})
|
||||
else:
|
||||
messages.append({'role': role, 'content': str(content) if content else ''})
|
||||
|
||||
@@ -312,7 +351,7 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
openai_payload['max_tokens'] = anthropic_payload['max_tokens']
|
||||
|
||||
# Common parameters
|
||||
for param in ('temperature', 'top_p', 'stop_sequences', 'stream'):
|
||||
for param in ('temperature', 'top_p', 'top_k', 'stop_sequences', 'stream', 'metadata', 'service_tier'):
|
||||
if param in anthropic_payload:
|
||||
if param == 'stop_sequences':
|
||||
openai_payload['stop'] = anthropic_payload[param]
|
||||
@@ -324,30 +363,33 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
|
||||
openai_tools = []
|
||||
for tool in anthropic_payload['tools']:
|
||||
openai_tools.append(
|
||||
{
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': tool.get('name', ''),
|
||||
'description': tool.get('description', ''),
|
||||
'parameters': tool.get('input_schema', {}),
|
||||
_copy_cache_control(
|
||||
tool,
|
||||
{
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': tool.get('name', ''),
|
||||
'description': tool.get('description', ''),
|
||||
'parameters': tool.get('input_schema', {}),
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
openai_payload['tools'] = openai_tools
|
||||
|
||||
# tool_choice
|
||||
if 'tool_choice' in anthropic_payload:
|
||||
tc = anthropic_payload['tool_choice']
|
||||
if isinstance(tc, dict):
|
||||
tc_type = tc.get('type', 'auto')
|
||||
if tc_type == 'auto':
|
||||
tool_choice = anthropic_payload['tool_choice']
|
||||
if isinstance(tool_choice, dict):
|
||||
tool_choice_type = tool_choice.get('type', 'auto')
|
||||
if tool_choice_type == 'auto':
|
||||
openai_payload['tool_choice'] = 'auto'
|
||||
elif tc_type == 'any':
|
||||
elif tool_choice_type == 'any':
|
||||
openai_payload['tool_choice'] = 'required'
|
||||
elif tc_type == 'tool':
|
||||
elif tool_choice_type == 'tool':
|
||||
openai_payload['tool_choice'] = {
|
||||
'type': 'function',
|
||||
'function': {'name': tc.get('name', '')},
|
||||
'function': {'name': tool_choice.get('name', '')},
|
||||
}
|
||||
|
||||
return openai_payload
|
||||
@@ -377,23 +419,23 @@ def convert_openai_to_anthropic_response(openai_response: dict, model: str = '')
|
||||
|
||||
# Build content blocks
|
||||
content = []
|
||||
msg_content = message.get('content')
|
||||
if msg_content:
|
||||
content.append({'type': 'text', 'text': msg_content})
|
||||
message_content = message.get('content')
|
||||
if message_content:
|
||||
content.append({'type': 'text', 'text': message_content})
|
||||
|
||||
# Tool calls → tool_use blocks
|
||||
tool_calls = message.get('tool_calls', [])
|
||||
for tc in tool_calls:
|
||||
func = tc.get('function', {})
|
||||
# Tool calls -> tool_use blocks
|
||||
tool_calls = message.get('tool_calls') or []
|
||||
for tool_call in tool_calls:
|
||||
function = tool_call.get('function', {})
|
||||
try:
|
||||
tool_input = json.loads(func.get('arguments', '{}'))
|
||||
tool_input = json.loads(function.get('arguments', '{}'))
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
tool_input = {}
|
||||
content.append(
|
||||
{
|
||||
'type': 'tool_use',
|
||||
'id': tc.get('id', f'toolu_{_uuid.uuid4().hex[:24]}'),
|
||||
'name': func.get('name', ''),
|
||||
'id': tool_call.get('id', f'toolu_{_uuid.uuid4().hex[:24]}'),
|
||||
'name': function.get('name', ''),
|
||||
'input': tool_input,
|
||||
}
|
||||
)
|
||||
@@ -404,6 +446,10 @@ def convert_openai_to_anthropic_response(openai_response: dict, model: str = '')
|
||||
'input_tokens': openai_usage.get('prompt_tokens', 0),
|
||||
'output_tokens': openai_usage.get('completion_tokens', 0),
|
||||
}
|
||||
if 'cache_creation_input_tokens' in openai_usage:
|
||||
usage['cache_creation_input_tokens'] = openai_usage['cache_creation_input_tokens']
|
||||
if 'cache_read_input_tokens' in openai_usage:
|
||||
usage['cache_read_input_tokens'] = openai_usage['cache_read_input_tokens']
|
||||
|
||||
return {
|
||||
'id': openai_response.get('id', f'msg_{_uuid.uuid4().hex[:24]}'),
|
||||
@@ -426,10 +472,14 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
|
||||
Handles text content, tool calls, and mixed content with proper
|
||||
multi-block indexing as required by Anthropic's streaming protocol.
|
||||
|
||||
Tool calls are tracked by their unique id (not OpenAI index) so that
|
||||
parallel calls sharing the same index get distinct Anthropic tool_use
|
||||
blocks. Each block follows the Anthropic lifecycle: start -> delta -> stop.
|
||||
"""
|
||||
import uuid as _uuid
|
||||
|
||||
msg_id = f'msg_{_uuid.uuid4().hex[:24]}'
|
||||
message_id = f'msg_{_uuid.uuid4().hex[:24]}'
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
stop_reason = 'end_turn'
|
||||
@@ -439,16 +489,21 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
current_block_index = 0
|
||||
text_block_open = False
|
||||
|
||||
# Track tool call state: maps OpenAI tool_call index -> Anthropic block index
|
||||
# This allows handling multiple concurrent tool calls.
|
||||
tool_call_blocks = {} # {openai_tc_index: anthropic_block_index}
|
||||
tool_call_started = {} # {openai_tc_index: bool}
|
||||
# Accumulated state for each tool call, keyed by tool call id.
|
||||
# Parallel calls that share the same OpenAI index get distinct entries.
|
||||
# Each entry: {id, name, arguments, block_index, started, stopped}
|
||||
tracked_tool_calls = {}
|
||||
# Map OpenAI tool call index -> tool call id for routing
|
||||
# argument-only deltas (deltas that carry arguments but no id).
|
||||
index_to_tool_id = {}
|
||||
# Whether any tool call block has been emitted (suppresses further text)
|
||||
has_tool_calls = False
|
||||
|
||||
# Emit message_start
|
||||
message_start = {
|
||||
'type': 'message_start',
|
||||
'message': {
|
||||
'id': msg_id,
|
||||
'id': message_id,
|
||||
'type': 'message',
|
||||
'role': 'assistant',
|
||||
'content': [],
|
||||
@@ -471,14 +526,14 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
if not line or not line.startswith('data:'):
|
||||
continue
|
||||
|
||||
data_str = line[5:].strip()
|
||||
if data_str == '[DONE]':
|
||||
data_string = line[5:].strip()
|
||||
if data_string == '[DONE]':
|
||||
continue
|
||||
if data_str == '{}':
|
||||
if data_string == '{}':
|
||||
continue
|
||||
|
||||
try:
|
||||
data = json.loads(data_str)
|
||||
data = json.loads(data_string)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
continue
|
||||
|
||||
@@ -492,6 +547,7 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
|
||||
delta = choices[0].get('delta', {})
|
||||
finish_reason = choices[0].get('finish_reason')
|
||||
message = choices[0].get('message') or {}
|
||||
|
||||
# Update usage if present
|
||||
if data.get('usage'):
|
||||
@@ -499,10 +555,11 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
output_tokens = data['usage'].get('completion_tokens', output_tokens)
|
||||
|
||||
# --- Handle text content ---
|
||||
# Anthropic expects text blocks before tool blocks, so skip
|
||||
# text deltas once any tool call has started.
|
||||
content = delta.get('content')
|
||||
if content is not None:
|
||||
if content and not has_tool_calls:
|
||||
if not text_block_open:
|
||||
# Start a new text content block
|
||||
block_start = {
|
||||
'type': 'content_block_start',
|
||||
'index': current_block_index,
|
||||
@@ -511,7 +568,6 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
yield f'event: content_block_start\ndata: {json.dumps(block_start)}\n\n'.encode()
|
||||
text_block_open = True
|
||||
|
||||
# Send text delta
|
||||
block_delta = {
|
||||
'type': 'content_block_delta',
|
||||
'index': current_block_index,
|
||||
@@ -520,7 +576,12 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
yield f'event: content_block_delta\ndata: {json.dumps(block_delta)}\n\n'.encode()
|
||||
|
||||
# --- Handle tool calls ---
|
||||
tool_calls = delta.get('tool_calls')
|
||||
# Some providers put tool_calls on the final message object
|
||||
# instead of the delta; fall back to that when needed.
|
||||
tool_calls = delta.get('tool_calls') or []
|
||||
if not tool_calls and message.get('tool_calls'):
|
||||
tool_calls = message['tool_calls']
|
||||
|
||||
if tool_calls:
|
||||
# Close text block if one is open (text comes before tools)
|
||||
if text_block_open:
|
||||
@@ -532,43 +593,95 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
text_block_open = False
|
||||
current_block_index += 1
|
||||
|
||||
for tc in tool_calls:
|
||||
tc_index = tc.get('index', 0)
|
||||
for tool_call in tool_calls:
|
||||
tool_call_index = tool_call.get('index', 0)
|
||||
tool_call_id = tool_call.get('id', '')
|
||||
tool_call_name = (tool_call.get('function') or {}).get('name', '')
|
||||
arguments_chunk = (tool_call.get('function') or {}).get('arguments', '')
|
||||
|
||||
if tc_index not in tool_call_started:
|
||||
# First time seeing this tool call — emit content_block_start
|
||||
tool_call_blocks[tc_index] = current_block_index
|
||||
tool_call_started[tc_index] = True
|
||||
# Resolve which tracked tool call this delta belongs to.
|
||||
# A delta with an id starts or identifies a specific tool.
|
||||
# A delta without an id carries arguments for the most
|
||||
# recent tool at this OpenAI index.
|
||||
if tool_call_id:
|
||||
if tool_call_id not in tracked_tool_calls:
|
||||
tracked_tool_calls[tool_call_id] = {
|
||||
'id': tool_call_id,
|
||||
'name': tool_call_name,
|
||||
'arguments': '',
|
||||
'block_index': -1,
|
||||
'started': False,
|
||||
'stopped': False,
|
||||
}
|
||||
index_to_tool_id[tool_call_index] = tool_call_id
|
||||
tool = tracked_tool_calls[tool_call_id]
|
||||
elif tool_call_index in index_to_tool_id:
|
||||
tool = tracked_tool_calls[index_to_tool_id[tool_call_index]]
|
||||
else:
|
||||
# First delta for this index with no id; create a
|
||||
# provisional entry with a generated fallback id.
|
||||
fallback_id = f'toolu_{_uuid.uuid4().hex[:24]}'
|
||||
tracked_tool_calls[fallback_id] = {
|
||||
'id': fallback_id,
|
||||
'name': tool_call_name,
|
||||
'arguments': '',
|
||||
'block_index': -1,
|
||||
'started': False,
|
||||
'stopped': False,
|
||||
}
|
||||
index_to_tool_id[tool_call_index] = fallback_id
|
||||
tool = tracked_tool_calls[fallback_id]
|
||||
|
||||
# Extract tool call ID and name from the first chunk
|
||||
tc_id = tc.get('id', f'toolu_{_uuid.uuid4().hex[:24]}')
|
||||
tc_name = tc.get('function', {}).get('name', '')
|
||||
# Update name if provided on a later delta
|
||||
if tool_call_name and not tool['name']:
|
||||
tool['name'] = tool_call_name
|
||||
|
||||
# Emit content_block_start once we have a name
|
||||
if not tool['started'] and tool['name']:
|
||||
tool['block_index'] = current_block_index
|
||||
tool['started'] = True
|
||||
has_tool_calls = True
|
||||
|
||||
block_start = {
|
||||
'type': 'content_block_start',
|
||||
'index': current_block_index,
|
||||
'content_block': {
|
||||
'type': 'tool_use',
|
||||
'id': tc_id,
|
||||
'name': tc_name,
|
||||
'id': tool['id'],
|
||||
'name': tool['name'],
|
||||
'input': {},
|
||||
},
|
||||
}
|
||||
yield f'event: content_block_start\ndata: {json.dumps(block_start)}\n\n'.encode()
|
||||
current_block_index += 1
|
||||
|
||||
# Emit argument chunks as input_json_delta
|
||||
args_chunk = tc.get('function', {}).get('arguments', '')
|
||||
if args_chunk:
|
||||
block_delta = {
|
||||
'type': 'content_block_delta',
|
||||
'index': tool_call_blocks[tc_index],
|
||||
'delta': {
|
||||
'type': 'input_json_delta',
|
||||
'partial_json': args_chunk,
|
||||
},
|
||||
}
|
||||
yield f'event: content_block_delta\ndata: {json.dumps(block_delta)}\n\n'.encode()
|
||||
# Buffer arguments and emit as input_json_delta
|
||||
if arguments_chunk:
|
||||
tool['arguments'] += arguments_chunk
|
||||
|
||||
if tool['started'] and not tool['stopped']:
|
||||
block_delta = {
|
||||
'type': 'content_block_delta',
|
||||
'index': tool['block_index'],
|
||||
'delta': {
|
||||
'type': 'input_json_delta',
|
||||
'partial_json': arguments_chunk,
|
||||
},
|
||||
}
|
||||
yield f'event: content_block_delta\ndata: {json.dumps(block_delta)}\n\n'.encode()
|
||||
|
||||
# Close the block once arguments form complete JSON
|
||||
if tool['started'] and not tool['stopped']:
|
||||
try:
|
||||
json.loads(tool['arguments'])
|
||||
tool['stopped'] = True
|
||||
block_stop = {
|
||||
'type': 'content_block_stop',
|
||||
'index': tool['block_index'],
|
||||
}
|
||||
yield f'event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n'.encode()
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
|
||||
# --- Handle finish reason ---
|
||||
if finish_reason is not None:
|
||||
@@ -582,15 +695,46 @@ async def openai_stream_to_anthropic_stream(openai_stream_generator, model: str
|
||||
except Exception as e:
|
||||
log.error(f'Error in Anthropic stream conversion: {e}')
|
||||
|
||||
# Flush any tools that buffered arguments but never emitted a block
|
||||
for tool in tracked_tool_calls.values():
|
||||
if not tool['started'] and tool['name']:
|
||||
tool['block_index'] = current_block_index
|
||||
tool['started'] = True
|
||||
|
||||
block_start = {
|
||||
'type': 'content_block_start',
|
||||
'index': current_block_index,
|
||||
'content_block': {
|
||||
'type': 'tool_use',
|
||||
'id': tool['id'],
|
||||
'name': tool['name'],
|
||||
'input': {},
|
||||
},
|
||||
}
|
||||
yield f'event: content_block_start\ndata: {json.dumps(block_start)}\n\n'.encode()
|
||||
current_block_index += 1
|
||||
|
||||
if tool['arguments']:
|
||||
block_delta = {
|
||||
'type': 'content_block_delta',
|
||||
'index': tool['block_index'],
|
||||
'delta': {
|
||||
'type': 'input_json_delta',
|
||||
'partial_json': tool['arguments'],
|
||||
},
|
||||
}
|
||||
yield f'event: content_block_delta\ndata: {json.dumps(block_delta)}\n\n'.encode()
|
||||
|
||||
# Close any open text block
|
||||
if text_block_open:
|
||||
block_stop = {'type': 'content_block_stop', 'index': current_block_index}
|
||||
yield f'event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n'.encode()
|
||||
|
||||
# Close any open tool call blocks
|
||||
for tc_index, block_index in tool_call_blocks.items():
|
||||
block_stop = {'type': 'content_block_stop', 'index': block_index}
|
||||
yield f'event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n'.encode()
|
||||
# Close any tool call blocks that are still open
|
||||
for tool in tracked_tool_calls.values():
|
||||
if tool['started'] and not tool['stopped']:
|
||||
block_stop = {'type': 'content_block_stop', 'index': tool['block_index']}
|
||||
yield f'event: content_block_stop\ndata: {json.dumps(block_stop)}\n\n'.encode()
|
||||
|
||||
# Emit message_delta with stop reason
|
||||
message_delta = {
|
||||
|
||||
@@ -39,6 +39,7 @@ from fastapi.responses import JSONResponse, RedirectResponse
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
from open_webui.env import CUSTOM_API_KEY_HEADER
|
||||
from open_webui.internal.db import ScopedSession
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.auth import get_http_authorization_cred
|
||||
from starlette.datastructures import MutableHeaders
|
||||
from starlette.requests import Request
|
||||
@@ -165,7 +166,7 @@ class AuthTokenMiddleware:
|
||||
token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=api_key)
|
||||
|
||||
request.state.token = token
|
||||
request.state.enable_api_keys = self._fastapi_app.state.config.ENABLE_API_KEYS
|
||||
request.state.enable_api_keys = await Config.get('auth.enable_api_keys')
|
||||
|
||||
async def send_with_timing(message: Message) -> None:
|
||||
if message['type'] == 'http.response.start':
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
@@ -25,6 +26,7 @@ from open_webui.env import (
|
||||
ENABLE_PASSWORD_VALIDATION,
|
||||
LICENSE_BLOB,
|
||||
OFFLINE_MODE,
|
||||
PASSWORD_HASH_ALGORITHM,
|
||||
PASSWORD_VALIDATION_HINT,
|
||||
PASSWORD_VALIDATION_REGEX_PATTERN,
|
||||
REDIS_KEY_PREFIX,
|
||||
@@ -35,6 +37,7 @@ from open_webui.env import (
|
||||
pk,
|
||||
)
|
||||
from open_webui.models.auths import Auths
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from pytz import UTC
|
||||
@@ -43,6 +46,7 @@ log = logging.getLogger(__name__)
|
||||
|
||||
SESSION_SECRET = WEBUI_SECRET_KEY
|
||||
ALGORITHM = 'HS256'
|
||||
PASSWORD_BCRYPT_MAX_BYTES = 72
|
||||
|
||||
##############
|
||||
# Auth Utils
|
||||
@@ -157,14 +161,21 @@ def get_license_data(app, key):
|
||||
bearer_security = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
def get_password_hash(password: str) -> str:
|
||||
"""Hash a password using bcrypt"""
|
||||
return bcrypt.hashpw(password.encode('utf-8'), bcrypt.gensalt()).decode('utf-8')
|
||||
async def get_password_hash(password: str) -> str:
|
||||
"""Hash a password using the configured algorithm in a thread pool."""
|
||||
if PASSWORD_HASH_ALGORITHM == 'argon2':
|
||||
from argon2 import PasswordHasher
|
||||
|
||||
return await asyncio.to_thread(PasswordHasher().hash, password)
|
||||
if PASSWORD_HASH_ALGORITHM == 'bcrypt':
|
||||
return (await asyncio.to_thread(bcrypt.hashpw, password.encode('utf-8'), bcrypt.gensalt())).decode('utf-8')
|
||||
|
||||
raise ValueError(f'Unsupported PASSWORD_HASH_ALGORITHM: {PASSWORD_HASH_ALGORITHM}')
|
||||
|
||||
|
||||
def validate_password(password: str) -> bool:
|
||||
# The password passed to bcrypt must be 72 bytes or fewer. If it is longer, it will be truncated before hashing.
|
||||
if len(password.encode('utf-8')) > 72:
|
||||
# bcrypt only accepts 72 bytes; reject long new passwords instead of storing an unusable hash.
|
||||
if PASSWORD_HASH_ALGORITHM == 'bcrypt' and len(password.encode('utf-8')) > PASSWORD_BCRYPT_MAX_BYTES:
|
||||
raise Exception(
|
||||
ERROR_MESSAGES.PASSWORD_TOO_LONG,
|
||||
)
|
||||
@@ -176,16 +187,29 @@ def validate_password(password: str) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def verify_password(plain_password: str, hashed_password: str) -> bool:
|
||||
"""Verify a password against its hash"""
|
||||
return (
|
||||
bcrypt.checkpw(
|
||||
plain_password.encode('utf-8'),
|
||||
async def verify_password(plain_password: str, hashed_password: str) -> bool:
|
||||
"""Verify a password using the algorithm encoded in its hash."""
|
||||
if not hashed_password:
|
||||
return False
|
||||
|
||||
if hashed_password.startswith('$argon2'):
|
||||
from argon2 import PasswordHasher
|
||||
from argon2.exceptions import InvalidHashError, VerificationError
|
||||
|
||||
try:
|
||||
return await asyncio.to_thread(PasswordHasher().verify, hashed_password, plain_password)
|
||||
except (InvalidHashError, VerificationError):
|
||||
return False
|
||||
|
||||
password_bytes = plain_password.encode('utf-8')[:PASSWORD_BCRYPT_MAX_BYTES]
|
||||
try:
|
||||
return await asyncio.to_thread(
|
||||
bcrypt.checkpw,
|
||||
password_bytes,
|
||||
hashed_password.encode('utf-8'),
|
||||
)
|
||||
if hashed_password
|
||||
else None
|
||||
)
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
# Let the one who signed this token be remembered at every gate,
|
||||
@@ -213,25 +237,25 @@ def decode_token(token: str) -> dict | None:
|
||||
return None
|
||||
|
||||
|
||||
async def is_valid_token(request, decoded) -> bool:
|
||||
async def is_valid_token(decoded, redis=None) -> bool:
|
||||
"""
|
||||
Check whether a JWT has been revoked. Two mechanisms:
|
||||
1. Per-token (jti) — used by user-initiated sign-out (known jti).
|
||||
2. Per-user (revoked_at) — used by OIDC back-channel logout when
|
||||
individual jti values are unknown; rejects tokens with iat <= revoked_at.
|
||||
"""
|
||||
if request.app.state.redis:
|
||||
if redis:
|
||||
# Per-token revocation
|
||||
jti = decoded.get('jti')
|
||||
if jti:
|
||||
revoked = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked')
|
||||
revoked = await redis.get(f'{REDIS_KEY_PREFIX}:auth:token:{jti}:revoked')
|
||||
if revoked:
|
||||
return False
|
||||
|
||||
# Per-user revocation (OIDC back-channel logout)
|
||||
user_id = decoded.get('id')
|
||||
if user_id:
|
||||
revoked_at = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at')
|
||||
revoked_at = await redis.get(f'{REDIS_KEY_PREFIX}:auth:user:{user_id}:revoked_at')
|
||||
if revoked_at:
|
||||
try:
|
||||
revoked_at_ts = int(revoked_at)
|
||||
@@ -341,7 +365,7 @@ async def get_current_user(
|
||||
)
|
||||
|
||||
if data is not None and 'id' in data:
|
||||
if data.get('jti') and not await is_valid_token(request, data):
|
||||
if not await is_valid_token(data, getattr(request.app.state, 'redis', None)):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail='Invalid token',
|
||||
@@ -409,12 +433,16 @@ async def get_current_user_by_api_key(request, api_key: str):
|
||||
detail=ERROR_MESSAGES.INVALID_TOKEN,
|
||||
)
|
||||
|
||||
user_permissions = await Config.get('user.permissions')
|
||||
enable_endpoint_restrictions = await Config.get('auth.api_key.endpoint_restrictions')
|
||||
allowed_endpoints = await Config.get('auth.api_key.allowed_endpoints', '')
|
||||
|
||||
if not request.state.enable_api_keys or (
|
||||
user.role != 'admin'
|
||||
and not await has_permission(
|
||||
user.id,
|
||||
'features.api_keys',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
user_permissions,
|
||||
)
|
||||
):
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.API_KEY_NOT_ALLOWED)
|
||||
@@ -422,10 +450,8 @@ async def get_current_user_by_api_key(request, api_key: str):
|
||||
# Enforce endpoint restrictions — checked here (not in middleware)
|
||||
# so it applies regardless of how the API key was transported
|
||||
# (Authorization header, cookie, x-api-key header, etc.).
|
||||
if request.app.state.config.ENABLE_API_KEYS_ENDPOINT_RESTRICTIONS:
|
||||
allowed_paths = [
|
||||
path.strip() for path in str(request.app.state.config.API_KEYS_ALLOWED_ENDPOINTS).split(',') if path.strip()
|
||||
]
|
||||
if enable_endpoint_restrictions:
|
||||
allowed_paths = [path.strip() for path in str(allowed_endpoints).split(',') if path.strip()]
|
||||
request_path = request.scope['path'] # Use raw ASGI path — not spoofable via Host header (CVE-2026-48710)
|
||||
is_allowed = any(request_path == allowed or request_path.startswith(allowed + '/') for allowed in allowed_paths)
|
||||
if not is_allowed:
|
||||
@@ -483,7 +509,7 @@ async def create_admin_user(email: str, password: str, name: str = 'Admin'):
|
||||
|
||||
log.info(f'Creating admin account from environment variables: {email}')
|
||||
try:
|
||||
hashed = get_password_hash(password)
|
||||
hashed = await get_password_hash(password)
|
||||
user = await Auths.insert_new_auth(
|
||||
email=email.lower(),
|
||||
password=hashed,
|
||||
|
||||
@@ -18,18 +18,23 @@ import logging
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from dateutil.rrule import rrulestr
|
||||
from fastapi import Request
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.events import EVENTS, publish_event
|
||||
from open_webui.internal.db import get_async_db
|
||||
from open_webui.models.automations import AutomationModel, AutomationRuns, Automations
|
||||
from open_webui.models.chats import ChatForm, Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.users import Users
|
||||
from open_webui.utils.auth import create_token
|
||||
from open_webui.utils.misc import parse_duration
|
||||
from open_webui.utils.task import prompt_template
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
@@ -172,7 +177,7 @@ async def scheduler_worker_loop(app) -> None:
|
||||
while True:
|
||||
try:
|
||||
# ── Automations ──
|
||||
if getattr(app.state.config, 'ENABLE_AUTOMATIONS', False):
|
||||
if await Config.get('automations.enable'):
|
||||
try:
|
||||
async with get_async_db() as db:
|
||||
batch = await Automations.claim_due(int(time.time_ns()), limit=10, db=db)
|
||||
@@ -184,7 +189,7 @@ async def scheduler_worker_loop(app) -> None:
|
||||
log.exception('Scheduler: automation error')
|
||||
|
||||
# ── Calendar Alerts ──
|
||||
if getattr(app.state.config, 'ENABLE_CALENDAR', False):
|
||||
if await Config.get('calendar.enable'):
|
||||
try:
|
||||
await _check_calendar_alerts(app)
|
||||
except Exception:
|
||||
@@ -202,11 +207,18 @@ async def scheduler_worker_loop(app) -> None:
|
||||
####################
|
||||
|
||||
|
||||
def _build_request(app) -> Request:
|
||||
def _build_request(
|
||||
app,
|
||||
token: Optional[str] = None,
|
||||
) -> Request:
|
||||
"""Build a minimal ASGI Request for chat_completion.
|
||||
|
||||
Mirrors the mock-request pattern used in main.py lifespan
|
||||
(model pre-fetch, tool server init) for consistency.
|
||||
|
||||
When token is provided, attach it as
|
||||
request.state.token so session-auth tool servers and terminals can
|
||||
authenticate headless scheduled runs as the automation owner.
|
||||
"""
|
||||
scope = {
|
||||
'type': 'http',
|
||||
@@ -222,7 +234,7 @@ def _build_request(app) -> Request:
|
||||
}
|
||||
request = Request(scope)
|
||||
# Ensure request.state is initialized with required attributes
|
||||
request.state.token = None
|
||||
request.state.token = HTTPAuthorizationCredentials(scheme='Bearer', credentials=token) if token else None
|
||||
request.state.enable_api_keys = False
|
||||
return request
|
||||
|
||||
@@ -239,7 +251,7 @@ def _resolve_model_tool_ids(app, model_id: str) -> list[str]:
|
||||
return list(tool_ids) if tool_ids else []
|
||||
|
||||
|
||||
def _resolve_model_features(app, model_id: str) -> dict:
|
||||
async def _resolve_model_features(app, model_id: str) -> dict:
|
||||
"""Read model default features from model config.
|
||||
|
||||
The frontend does this in Chat.svelte (model.info.meta.defaultFeatureIds
|
||||
@@ -256,14 +268,13 @@ def _resolve_model_features(app, model_id: str) -> dict:
|
||||
return {}
|
||||
|
||||
capabilities = meta.get('capabilities', {})
|
||||
config = app.state.config
|
||||
features = {}
|
||||
|
||||
# code_interpreter is excluded: it requires the frontend event emitter
|
||||
# and does not work in headless backend execution.
|
||||
feature_checks = {
|
||||
'web_search': getattr(config, 'ENABLE_WEB_SEARCH', False),
|
||||
'image_generation': getattr(config, 'ENABLE_IMAGE_GENERATION', False),
|
||||
'web_search': await Config.get('web.search.enable'),
|
||||
'image_generation': await Config.get('image_generation.enable'),
|
||||
}
|
||||
|
||||
for feature_id in default_feature_ids:
|
||||
@@ -357,6 +368,30 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
user = await Users.get_user_by_id(automation.user_id)
|
||||
if not user:
|
||||
await _record_run(automation.id, 'error', error='User not found')
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.AUTOMATION_RUN_FAILED,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name, 'error': 'User not found'},
|
||||
)
|
||||
return
|
||||
|
||||
# Re-gate the rehydrated owner: a demoted/deactivated or de-permissioned owner must not run.
|
||||
from open_webui.utils.access_control import has_permission
|
||||
|
||||
if user.role not in ('user', 'admin') or (
|
||||
user.role != 'admin'
|
||||
and not await has_permission(user.id, 'features.automations', await Config.get('user.permissions'))
|
||||
):
|
||||
error = 'Owner no longer permitted to run automations'
|
||||
await _record_run(automation.id, 'error', error=error)
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.AUTOMATION_RUN_FAILED,
|
||||
actor=user,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name, 'error': error},
|
||||
)
|
||||
return
|
||||
|
||||
prompt = await prompt_template(automation.data['prompt'], user)
|
||||
@@ -408,7 +443,15 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
)
|
||||
|
||||
if not chat:
|
||||
await _record_run(automation.id, 'error', error='Failed to create chat')
|
||||
error = 'Failed to create chat'
|
||||
await _record_run(automation.id, 'error', error=error)
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.AUTOMATION_RUN_FAILED,
|
||||
actor=user,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name, 'error': error},
|
||||
)
|
||||
return
|
||||
|
||||
# Notify frontend to refresh chat list
|
||||
@@ -426,7 +469,7 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
|
||||
# Resolve model defaults (frontend does this, backend doesn't)
|
||||
tool_ids = _resolve_model_tool_ids(app, model_id)
|
||||
features = _resolve_model_features(app, model_id)
|
||||
features = await _resolve_model_features(app, model_id)
|
||||
filter_ids = _resolve_model_filter_ids(app, model_id)
|
||||
|
||||
# Resolve terminal from model config
|
||||
@@ -460,7 +503,15 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
|
||||
# Call the full chat completion pipeline (same as POST /api/chat/completions).
|
||||
# The handler reference is stored on app.state to avoid circular imports.
|
||||
request = _build_request(app)
|
||||
try:
|
||||
expires_delta = parse_duration(str(await Config.get('automations.auth_token_expires_in', '1h')))
|
||||
except ValueError:
|
||||
expires_delta = None
|
||||
token = create_token(
|
||||
data={'id': user.id, 'typ': 'automation'},
|
||||
expires_delta=expires_delta or timedelta(hours=1),
|
||||
)
|
||||
request = _build_request(app, token=token)
|
||||
await app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user)
|
||||
|
||||
# Notify user
|
||||
@@ -478,10 +529,24 @@ async def execute_automation(app, automation: AutomationModel) -> None:
|
||||
)
|
||||
|
||||
await _record_run(automation.id, 'success', chat_id=chat.id)
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.AUTOMATION_RUN_COMPLETED,
|
||||
actor=user,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name, 'chat_id': chat.id},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
log.exception(f'Automation {automation.id} failed')
|
||||
await _record_run(automation.id, 'error', error=str(e)[:4000])
|
||||
error = str(e)[:4000]
|
||||
await _record_run(automation.id, 'error', error=error)
|
||||
await publish_event(
|
||||
app,
|
||||
EVENTS.AUTOMATION_RUN_FAILED,
|
||||
subject_id=automation.id,
|
||||
data={'name': automation.name, 'error': error},
|
||||
)
|
||||
|
||||
|
||||
####################
|
||||
@@ -551,7 +616,7 @@ async def _check_calendar_alerts(app) -> None:
|
||||
# Send webhook notification if user has one configured
|
||||
try:
|
||||
webui_name = getattr(app.state, 'WEBUI_NAME', 'Open WebUI')
|
||||
enable_user_webhooks = getattr(app.state.config, 'ENABLE_USER_WEBHOOKS', False)
|
||||
enable_user_webhooks = await Config.get('ui.enable_user_webhooks')
|
||||
|
||||
if enable_user_webhooks:
|
||||
user = await Users.get_user_by_id(event.user_id)
|
||||
|
||||
@@ -36,14 +36,16 @@ def expand_recurring_event(
|
||||
range_end_dt = datetime.fromtimestamp(range_end_ns / 1_000_000_000)
|
||||
scan_start = range_start_dt - timedelta(days=1)
|
||||
|
||||
original_start_ns = event_dict['start_at']
|
||||
original_start_dt = datetime.fromtimestamp(original_start_ns / 1_000_000_000)
|
||||
|
||||
try:
|
||||
# Parse with dtstart near the range so we never iterate from epoch
|
||||
rule = rrulestr(rrule_str, dtstart=scan_start, ignoretz=True)
|
||||
# Anchor to the event's real start so day-of-week / day-of-month are correct
|
||||
rule = rrulestr(rrule_str, dtstart=original_start_dt, ignoretz=True)
|
||||
except Exception:
|
||||
log.warning(f'Failed to parse RRULE for event {event_dict.get("id")}: {rrule_str}')
|
||||
return [event_dict]
|
||||
|
||||
original_start_ns = event_dict['start_at']
|
||||
original_end_ns = event_dict.get('end_at')
|
||||
duration_ns = (original_end_ns - original_start_ns) if original_end_ns else None
|
||||
|
||||
|
||||
@@ -237,6 +237,12 @@ async def generate_chat_completion(
|
||||
|
||||
form_data['model'] = selected_model_id
|
||||
|
||||
# bypass_filter recursion below skips the line-200 check; gate the resolved model here.
|
||||
if not bypass_filter and user.role == 'user':
|
||||
selected_model = request.app.state.MODELS.get(selected_model_id)
|
||||
if selected_model:
|
||||
await check_model_access(user, selected_model)
|
||||
|
||||
if selected_model_id:
|
||||
if form_data.get('stream') == True:
|
||||
|
||||
|
||||
@@ -0,0 +1,379 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.utils.misc import get_content_from_message, get_last_user_message, get_message_list
|
||||
from open_webui.utils.task import (
|
||||
get_task_model_id,
|
||||
prompt_template,
|
||||
prompt_variables_template,
|
||||
replace_messages_variable,
|
||||
replace_prompt_variable,
|
||||
)
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_CONTEXT_COMPACTION_PROMPT = """### Task:
|
||||
Summarize the conversation history that will be compacted out of the active chat context.
|
||||
|
||||
### Instructions:
|
||||
- Preserve key decisions, user preferences, and constraints.
|
||||
- Preserve files, artifacts, tool results, and code changes that matter going forward.
|
||||
- Preserve the current task state, unresolved questions, and next steps.
|
||||
- Be factual and specific. Do not invent details.
|
||||
- Keep the summary concise, but complete enough for the assistant to continue without the removed messages.
|
||||
|
||||
### Previous Summary:
|
||||
{{PREVIOUS_SUMMARY}}
|
||||
|
||||
### Messages Being Compacted:
|
||||
{{COMPACTED_MESSAGES}}
|
||||
|
||||
### Recent Messages Kept In Context:
|
||||
{{RECENT_MESSAGES}}"""
|
||||
|
||||
|
||||
async def compact_messages_for_request(
|
||||
request,
|
||||
user,
|
||||
messages: list[dict],
|
||||
metadata: dict,
|
||||
model_id: str,
|
||||
models: dict,
|
||||
system_prompt: str = '',
|
||||
) -> tuple[list[dict], str | None, bool]:
|
||||
config = await _load_config()
|
||||
if not config['enable']:
|
||||
return messages, None, False
|
||||
|
||||
messages, previous_summary = _apply_latest_summary_checkpoint(messages)
|
||||
token_threshold = _resolve_token_threshold(config['token_threshold'], metadata)
|
||||
if not _exceeds_token_threshold(messages, system_prompt, previous_summary, token_threshold) or len(messages) <= 3:
|
||||
return messages, previous_summary, False
|
||||
|
||||
boundary = _find_compaction_boundary(messages)
|
||||
compacted_messages = messages[:boundary]
|
||||
recent_messages = messages[boundary:]
|
||||
if not compacted_messages or not recent_messages:
|
||||
return messages, previous_summary, False
|
||||
|
||||
event_emitter = None
|
||||
if metadata.get('chat_id') and metadata.get('message_id'):
|
||||
from open_webui.socket.main import get_event_emitter
|
||||
|
||||
event_emitter = await get_event_emitter(metadata)
|
||||
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'context_compaction',
|
||||
'data': {
|
||||
'action': 'context_compaction',
|
||||
'description': 'Compacting context',
|
||||
'done': False,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
summary = await _generate_summary(
|
||||
request,
|
||||
user,
|
||||
model_id,
|
||||
models,
|
||||
compacted_messages,
|
||||
recent_messages,
|
||||
previous_summary,
|
||||
config['prompt_template'],
|
||||
)
|
||||
except Exception:
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'context_compaction',
|
||||
'data': {
|
||||
'action': 'context_compaction',
|
||||
'description': 'Context compaction failed',
|
||||
'done': True,
|
||||
'error': True,
|
||||
},
|
||||
}
|
||||
)
|
||||
raise
|
||||
|
||||
chat_id = metadata.get('chat_id')
|
||||
checkpoint_message_id = metadata.get('user_message_id') or metadata.get('message_id')
|
||||
if chat_id and checkpoint_message_id and not chat_id.startswith(('local:', 'channel:')):
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
||||
chat_id,
|
||||
checkpoint_message_id,
|
||||
{'contextSummary': summary},
|
||||
)
|
||||
|
||||
log.info(
|
||||
'Compacted chat context for chat=%s checkpoint=%s response=%s dropped=%d kept=%d summary_chars=%d',
|
||||
chat_id,
|
||||
checkpoint_message_id,
|
||||
metadata.get('message_id'),
|
||||
len(compacted_messages),
|
||||
len(recent_messages),
|
||||
len(summary),
|
||||
)
|
||||
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'context_compaction',
|
||||
'data': {
|
||||
'action': 'context_compaction',
|
||||
'description': 'Context compacted',
|
||||
'done': True,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
return recent_messages, summary, True
|
||||
|
||||
|
||||
async def compact_chat_branch(request, user, chat: Any, model_id: str, models: dict) -> dict:
|
||||
config = await _load_config()
|
||||
if not config['enable']:
|
||||
return {'ok': True, 'compacted': False, 'reason': 'disabled'}
|
||||
|
||||
history = (chat.chat or {}).get('history') or {}
|
||||
current_id = history.get('currentId')
|
||||
if not current_id:
|
||||
return {'ok': True, 'compacted': False, 'reason': 'empty'}
|
||||
|
||||
messages_map = await Chats.get_messages_map_by_chat_id(chat.id)
|
||||
if not messages_map:
|
||||
messages_map = history.get('messages') or {}
|
||||
|
||||
messages, previous_summary = _apply_latest_summary_checkpoint(get_message_list(messages_map, current_id))
|
||||
if len(messages) <= 2:
|
||||
return {'ok': True, 'compacted': False, 'reason': 'too_short'}
|
||||
|
||||
compacted_messages = messages[:-1]
|
||||
recent_messages = messages[-1:]
|
||||
summary = await _generate_summary(
|
||||
request,
|
||||
user,
|
||||
model_id,
|
||||
models,
|
||||
compacted_messages,
|
||||
recent_messages,
|
||||
previous_summary,
|
||||
config['prompt_template'],
|
||||
)
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(chat.id, current_id, {'contextSummary': summary})
|
||||
|
||||
return {
|
||||
'ok': True,
|
||||
'compacted': True,
|
||||
'dropped_messages': len(compacted_messages),
|
||||
'kept_messages': len(recent_messages),
|
||||
'summary_chars': len(summary),
|
||||
}
|
||||
|
||||
|
||||
async def _load_config() -> dict:
|
||||
values = await Config.get_many(
|
||||
'chat.context_compaction.enable',
|
||||
'chat.context_compaction.token_threshold',
|
||||
'chat.context_compaction.prompt_template',
|
||||
)
|
||||
return {
|
||||
'enable': bool(values.get('chat.context_compaction.enable', False)),
|
||||
'token_threshold': int(values.get('chat.context_compaction.token_threshold', 80000) or 80000),
|
||||
'prompt_template': values.get('chat.context_compaction.prompt_template', '') or '',
|
||||
}
|
||||
|
||||
|
||||
def _parse_positive_int(value: Any) -> int | None:
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return parsed if parsed > 0 else None
|
||||
|
||||
|
||||
def _resolve_token_threshold(global_threshold: int, metadata: dict) -> int:
|
||||
configured_threshold = _parse_positive_int((metadata.get('params') or {}).get('compact_token_threshold'))
|
||||
if configured_threshold is None:
|
||||
return global_threshold
|
||||
return min(configured_threshold, global_threshold)
|
||||
|
||||
|
||||
def _apply_latest_summary_checkpoint(messages: list[dict]) -> tuple[list[dict], str | None]:
|
||||
summary = None
|
||||
summary_idx = None
|
||||
|
||||
for idx, message in enumerate(messages):
|
||||
value = message.get('contextSummary') or message.get('context_summary')
|
||||
if isinstance(value, str) and value.strip():
|
||||
summary = value
|
||||
summary_idx = idx
|
||||
|
||||
if summary_idx is None:
|
||||
return messages, None
|
||||
return messages[summary_idx:], summary
|
||||
|
||||
|
||||
def _exceeds_token_threshold(messages: list[dict], system_prompt: str, summary: str | None, threshold: int) -> bool:
|
||||
if threshold <= 0:
|
||||
return False
|
||||
|
||||
for idx in range(len(messages) - 1, -1, -1):
|
||||
usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage')
|
||||
if isinstance(usage, dict) and usage.get('input_tokens'):
|
||||
total = int(usage.get('input_tokens') or 0) + int(usage.get('output_tokens') or 0)
|
||||
return total + _estimate_messages_tokens(messages[idx + 1 :]) > threshold
|
||||
|
||||
estimated = _estimate_tokens(system_prompt) + _estimate_tokens(summary or '') + _estimate_messages_tokens(messages)
|
||||
return estimated > threshold
|
||||
|
||||
|
||||
def _find_compaction_boundary(messages: list[dict]) -> int:
|
||||
keep_count = max(2, len(messages) * 2 // 5)
|
||||
split = max(1, len(messages) - keep_count)
|
||||
|
||||
while split < len(messages) - 1:
|
||||
previous = messages[split - 1] if split > 0 else {}
|
||||
current = messages[split]
|
||||
if current.get('role') == 'tool' or previous.get('tool_calls') or previous.get('output'):
|
||||
split += 1
|
||||
continue
|
||||
break
|
||||
|
||||
return min(split, len(messages) - 2)
|
||||
|
||||
|
||||
async def _generate_summary(
|
||||
request,
|
||||
user,
|
||||
model_id: str,
|
||||
models: dict,
|
||||
compacted_messages: list[dict],
|
||||
recent_messages: list[dict],
|
||||
previous_summary: str | None,
|
||||
summary_prompt_template: str,
|
||||
) -> str:
|
||||
from open_webui.utils.chat import generate_chat_completion
|
||||
|
||||
task_model_id = get_task_model_id(
|
||||
model_id,
|
||||
await Config.get('task.model.default'),
|
||||
await Config.get('task.model.external'),
|
||||
models,
|
||||
)
|
||||
if task_model_id not in models:
|
||||
task_model_id = model_id
|
||||
if task_model_id not in models:
|
||||
raise ValueError('No available model for context compaction')
|
||||
|
||||
summary_prompt_template = summary_prompt_template.strip() or DEFAULT_CONTEXT_COMPACTION_PROMPT
|
||||
all_messages = [*compacted_messages, *recent_messages]
|
||||
prompt = replace_prompt_variable(summary_prompt_template, get_last_user_message(all_messages) or '')
|
||||
prompt = replace_messages_variable(prompt, all_messages)
|
||||
prompt = replace_messages_variable(prompt, compacted_messages, 'COMPACTED_MESSAGES')
|
||||
prompt = replace_messages_variable(prompt, recent_messages, 'RECENT_MESSAGES')
|
||||
prompt = prompt_variables_template(prompt, {'{{PREVIOUS_SUMMARY}}': previous_summary or ''})
|
||||
prompt = await prompt_template(prompt, user)
|
||||
|
||||
max_tokens = models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000)
|
||||
payload = {
|
||||
'model': task_model_id,
|
||||
'messages': [{'role': 'user', 'content': prompt}],
|
||||
'stream': False,
|
||||
**(
|
||||
{'max_tokens': max_tokens}
|
||||
if models[task_model_id].get('owned_by') == 'ollama'
|
||||
else {'max_completion_tokens': max_tokens}
|
||||
),
|
||||
'metadata': {
|
||||
**(request.state.metadata if hasattr(request.state, 'metadata') else {}),
|
||||
'task': 'context_compaction',
|
||||
},
|
||||
}
|
||||
|
||||
response = await generate_chat_completion(request, form_data=payload, user=user)
|
||||
summary = _response_text(response).strip()
|
||||
if summary:
|
||||
return summary
|
||||
|
||||
parts = [previous_summary] if previous_summary else []
|
||||
for message in compacted_messages:
|
||||
content = get_content_from_message(message)
|
||||
if content:
|
||||
parts.append(f'- {message.get("role", "unknown")}: {content[:500]}')
|
||||
return '\n'.join(parts)[:4000]
|
||||
|
||||
|
||||
def _response_text(response: Any) -> str:
|
||||
if isinstance(response, list) and len(response) == 1:
|
||||
response = response[0]
|
||||
|
||||
if isinstance(response, JSONResponse):
|
||||
try:
|
||||
response = json.loads(response.body.decode('utf-8', 'replace'))
|
||||
except Exception:
|
||||
return ''
|
||||
|
||||
if not isinstance(response, dict):
|
||||
return ''
|
||||
|
||||
choices = response.get('choices') or []
|
||||
if choices:
|
||||
message = choices[0].get('message') or {}
|
||||
return message.get('content') or message.get('reasoning_content') or ''
|
||||
|
||||
parts = []
|
||||
for item in response.get('output') or []:
|
||||
for content in item.get('content') or []:
|
||||
if isinstance(content, dict):
|
||||
parts.append(content.get('text') or content.get('content') or '')
|
||||
return '\n'.join(part for part in parts if part)
|
||||
|
||||
|
||||
def _estimate_messages_tokens(messages: list[dict]) -> int:
|
||||
total = 0
|
||||
for message in messages:
|
||||
total += 4
|
||||
content = message.get('content')
|
||||
if isinstance(content, list):
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
total += _estimate_tokens(item)
|
||||
elif item.get('type') in {'image', 'image_url'}:
|
||||
total += 1000
|
||||
else:
|
||||
total += _estimate_tokens(item.get('text') or item.get('content') or item)
|
||||
else:
|
||||
total += _estimate_tokens(content)
|
||||
|
||||
total += _estimate_tokens(message.get('output'))
|
||||
total += _estimate_tokens(message.get('tool_calls'))
|
||||
total += _estimate_tokens(message.get('files'))
|
||||
return total
|
||||
|
||||
|
||||
def _estimate_tokens(value: Any) -> int:
|
||||
if value is None:
|
||||
return 0
|
||||
|
||||
if not isinstance(value, str):
|
||||
try:
|
||||
value = json.dumps(value, ensure_ascii=False)
|
||||
except Exception:
|
||||
value = str(value)
|
||||
|
||||
if not value:
|
||||
return 0
|
||||
|
||||
return max(1, len(value) // 4)
|
||||
@@ -20,7 +20,7 @@ from open_webui.env import (
|
||||
)
|
||||
from open_webui.models.chats import Chats
|
||||
from open_webui.models.files import Files
|
||||
from open_webui.retrieval.web.utils import validate_url
|
||||
from open_webui.retrieval.web.utils import get_ssrf_safe_session, validate_url
|
||||
from open_webui.routers.files import upload_file_handler
|
||||
from open_webui.utils.access_control.files import has_access_to_file
|
||||
from open_webui.routers.images import (
|
||||
@@ -28,7 +28,6 @@ from open_webui.routers.images import (
|
||||
upload_image,
|
||||
)
|
||||
from open_webui.storage.provider import Storage
|
||||
from open_webui.utils.session_pool import get_session
|
||||
|
||||
BASE64_IMAGE_URL_PREFIX = re.compile(r'data:image/\w+;base64,', re.IGNORECASE)
|
||||
MARKDOWN_IMAGE_URL_PATTERN = re.compile(r'!\[(.*?)\]\((.+?)\)', re.IGNORECASE)
|
||||
@@ -59,17 +58,18 @@ async def get_image_base64_from_url(url: str, user=None) -> Optional[str]:
|
||||
# called only on the originally-submitted URL; following 3xx redirects
|
||||
# without re-validation would let an attacker reach private IPs via a
|
||||
# public host that redirects internally (e.g. cloud-metadata exfil).
|
||||
validate_url(url)
|
||||
# Download the image from the URL
|
||||
session = await get_session()
|
||||
async with session.get(
|
||||
url, ssl=AIOHTTP_CLIENT_SESSION_SSL, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
image_data = await response.read()
|
||||
encoded_string = base64.b64encode(image_data).decode('utf-8')
|
||||
content_type = response.headers.get('Content-Type', 'image/png')
|
||||
return f'data:{content_type};base64,{encoded_string}'
|
||||
await asyncio.to_thread(validate_url, url)
|
||||
# Fetch through an SSRF-safe session that re-checks the connect-time IP, so a
|
||||
# rebinding DNS answer that passed validate_url cannot reach an internal address.
|
||||
async with get_ssrf_safe_session() as session:
|
||||
async with session.get(
|
||||
url, ssl=AIOHTTP_CLIENT_SESSION_SSL, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
image_data = await response.read()
|
||||
encoded_string = base64.b64encode(image_data).decode('utf-8')
|
||||
content_type = response.headers.get('Content-Type', 'image/png')
|
||||
return f'data:{content_type};base64,{encoded_string}'
|
||||
else:
|
||||
# Non-URL string — treat as file_id. Delegate to the canonical
|
||||
# file-ID resolver which enforces ownership/access checks.
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user