mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-13 09:20:14 -03:00
Compare commits
20 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8e45c22d7a | |||
| 191c4e03cd | |||
| ab4154c57d | |||
| 28e93d12ff | |||
| 75e63c758b | |||
| 823f71f269 | |||
| 042dd4088d | |||
| eaa791a9eb | |||
| 2228627ff4 | |||
| 4c647ad9c8 | |||
| 8ca3e6c33f | |||
| dd6bdbf297 | |||
| b47dde87e4 | |||
| 99e65cccd8 | |||
| 3bdacb8f46 | |||
| b4f9c224d3 | |||
| 5ec0399c81 | |||
| b464fdc333 | |||
| 53825500db | |||
| f2ac790752 |
+313
-291
File diff suppressed because it is too large
Load Diff
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "Lokale Versionen",
|
"versionsCount": "Lokale Versionen",
|
||||||
"versionsCountDesc": "Meiste Versionen zuerst",
|
"versionsCountDesc": "Meiste Versionen zuerst",
|
||||||
"versionsCountAsc": "Wenigste Versionen zuerst",
|
"versionsCountAsc": "Wenigste Versionen zuerst",
|
||||||
"versionIdDesc": "Neueste Version zuerst"
|
"versionIdDesc": "Neueste Version zuerst",
|
||||||
|
"random": "Zufällig",
|
||||||
|
"randomAction": "Zufällig mischen"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Modelliste aktualisieren",
|
"title": "Modelliste aktualisieren",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "Zu Nightly-Kanal wechseln",
|
"nightlyTitle": "Zu Nightly-Kanal wechseln",
|
||||||
"nightlyMessage": "Der Wechsel zu Nightly initialisiert ein Git-Repository und verfolgt die neuesten Commits des main-Branches. Updates sind häufiger, können aber instabil sein. Sie können jederzeit zu Release zurückwechseln.",
|
"nightlyMessage": "Der Wechsel zu Nightly initialisiert ein Git-Repository und verfolgt die neuesten Commits des main-Branches. Updates sind häufiger, können aber instabil sein. Sie können jederzeit zu Release zurückwechseln.",
|
||||||
"releaseTitle": "Zu Release-Kanal wechseln",
|
"releaseTitle": "Zu Release-Kanal wechseln",
|
||||||
"releaseMessage": "Der Wechsel zu Release entfernt das Git-Repository und installiert die neueste stabile Version. Zukünftige Updates verwenden nur stabile Versionen.",
|
"releaseMessage": "Der Wechsel zu Release checkt den neuesten stabilen Versions-Tag aus. Sie können jederzeit zu Nightly zurückwechseln.",
|
||||||
"switching": "Wechsle zu {channel}-Kanal...",
|
"switching": "Wechsle zu {channel}-Kanal...",
|
||||||
"completed": "Erfolgreich zu {channel}-Kanal gewechselt",
|
"completed": "Erfolgreich zu {channel}-Kanal gewechselt",
|
||||||
"failed": "Kanalwechsel fehlgeschlagen"
|
"failed": "Kanalwechsel fehlgeschlagen"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "Local Versions",
|
"versionsCount": "Local Versions",
|
||||||
"versionsCountDesc": "Most versions first",
|
"versionsCountDesc": "Most versions first",
|
||||||
"versionsCountAsc": "Fewest versions first",
|
"versionsCountAsc": "Fewest versions first",
|
||||||
"versionIdDesc": "Newest version first"
|
"versionIdDesc": "Newest version first",
|
||||||
|
"random": "Random",
|
||||||
|
"randomAction": "Randomize (shuffle)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Refresh model list",
|
"title": "Refresh model list",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "Switch to Nightly Channel",
|
"nightlyTitle": "Switch to Nightly Channel",
|
||||||
"nightlyMessage": "Switching to Nightly will initialize a Git repository and track the latest main branch commits. Updates will be more frequent but may be unstable. You can switch back to Release at any time.",
|
"nightlyMessage": "Switching to Nightly will initialize a Git repository and track the latest main branch commits. Updates will be more frequent but may be unstable. You can switch back to Release at any time.",
|
||||||
"releaseTitle": "Switch to Release Channel",
|
"releaseTitle": "Switch to Release Channel",
|
||||||
"releaseMessage": "Switching to Release will remove the Git repository and install the latest stable release. Future updates will use stable releases only.",
|
"releaseMessage": "Switching to Release will checkout the latest stable release tag. You can switch back to Nightly at any time.",
|
||||||
"switching": "Switching to {channel} channel...",
|
"switching": "Switching to {channel} channel...",
|
||||||
"completed": "Successfully switched to {channel} channel",
|
"completed": "Successfully switched to {channel} channel",
|
||||||
"failed": "Failed to switch channel"
|
"failed": "Failed to switch channel"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "Versiones locales",
|
"versionsCount": "Versiones locales",
|
||||||
"versionsCountDesc": "Más versiones primero",
|
"versionsCountDesc": "Más versiones primero",
|
||||||
"versionsCountAsc": "Menos versiones primero",
|
"versionsCountAsc": "Menos versiones primero",
|
||||||
"versionIdDesc": "Versión más nueva primero"
|
"versionIdDesc": "Versión más nueva primero",
|
||||||
|
"random": "Aleatorio",
|
||||||
|
"randomAction": "Aleatorizar (barajar)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Actualizar lista de modelos",
|
"title": "Actualizar lista de modelos",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "Cambiar a canal Nightly",
|
"nightlyTitle": "Cambiar a canal Nightly",
|
||||||
"nightlyMessage": "Cambiar a Nightly inicializara un repositorio Git y seguira los ultimos commits de la rama main. Las actualizaciones son mas frecuentes pero pueden ser inestables. Puede volver a Release en cualquier momento.",
|
"nightlyMessage": "Cambiar a Nightly inicializara un repositorio Git y seguira los ultimos commits de la rama main. Las actualizaciones son mas frecuentes pero pueden ser inestables. Puede volver a Release en cualquier momento.",
|
||||||
"releaseTitle": "Cambiar a canal Release",
|
"releaseTitle": "Cambiar a canal Release",
|
||||||
"releaseMessage": "Cambiar a Release eliminara el repositorio Git e instalara la ultima version estable. Las futuras actualizaciones usaran solo versiones estables.",
|
"releaseMessage": "Cambiar a Release hara checkout de la ultima etiqueta de version estable. Puede volver a Nightly en cualquier momento.",
|
||||||
"switching": "Cambiando a canal {channel}...",
|
"switching": "Cambiando a canal {channel}...",
|
||||||
"completed": "Cambio a canal {channel} exitoso",
|
"completed": "Cambio a canal {channel} exitoso",
|
||||||
"failed": "Error al cambiar de canal"
|
"failed": "Error al cambiar de canal"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "Versions locales",
|
"versionsCount": "Versions locales",
|
||||||
"versionsCountDesc": "Plus de versions d'abord",
|
"versionsCountDesc": "Plus de versions d'abord",
|
||||||
"versionsCountAsc": "Moins de versions d'abord",
|
"versionsCountAsc": "Moins de versions d'abord",
|
||||||
"versionIdDesc": "Version la plus récente d'abord"
|
"versionIdDesc": "Version la plus récente d'abord",
|
||||||
|
"random": "Aléatoire",
|
||||||
|
"randomAction": "Aléatoire (mélanger)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Actualiser la liste des modèles",
|
"title": "Actualiser la liste des modèles",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "Passer au canal Nightly",
|
"nightlyTitle": "Passer au canal Nightly",
|
||||||
"nightlyMessage": "Passer a Nightly initialisera un depot Git et suivra les derniers commits de la branche main. Les mises a jour sont plus frequentes mais peuvent etre instables. Vous pouvez revenir a Release a tout moment.",
|
"nightlyMessage": "Passer a Nightly initialisera un depot Git et suivra les derniers commits de la branche main. Les mises a jour sont plus frequentes mais peuvent etre instables. Vous pouvez revenir a Release a tout moment.",
|
||||||
"releaseTitle": "Passer au canal Release",
|
"releaseTitle": "Passer au canal Release",
|
||||||
"releaseMessage": "Passer a Release supprimera le depot Git et installera la derniere version stable. Les futures mises a jour utiliseront uniquement des versions stables.",
|
"releaseMessage": "Passer a Release passera au dernier tag de version stable. Vous pouvez revenir a Nightly a tout moment.",
|
||||||
"switching": "Passage au canal {channel}...",
|
"switching": "Passage au canal {channel}...",
|
||||||
"completed": "Basculement vers le canal {channel} reussi",
|
"completed": "Basculement vers le canal {channel} reussi",
|
||||||
"failed": "Echec du changement de canal"
|
"failed": "Echec du changement de canal"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "גרסאות מקומיות",
|
"versionsCount": "גרסאות מקומיות",
|
||||||
"versionsCountDesc": "הכי הרבה גרסאות ראשונות",
|
"versionsCountDesc": "הכי הרבה גרסאות ראשונות",
|
||||||
"versionsCountAsc": "הכי מעט גרסאות ראשונות",
|
"versionsCountAsc": "הכי מעט גרסאות ראשונות",
|
||||||
"versionIdDesc": "גרסה חדשה ביותר ראשונה"
|
"versionIdDesc": "גרסה חדשה ביותר ראשונה",
|
||||||
|
"random": "אקראי",
|
||||||
|
"randomAction": "ערבוב אקראי"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "רענן רשימת מודלים",
|
"title": "רענן רשימת מודלים",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "מעבר לערוץ Nightly",
|
"nightlyTitle": "מעבר לערוץ Nightly",
|
||||||
"nightlyMessage": "מעבר ל-Nightly יאתחל מאגר Git ויעקוב אחר הקומיטים האחרונים בענף main. העדכונים תכופים יותר אך עשויים להיות לא יציבים. ניתן לחזור ל-Release בכל עת.",
|
"nightlyMessage": "מעבר ל-Nightly יאתחל מאגר Git ויעקוב אחר הקומיטים האחרונים בענף main. העדכונים תכופים יותר אך עשויים להיות לא יציבים. ניתן לחזור ל-Release בכל עת.",
|
||||||
"releaseTitle": "מעבר לערוץ Release",
|
"releaseTitle": "מעבר לערוץ Release",
|
||||||
"releaseMessage": "מעבר ל-Release יסיר את מאגר ה-Git ויתקין את הגרסה היציבה האחרונה. עדכונים עתידיים ישתמשו בגרסאות יציבות בלבד.",
|
"releaseMessage": "מעבר ל-Release יעבור לתגית הגרסה היציבה האחרונה. ניתן לחזור ל-Nightly בכל עת.",
|
||||||
"switching": "מעבר לערוץ {channel}...",
|
"switching": "מעבר לערוץ {channel}...",
|
||||||
"completed": "המעבר לערוץ {channel} הושלם",
|
"completed": "המעבר לערוץ {channel} הושלם",
|
||||||
"failed": "החלפת ערוץ נכשלה"
|
"failed": "החלפת ערוץ נכשלה"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "ローカルバージョン数",
|
"versionsCount": "ローカルバージョン数",
|
||||||
"versionsCountDesc": "バージョン数の多い順",
|
"versionsCountDesc": "バージョン数の多い順",
|
||||||
"versionsCountAsc": "バージョン数の少ない順",
|
"versionsCountAsc": "バージョン数の少ない順",
|
||||||
"versionIdDesc": "最新バージョン順"
|
"versionIdDesc": "最新バージョン順",
|
||||||
|
"random": "ランダム",
|
||||||
|
"randomAction": "シャッフル(ランダム)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "モデルリストを更新",
|
"title": "モデルリストを更新",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "ナイトリーチャンネルに切り替え",
|
"nightlyTitle": "ナイトリーチャンネルに切り替え",
|
||||||
"nightlyMessage": "ナイトリーに切り替えると、Gitリポジトリが初期化され、mainブランチの最新コミットを追跡します。更新頻度は高くなりますが、不安定な場合があります。いつでもリリース版に戻せます。",
|
"nightlyMessage": "ナイトリーに切り替えると、Gitリポジトリが初期化され、mainブランチの最新コミットを追跡します。更新頻度は高くなりますが、不安定な場合があります。いつでもリリース版に戻せます。",
|
||||||
"releaseTitle": "リリースチャンネルに切り替え",
|
"releaseTitle": "リリースチャンネルに切り替え",
|
||||||
"releaseMessage": "リリースに切り替えると、Gitリポジトリが削除され、最新の安定版がインストールされます。以降の更新は安定版のみが使用されます。",
|
"releaseMessage": "リリースに切り替えると、最新の安定版タグにチェックアウトされます。いつでもNightlyに戻せます。",
|
||||||
"switching": "{channel} チャンネルに切り替え中...",
|
"switching": "{channel} チャンネルに切り替え中...",
|
||||||
"completed": "{channel} チャンネルに切り替えました",
|
"completed": "{channel} チャンネルに切り替えました",
|
||||||
"failed": "チャンネルの切り替えに失敗しました"
|
"failed": "チャンネルの切り替えに失敗しました"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "로컬 버전 수",
|
"versionsCount": "로컬 버전 수",
|
||||||
"versionsCountDesc": "버전 수 많은 순",
|
"versionsCountDesc": "버전 수 많은 순",
|
||||||
"versionsCountAsc": "버전 수 적은 순",
|
"versionsCountAsc": "버전 수 적은 순",
|
||||||
"versionIdDesc": "최신 버전순"
|
"versionIdDesc": "최신 버전순",
|
||||||
|
"random": "랜덤",
|
||||||
|
"randomAction": "셔플 (무작위)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "모델 목록 새로고침",
|
"title": "모델 목록 새로고침",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "나이틀리 채널로 전환",
|
"nightlyTitle": "나이틀리 채널로 전환",
|
||||||
"nightlyMessage": "나이틀리로 전환하면 Git 저장소가 초기화되고 main 브랜치의 최신 커밋을 추적합니다. 업데이트 빈도는 높지만 불안정할 수 있습니다. 언제든지 릴리스로 돌아갈 수 있습니다.",
|
"nightlyMessage": "나이틀리로 전환하면 Git 저장소가 초기화되고 main 브랜치의 최신 커밋을 추적합니다. 업데이트 빈도는 높지만 불안정할 수 있습니다. 언제든지 릴리스로 돌아갈 수 있습니다.",
|
||||||
"releaseTitle": "릴리스 채널로 전환",
|
"releaseTitle": "릴리스 채널로 전환",
|
||||||
"releaseMessage": "릴리스로 전환하면 Git 저장소가 제거되고 최신 안정 버전이 설치됩니다. 이후 업데이트는 안정 버전만 사용됩니다.",
|
"releaseMessage": "릴리스로 전환하면 최신 안정 버전 태그로 체크아웃됩니다. 언제든지 나이틀리로 돌아갈 수 있습니다.",
|
||||||
"switching": "{channel} 채널로 전환 중...",
|
"switching": "{channel} 채널로 전환 중...",
|
||||||
"completed": "{channel} 채널로 전환 완료",
|
"completed": "{channel} 채널로 전환 완료",
|
||||||
"failed": "채널 전환 실패"
|
"failed": "채널 전환 실패"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "Локальные версии",
|
"versionsCount": "Локальные версии",
|
||||||
"versionsCountDesc": "Сначала больше версий",
|
"versionsCountDesc": "Сначала больше версий",
|
||||||
"versionsCountAsc": "Сначала меньше версий",
|
"versionsCountAsc": "Сначала меньше версий",
|
||||||
"versionIdDesc": "Сначала новые версии"
|
"versionIdDesc": "Сначала новые версии",
|
||||||
|
"random": "Случайно",
|
||||||
|
"randomAction": "Перемешать"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Обновить список моделей",
|
"title": "Обновить список моделей",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "Переключиться на Nightly",
|
"nightlyTitle": "Переключиться на Nightly",
|
||||||
"nightlyMessage": "Переключение на Nightly инициализирует Git-репозиторий и отслеживает последние коммиты ветки main. Обновления чаще, но могут быть нестабильными. Вы можете вернуться к Release в любое время.",
|
"nightlyMessage": "Переключение на Nightly инициализирует Git-репозиторий и отслеживает последние коммиты ветки main. Обновления чаще, но могут быть нестабильными. Вы можете вернуться к Release в любое время.",
|
||||||
"releaseTitle": "Переключиться на Release",
|
"releaseTitle": "Переключиться на Release",
|
||||||
"releaseMessage": "Переключение на Release удалит Git-репозиторий и установит последнюю стабильную версию. Будущие обновления будут использовать только стабильные версии.",
|
"releaseMessage": "Переключение на Release выполнит checkout последнего стабильного тега. Вы можете вернуться к Nightly в любое время.",
|
||||||
"switching": "Переключение на канал {channel}...",
|
"switching": "Переключение на канал {channel}...",
|
||||||
"completed": "Успешно переключено на канал {channel}",
|
"completed": "Успешно переключено на канал {channel}",
|
||||||
"failed": "Не удалось переключить канал"
|
"failed": "Не удалось переключить канал"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "本地版本数",
|
"versionsCount": "本地版本数",
|
||||||
"versionsCountDesc": "版本数从多到少",
|
"versionsCountDesc": "版本数从多到少",
|
||||||
"versionsCountAsc": "版本数从少到多",
|
"versionsCountAsc": "版本数从少到多",
|
||||||
"versionIdDesc": "最新版本优先"
|
"versionIdDesc": "最新版本优先",
|
||||||
|
"random": "随机",
|
||||||
|
"randomAction": "随机排序(洗牌)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "刷新模型列表",
|
"title": "刷新模型列表",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "切换到 Nightly",
|
"nightlyTitle": "切换到 Nightly",
|
||||||
"nightlyMessage": "切换到 Nightly 将初始化 Git 仓库并跟踪 main 分支的最新提交。更新更频繁但可能不稳定,可随时切回稳定版。",
|
"nightlyMessage": "切换到 Nightly 将初始化 Git 仓库并跟踪 main 分支的最新提交。更新更频繁但可能不稳定,可随时切回稳定版。",
|
||||||
"releaseTitle": "切换到稳定版",
|
"releaseTitle": "切换到稳定版",
|
||||||
"releaseMessage": "切换到稳定版将移除 Git 仓库并安装最新的稳定发布版本,后续仅使用稳定版更新。",
|
"releaseMessage": "切换到稳定版将检出最新的发布标签。可随时切换回每日构建版。",
|
||||||
"switching": "正在切换到 {channel} 频道...",
|
"switching": "正在切换到 {channel} 频道...",
|
||||||
"completed": "已切换到 {channel} 频道",
|
"completed": "已切换到 {channel} 频道",
|
||||||
"failed": "切换频道失败"
|
"failed": "切换频道失败"
|
||||||
|
|||||||
+4
-2
@@ -714,7 +714,9 @@
|
|||||||
"versionsCount": "本地版本數",
|
"versionsCount": "本地版本數",
|
||||||
"versionsCountDesc": "版本數從多到少",
|
"versionsCountDesc": "版本數從多到少",
|
||||||
"versionsCountAsc": "版本數從少到多",
|
"versionsCountAsc": "版本數從少到多",
|
||||||
"versionIdDesc": "最新版本優先"
|
"versionIdDesc": "最新版本優先",
|
||||||
|
"random": "隨機",
|
||||||
|
"randomAction": "隨機排序(洗牌)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "重新整理模型列表",
|
"title": "重新整理模型列表",
|
||||||
@@ -1782,7 +1784,7 @@
|
|||||||
"nightlyTitle": "切换到 Nightly",
|
"nightlyTitle": "切换到 Nightly",
|
||||||
"nightlyMessage": "切换到 Nightly 将初始化 Git 仓库并跟踪 main 分支的最新提交。更新更频繁但可能不稳定,可随时切回稳定版。",
|
"nightlyMessage": "切换到 Nightly 将初始化 Git 仓库并跟踪 main 分支的最新提交。更新更频繁但可能不稳定,可随时切回稳定版。",
|
||||||
"releaseTitle": "切换到稳定版",
|
"releaseTitle": "切换到稳定版",
|
||||||
"releaseMessage": "切换到稳定版将移除 Git 仓库并安装最新的稳定发布版本,后续仅使用稳定版更新。",
|
"releaseMessage": "切換到穩定版將檢出最新的發布標籤。可隨時切換回每日構建版。",
|
||||||
"switching": "正在切換到 {channel} 頻道...",
|
"switching": "正在切換到 {channel} 頻道...",
|
||||||
"completed": "已切換到 {channel} 頻道",
|
"completed": "已切換到 {channel} 頻道",
|
||||||
"failed": "切換頻道失敗"
|
"failed": "切換頻道失敗"
|
||||||
|
|||||||
@@ -2,7 +2,8 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
|
||||||
from .constants import CLIP_SKIP_SENTINEL, MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE, METADATA_OVERWRITE_FIELDS
|
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
|
||||||
|
from .overwrite_utils import collect_overwrite_params
|
||||||
|
|
||||||
|
|
||||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||||
@@ -1233,14 +1234,7 @@ class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
|||||||
if not inputs:
|
if not inputs:
|
||||||
return
|
return
|
||||||
|
|
||||||
overwrite_params = {}
|
overwrite_params = collect_overwrite_params(inputs)
|
||||||
for key in METADATA_OVERWRITE_FIELDS:
|
|
||||||
value = inputs.get(key)
|
|
||||||
if key == "clip_skip":
|
|
||||||
if value != CLIP_SKIP_SENTINEL:
|
|
||||||
overwrite_params[key] = value
|
|
||||||
elif value: # truthy — only overwrite when user provided a real value
|
|
||||||
overwrite_params[key] = value
|
|
||||||
|
|
||||||
if overwrite_params:
|
if overwrite_params:
|
||||||
metadata.setdefault(OVERWRITE, {})
|
metadata.setdefault(OVERWRITE, {})
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
"""Shared helpers for Metadata Overwrite node metadata collection.
|
||||||
|
|
||||||
|
Used by both the MetadataOverwriteLM node (execution time) and the
|
||||||
|
MetadataOverwriteExtractor (hook time) so the conversion/filtering logic
|
||||||
|
cannot drift between the two paths.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
|
from ..utils.utils import model_patcher_to_name
|
||||||
|
from .constants import CLIP_SKIP_SENTINEL, METADATA_OVERWRITE_FIELDS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def collect_overwrite_params(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Convert node input values into non-default overwrite parameters.
|
||||||
|
|
||||||
|
For most fields, a falsy value (empty string, 0) means "not set" and is
|
||||||
|
skipped. clip_skip uses a dedicated sentinel (-25) so that a wired value
|
||||||
|
of 0 is preserved. The ``model`` field accepts either a manual string or
|
||||||
|
a wired MODEL (ModelPatcher) connection; in the latter case the source
|
||||||
|
model name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path.
|
||||||
|
"""
|
||||||
|
result: Dict[str, Any] = {}
|
||||||
|
for key in METADATA_OVERWRITE_FIELDS:
|
||||||
|
value = values.get(key)
|
||||||
|
if key == "model" and not isinstance(value, str):
|
||||||
|
value = model_patcher_to_name(value)
|
||||||
|
if value is None:
|
||||||
|
logger.warning(
|
||||||
|
"Could not extract model name from wired MODEL input "
|
||||||
|
"(no cached_patcher_init); model metadata overwrite skipped"
|
||||||
|
)
|
||||||
|
if key == "clip_skip":
|
||||||
|
if value != CLIP_SKIP_SENTINEL:
|
||||||
|
result[key] = value
|
||||||
|
elif value:
|
||||||
|
result[key] = value
|
||||||
|
return result
|
||||||
@@ -1,26 +1,102 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
|
||||||
|
|
||||||
|
|
||||||
|
def _is_stack_input(name: str) -> bool:
|
||||||
|
return bool(_STACK_INPUT_PATTERN.match(name))
|
||||||
|
|
||||||
|
|
||||||
|
def _stack_slot_number(name: str) -> int:
|
||||||
|
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
|
||||||
|
match = _STACK_INPUT_PATTERN.match(name)
|
||||||
|
if not match:
|
||||||
|
return -1
|
||||||
|
letter, digits = match.group(1), match.group(2)
|
||||||
|
if digits is not None:
|
||||||
|
return int(digits)
|
||||||
|
return 1 if letter == "a" else 2
|
||||||
|
|
||||||
|
|
||||||
|
class _LoraStackOptionalInputs:
|
||||||
|
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
|
||||||
|
|
||||||
|
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
|
||||||
|
self._explicit_inputs = explicit_inputs
|
||||||
|
|
||||||
|
def __contains__(self, item: object) -> bool:
|
||||||
|
if not isinstance(item, str):
|
||||||
|
return False
|
||||||
|
return item in self._explicit_inputs or _is_stack_input(item)
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
|
||||||
|
if key in self._explicit_inputs:
|
||||||
|
return self._explicit_inputs[key]
|
||||||
|
if _is_stack_input(key):
|
||||||
|
return (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
raise KeyError(key)
|
||||||
|
|
||||||
|
|
||||||
class LoraStackCombinerLM:
|
class LoraStackCombinerLM:
|
||||||
NAME = "Lora Stack Combiner (LoraManager)"
|
NAME = "Lora Stack Combiner (LoraManager)"
|
||||||
CATEGORY = "Lora Manager/stackers"
|
CATEGORY = "Lora Manager/stackers"
|
||||||
|
DESCRIPTION = (
|
||||||
|
"Combines multiple LoRA stacks into a single stack. "
|
||||||
|
"Supports dynamic inputs: connect a stack to add more inputs."
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
|
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
|
||||||
|
"lora_stack1": (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"lora_stack2": (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
stack = inspect.stack()
|
||||||
|
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||||
|
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment]
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {},
|
||||||
"lora_stack_a": ("LORA_STACK",),
|
"optional": optional_inputs,
|
||||||
"lora_stack_b": ("LORA_STACK",),
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("LORA_STACK",)
|
RETURN_TYPES = ("LORA_STACK",)
|
||||||
RETURN_NAMES = ("LORA_STACK",)
|
RETURN_NAMES = ("LORA_STACK",)
|
||||||
FUNCTION = "combine_stacks"
|
FUNCTION = "combine_stacks"
|
||||||
|
|
||||||
def combine_stacks(self, lora_stack_a, lora_stack_b):
|
def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
|
||||||
combined_stack = []
|
stacks = {
|
||||||
|
"lora_stack1": lora_stack1,
|
||||||
|
"lora_stack2": lora_stack2,
|
||||||
|
}
|
||||||
|
for key, value in kwargs.items():
|
||||||
|
if _is_stack_input(key) and value is not None:
|
||||||
|
stacks[key] = value
|
||||||
|
|
||||||
if lora_stack_a:
|
combined_stack = []
|
||||||
combined_stack.extend(lora_stack_a)
|
for key in sorted(stacks, key=_stack_slot_number):
|
||||||
if lora_stack_b:
|
stack = stacks[key]
|
||||||
combined_stack.extend(lora_stack_b)
|
if stack:
|
||||||
|
combined_stack.extend(stack)
|
||||||
|
|
||||||
return (combined_stack,)
|
return (combined_stack,)
|
||||||
|
|||||||
@@ -9,10 +9,8 @@ but users may wire 0 to express "no clip skip / default".
|
|||||||
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ..metadata_collector.constants import (
|
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
|
||||||
CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL,
|
from ..metadata_collector.overwrite_utils import collect_overwrite_params
|
||||||
METADATA_OVERWRITE_FIELDS,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class MetadataOverwriteLM:
|
class MetadataOverwriteLM:
|
||||||
@@ -87,12 +85,16 @@ class MetadataOverwriteLM:
|
|||||||
},
|
},
|
||||||
),
|
),
|
||||||
"model": (
|
"model": (
|
||||||
"STRING",
|
"STRING,MODEL",
|
||||||
{
|
{
|
||||||
"default": "",
|
"default": "",
|
||||||
|
"widgetType": "STRING",
|
||||||
"tooltip": (
|
"tooltip": (
|
||||||
"The checkpoint or diffusion model (UNet) used "
|
"The checkpoint or diffusion model (UNet) used "
|
||||||
"for generation. Only overwrites when non-empty."
|
"for generation. Fill in the name manually or "
|
||||||
|
"connect a MODEL output — the model name is then "
|
||||||
|
"extracted automatically. Only overwrites when "
|
||||||
|
"non-empty."
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
@@ -158,13 +160,10 @@ class MetadataOverwriteLM:
|
|||||||
For most fields, a falsy value (empty string, 0) means "not set"
|
For most fields, a falsy value (empty string, 0) means "not set"
|
||||||
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
||||||
a wired value of 0 is preserved and reaches the metadata pipeline.
|
a wired value of 0 is preserved and reaches the metadata pipeline.
|
||||||
|
|
||||||
|
The ``model`` field accepts either a manual string or a wired MODEL
|
||||||
|
(ModelPatcher) connection; in the latter case the underlying model
|
||||||
|
name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path.
|
||||||
"""
|
"""
|
||||||
result: dict[str, Any] = {}
|
return (collect_overwrite_params(kwargs),)
|
||||||
for key in METADATA_OVERWRITE_FIELDS:
|
|
||||||
value = kwargs.get(key)
|
|
||||||
if key == "clip_skip":
|
|
||||||
if value != _CLIP_SKIP_SENTINEL:
|
|
||||||
result[key] = value
|
|
||||||
elif value:
|
|
||||||
result[key] = value
|
|
||||||
return (result,)
|
|
||||||
|
|||||||
@@ -7,6 +7,21 @@ from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_c
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _reload_gguf_unet(
|
||||||
|
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
||||||
|
) -> object:
|
||||||
|
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
||||||
|
|
||||||
|
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
|
||||||
|
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||||
|
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||||
|
with core ComfyUI loaders.
|
||||||
|
"""
|
||||||
|
loader = UNETLoaderLM()
|
||||||
|
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
class UNETLoaderLM:
|
class UNETLoaderLM:
|
||||||
"""UNET Loader with support for extra folder paths
|
"""UNET Loader with support for extra folder paths
|
||||||
|
|
||||||
@@ -196,6 +211,12 @@ class UNETLoaderLM:
|
|||||||
# Wrap with GGUFModelPatcher
|
# Wrap with GGUFModelPatcher
|
||||||
model = GGUFModelPatcher.clone(model)
|
model = GGUFModelPatcher.clone(model)
|
||||||
|
|
||||||
|
# Register a reload factory so the MODEL carries its source path
|
||||||
|
# (cached_patcher_init) like core ComfyUI loaders do — required
|
||||||
|
# for model-name extraction downstream and for ModelPatcher
|
||||||
|
# deepclone/dynamic machinery.
|
||||||
|
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
||||||
|
|
||||||
return (model,)
|
return (model,)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -1562,6 +1562,11 @@ class SettingsHandler:
|
|||||||
{"success": False, "error": validation_error}
|
{"success": False, "error": validation_error}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if key == "update_channel" and value not in ("release", "nightly"):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
|
||||||
|
)
|
||||||
|
|
||||||
if value == "__DELETE__" and key in (
|
if value == "__DELETE__" and key in (
|
||||||
"proxy_username",
|
"proxy_username",
|
||||||
"proxy_password",
|
"proxy_password",
|
||||||
@@ -2585,6 +2590,8 @@ class ModelLibraryHandler:
|
|||||||
status=400,
|
status=400,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
cursor = request.query.get("cursor")
|
||||||
|
|
||||||
metadata_provider = await self._metadata_provider_factory()
|
metadata_provider = await self._metadata_provider_factory()
|
||||||
if not metadata_provider:
|
if not metadata_provider:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2593,7 +2600,7 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
models = await metadata_provider.get_user_models(username)
|
result = await metadata_provider.get_user_models(username, cursor)
|
||||||
except NotImplementedError:
|
except NotImplementedError:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
@@ -2603,14 +2610,35 @@ class ModelLibraryHandler:
|
|||||||
status=501,
|
status=501,
|
||||||
)
|
)
|
||||||
|
|
||||||
if models is None:
|
if result is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "Failed to fetch user models"},
|
{"success": False, "error": "Failed to fetch user models"},
|
||||||
status=502,
|
status=502,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if isinstance(result, dict):
|
||||||
|
models = result.get("items")
|
||||||
|
next_cursor = result.get("nextCursor")
|
||||||
|
else:
|
||||||
|
# Defensive: tolerate providers that still return a raw list
|
||||||
|
models = result
|
||||||
|
next_cursor = None
|
||||||
|
|
||||||
if not isinstance(models, list):
|
if not isinstance(models, list):
|
||||||
models = []
|
models = []
|
||||||
|
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||||
|
next_cursor = str(next_cursor)
|
||||||
|
|
||||||
|
estimated_total = None
|
||||||
|
if cursor is None:
|
||||||
|
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||||
|
if get_count is not None:
|
||||||
|
try:
|
||||||
|
estimated_total = await get_count(username)
|
||||||
|
except Exception: # best-effort only
|
||||||
|
estimated_total = None
|
||||||
|
if not isinstance(estimated_total, int):
|
||||||
|
estimated_total = None
|
||||||
|
|
||||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||||
@@ -2630,6 +2658,7 @@ class ModelLibraryHandler:
|
|||||||
versions: list[dict] = []
|
versions: list[dict] = []
|
||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
model_ids: list[int] = []
|
model_ids: list[int] = []
|
||||||
|
model_count = 0
|
||||||
for model in models:
|
for model in models:
|
||||||
try:
|
try:
|
||||||
model_ids.append(int(model.get("id")))
|
model_ids.append(int(model.get("id")))
|
||||||
@@ -2663,6 +2692,8 @@ class ModelLibraryHandler:
|
|||||||
if model_type not in normalized_allowed_types:
|
if model_type not in normalized_allowed_types:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
model_count += 1
|
||||||
|
|
||||||
scanner = type_scanner_map.get(model_type)
|
scanner = type_scanner_map.get(model_type)
|
||||||
if scanner is None:
|
if scanner is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2728,7 +2759,15 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": True, "username": username, "versions": versions}
|
{
|
||||||
|
"success": True,
|
||||||
|
"username": username,
|
||||||
|
"versions": versions,
|
||||||
|
"modelCount": model_count,
|
||||||
|
"nextCursor": next_cursor,
|
||||||
|
"hasMore": next_cursor is not None,
|
||||||
|
"estimatedTotal": estimated_total,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||||
|
|||||||
+127
-46
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
|
|||||||
return excludes
|
return excludes
|
||||||
|
|
||||||
|
|
||||||
|
def _stage_preserved_items(plugin_root: str) -> tuple[str, list[str]]:
|
||||||
|
"""Move preserved user-data items to a temp directory outside *plugin_root*.
|
||||||
|
|
||||||
|
This ensures that ``git reset --hard``, ``git clean -fd``, and ZIP-based
|
||||||
|
replacement cannot touch these files even when ``-e`` exclusion patterns
|
||||||
|
are mishandled (e.g. on Windows where forward-slash patterns may not
|
||||||
|
match backslash-prefixed paths in some Git builds, or where file locks
|
||||||
|
prevent deletion/recreation).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``(backup_root, staged_names)``: the temp directory path and the
|
||||||
|
list of item names that were successfully moved.
|
||||||
|
"""
|
||||||
|
backup_root = tempfile.mkdtemp(prefix='lora_manager_update_')
|
||||||
|
staged: list[str] = []
|
||||||
|
for name in _PRESERVE_DIRS:
|
||||||
|
src = os.path.join(plugin_root, name)
|
||||||
|
if not os.path.lexists(src):
|
||||||
|
continue
|
||||||
|
dst = os.path.join(backup_root, name)
|
||||||
|
try:
|
||||||
|
shutil.move(src, dst)
|
||||||
|
staged.append(name)
|
||||||
|
logger.debug("Staged '%s' for update safety", name)
|
||||||
|
except OSError:
|
||||||
|
# ``shutil.move`` may fail on Windows if a file handle inside
|
||||||
|
# the directory is still open (e.g. a SQLite WAL file). Fall
|
||||||
|
# back to copy-then-remove.
|
||||||
|
logger.debug("Move failed for '%s', falling back to copy", name)
|
||||||
|
try:
|
||||||
|
if os.path.isdir(src) and not os.path.islink(src):
|
||||||
|
shutil.copytree(src, dst, symlinks=True)
|
||||||
|
shutil.rmtree(src, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
shutil.copy2(src, dst)
|
||||||
|
os.remove(src)
|
||||||
|
staged.append(name)
|
||||||
|
logger.info("Copied (then removed) '%s' for update safety", name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Could not stage '%s': %s (will rely on git -e / skip lists)", name, exc
|
||||||
|
)
|
||||||
|
return backup_root, staged
|
||||||
|
|
||||||
|
|
||||||
|
def _restore_preserved_items(plugin_root: str, backup_root: str, staged: list[str]) -> None:
|
||||||
|
"""Move staged items back from *backup_root* into *plugin_root*.
|
||||||
|
|
||||||
|
Any leftover placeholder at the destination (created by git checkout or
|
||||||
|
ZIP extraction) is removed before the move.
|
||||||
|
"""
|
||||||
|
for name in staged:
|
||||||
|
src = os.path.join(backup_root, name)
|
||||||
|
dst = os.path.join(plugin_root, name)
|
||||||
|
try:
|
||||||
|
if os.path.lexists(dst):
|
||||||
|
if os.path.isdir(dst) and not os.path.islink(dst):
|
||||||
|
shutil.rmtree(dst, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
os.remove(dst)
|
||||||
|
shutil.move(src, dst)
|
||||||
|
logger.debug("Restored '%s' after update", name)
|
||||||
|
except OSError:
|
||||||
|
logger.debug("Move failed restoring '%s', falling back to copy", name)
|
||||||
|
try:
|
||||||
|
if os.path.isdir(src) and not os.path.islink(src):
|
||||||
|
shutil.copytree(src, dst, symlinks=True, dirs_exist_ok=True)
|
||||||
|
shutil.rmtree(src, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
shutil.copy2(src, dst)
|
||||||
|
os.remove(src)
|
||||||
|
logger.info("Copied '%s' back after update", name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to restore '%s': %s", name, exc)
|
||||||
|
shutil.rmtree(backup_root, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class UpdateRoutes:
|
class UpdateRoutes:
|
||||||
"""Routes for handling plugin update checks"""
|
"""Routes for handling plugin update checks"""
|
||||||
|
|
||||||
@@ -173,20 +251,22 @@ class UpdateRoutes:
|
|||||||
if os.path.exists(settings_path):
|
if os.path.exists(settings_path):
|
||||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||||
settings_backup = f.read()
|
settings_backup = f.read()
|
||||||
logger.info("Backed up settings.json")
|
logger.debug("Backed up settings.json (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
git_folder = os.path.join(plugin_root, '.git')
|
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||||
if os.path.exists(git_folder):
|
try:
|
||||||
# Git update
|
git_folder = os.path.join(plugin_root, '.git')
|
||||||
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
if os.path.exists(git_folder):
|
||||||
else:
|
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
||||||
# Fallback: Download ZIP and replace files
|
else:
|
||||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||||
|
finally:
|
||||||
|
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||||
|
|
||||||
if settings_backup and success:
|
if settings_backup and success:
|
||||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||||
f.write(settings_backup)
|
f.write(settings_backup)
|
||||||
logger.info("Restored settings.json")
|
logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
if success:
|
if success:
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -211,9 +291,11 @@ class UpdateRoutes:
|
|||||||
async def switch_channel(request):
|
async def switch_channel(request):
|
||||||
"""
|
"""
|
||||||
Switch between release and nightly update channels.
|
Switch between release and nightly update channels.
|
||||||
|
|
||||||
Release → Nightly: Initialize a Git repository (from ZIP/CM stable mode)
|
ZIP/CNR install → Nightly: git init + checkout main (one-way upgrade)
|
||||||
Nightly → Release: Remove .git, download latest release ZIP, write .tracking
|
Git install → Release: git checkout latest tag (.git preserved)
|
||||||
|
ZIP/CNR install → Release: ZIP download (no .git, stays in ZIP mode)
|
||||||
|
Git install → Nightly: git checkout main + pull
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
body = await request.json() if request.has_body else {}
|
body = await request.json() if request.has_body else {}
|
||||||
@@ -233,47 +315,47 @@ class UpdateRoutes:
|
|||||||
if os.path.exists(settings_path):
|
if os.path.exists(settings_path):
|
||||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||||
settings_backup = f.read()
|
settings_backup = f.read()
|
||||||
logger.info("Backed up settings.json before channel switch")
|
logger.debug("Backed up settings.json before channel switch (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
git_folder = os.path.join(plugin_root, '.git')
|
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||||
|
try:
|
||||||
|
git_folder = os.path.join(plugin_root, '.git')
|
||||||
|
|
||||||
if channel == 'nightly':
|
if channel == 'nightly':
|
||||||
git_backup = None
|
git_backup = None
|
||||||
if os.path.exists(git_folder):
|
if os.path.exists(git_folder):
|
||||||
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
|
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
|
||||||
|
|
||||||
success = False
|
success = False
|
||||||
new_version = ''
|
new_version = ''
|
||||||
try:
|
try:
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
success, new_version = await UpdateRoutes._perform_git_update(
|
||||||
|
plugin_root, nightly=True
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
|
||||||
|
finally:
|
||||||
|
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
|
||||||
|
else:
|
||||||
|
success = False
|
||||||
|
new_version = ''
|
||||||
if os.path.exists(git_folder):
|
if os.path.exists(git_folder):
|
||||||
success, new_version = await UpdateRoutes._perform_git_update(
|
success, new_version = await UpdateRoutes._perform_git_update(
|
||||||
plugin_root, nightly=True
|
plugin_root, nightly=False
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
|
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||||
finally:
|
if os.path.exists(tracking_file):
|
||||||
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
|
os.remove(tracking_file)
|
||||||
else:
|
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||||
git_backup = None
|
finally:
|
||||||
if os.path.exists(git_folder):
|
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||||
git_backup = UpdateRoutes._backup_git(git_folder, 'release')
|
|
||||||
|
|
||||||
success = False
|
|
||||||
new_version = ''
|
|
||||||
try:
|
|
||||||
if os.path.exists(git_folder):
|
|
||||||
shutil.rmtree(git_folder)
|
|
||||||
tracking_file = os.path.join(plugin_root, '.tracking')
|
|
||||||
if os.path.exists(tracking_file):
|
|
||||||
os.remove(tracking_file)
|
|
||||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
|
||||||
finally:
|
|
||||||
UpdateRoutes._restore_git(git_backup, git_folder, success, 'release')
|
|
||||||
|
|
||||||
if settings_backup and success:
|
if settings_backup and success:
|
||||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||||
f.write(settings_backup)
|
f.write(settings_backup)
|
||||||
logger.info("Restored settings.json after channel switch")
|
logger.debug("Restored settings.json content after channel switch (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
if success:
|
if success:
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -417,8 +499,7 @@ class UpdateRoutes:
|
|||||||
except Exception:
|
except Exception:
|
||||||
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
||||||
|
|
||||||
# Skip settings.json, civitai, model cache and runtime cache folders
|
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=list(_PRESERVE_DIRS))
|
||||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
|
|
||||||
|
|
||||||
# Extract ZIP to temp dir
|
# Extract ZIP to temp dir
|
||||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
@@ -428,7 +509,7 @@ class UpdateRoutes:
|
|||||||
extracted_root = next(os.scandir(tmp_dir)).path
|
extracted_root = next(os.scandir(tmp_dir)).path
|
||||||
|
|
||||||
# Copy files, skipping user data that should be preserved
|
# Copy files, skipping user data that should be preserved
|
||||||
skip_items = {'settings.json', 'civitai', 'wildcards', 'backups', 'stats'}
|
skip_items = set(_PRESERVE_DIRS)
|
||||||
for item in os.listdir(extracted_root):
|
for item in os.listdir(extracted_root):
|
||||||
if item in skip_items:
|
if item in skip_items:
|
||||||
continue
|
continue
|
||||||
@@ -445,7 +526,7 @@ class UpdateRoutes:
|
|||||||
# for ComfyUI Manager to work properly
|
# for ComfyUI Manager to work properly
|
||||||
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
||||||
tracking_files = []
|
tracking_files = []
|
||||||
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'}
|
skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
|
||||||
for root, dirs, files in os.walk(extracted_root):
|
for root, dirs, files in os.walk(extracted_root):
|
||||||
# Skip user data directories and their contents
|
# Skip user data directories and their contents
|
||||||
rel_root = os.path.relpath(root, extracted_root)
|
rel_root = os.path.relpath(root, extracted_root)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
|
import random
|
||||||
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
|
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -390,6 +391,12 @@ class BaseModelService(ABC):
|
|||||||
(item.get("model_name") or item.get("file_name") or "").lower(),
|
(item.get("model_name") or item.get("file_name") or "").lower(),
|
||||||
item.get("file_path", "").lower(),
|
item.get("file_path", "").lower(),
|
||||||
)
|
)
|
||||||
|
elif key_name == "random":
|
||||||
|
# Seeded random shuffle: same seed -> same order (stable pagination)
|
||||||
|
rng = random.Random(sort_params.seed or "random")
|
||||||
|
result = list(data)
|
||||||
|
rng.shuffle(result)
|
||||||
|
return result
|
||||||
elif key_name == "size":
|
elif key_name == "size":
|
||||||
key_fn = lambda item: (
|
key_fn = lambda item: (
|
||||||
int(item.get("size", 0) or 0),
|
int(item.get("size", 0) or 0),
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import asyncio
|
|||||||
import copy
|
import copy
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import time
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
||||||
from .connectivity_guard import (
|
from .connectivity_guard import (
|
||||||
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Best-effort cache for creator model counts, keyed by lowercase username.
|
||||||
|
# Values are (monotonic timestamp, count or None); None results are cached
|
||||||
|
# too so repeated failures don't hammer the API.
|
||||||
|
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
|
||||||
|
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
|
||||||
|
|
||||||
|
|
||||||
class CivitaiClient:
|
class CivitaiClient:
|
||||||
_instance = None
|
_instance = None
|
||||||
@@ -743,17 +750,34 @@ class CivitaiClient:
|
|||||||
|
|
||||||
return all_versions if all_versions else None
|
return all_versions if all_versions else None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(
|
||||||
"""Fetch all models for a specific Civitai user."""
|
self, username: str, cursor: Optional[str] = None
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Fetch one page (up to 100 models) for a specific Civitai user.
|
||||||
|
|
||||||
|
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||||
|
or None on failure. Pass ``cursor`` (from a previous response's
|
||||||
|
``nextCursor``) to fetch subsequent pages.
|
||||||
|
"""
|
||||||
if not username:
|
if not username:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"username": username,
|
||||||
|
"nsfw": "true",
|
||||||
|
"limit": 100,
|
||||||
|
"sort": "Newest",
|
||||||
|
"period": "AllTime",
|
||||||
|
}
|
||||||
|
if cursor:
|
||||||
|
params["cursor"] = cursor
|
||||||
|
|
||||||
try:
|
try:
|
||||||
success, result = await self._make_request(
|
success, result = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
f"{self.base_url}/models",
|
f"{self.base_url}/models",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
params={"username": username, "nsfw": "true"},
|
params=params,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not success:
|
if not success:
|
||||||
@@ -765,7 +789,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
items = result.get("items") if isinstance(result, dict) else None
|
items = result.get("items") if isinstance(result, dict) else None
|
||||||
if not isinstance(items, list):
|
if not isinstance(items, list):
|
||||||
return []
|
items = []
|
||||||
|
|
||||||
for model in items:
|
for model in items:
|
||||||
versions = model.get("modelVersions")
|
versions = model.get("modelVersions")
|
||||||
@@ -774,9 +798,68 @@ class CivitaiClient:
|
|||||||
for version in versions:
|
for version in versions:
|
||||||
self._remove_comfy_metadata(version)
|
self._remove_comfy_metadata(version)
|
||||||
|
|
||||||
return items
|
next_cursor: Optional[str] = None
|
||||||
|
metadata = result.get("metadata") if isinstance(result, dict) else None
|
||||||
|
if isinstance(metadata, dict):
|
||||||
|
raw_cursor = metadata.get("nextCursor")
|
||||||
|
if raw_cursor is not None:
|
||||||
|
next_cursor = str(raw_cursor)
|
||||||
|
|
||||||
|
return {"items": items, "nextCursor": next_cursor}
|
||||||
except RateLimitError:
|
except RateLimitError:
|
||||||
raise
|
raise
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Error fetching models for %s: %s", username, exc)
|
logger.error("Error fetching models for %s: %s", username, exc)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
"""Best-effort lookup of a creator's published model count.
|
||||||
|
|
||||||
|
Uses the ``/creators`` endpoint (a contains-match query), picking the
|
||||||
|
entry whose username matches exactly (case-insensitive). Returns None
|
||||||
|
on any failure; never raises. Results (including None) are cached
|
||||||
|
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
|
||||||
|
"""
|
||||||
|
if not username:
|
||||||
|
return None
|
||||||
|
|
||||||
|
cache_key = username.lower()
|
||||||
|
cached = _creator_model_count_cache.get(cache_key)
|
||||||
|
if cached is not None:
|
||||||
|
cached_at, cached_count = cached
|
||||||
|
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
|
||||||
|
return cached_count
|
||||||
|
|
||||||
|
count: Optional[int] = None
|
||||||
|
try:
|
||||||
|
success, result = await self._make_request(
|
||||||
|
"GET",
|
||||||
|
f"{self.base_url}/creators",
|
||||||
|
use_auth=True,
|
||||||
|
params={"query": username, "limit": 10},
|
||||||
|
)
|
||||||
|
|
||||||
|
if success and isinstance(result, dict):
|
||||||
|
creators = result.get("items")
|
||||||
|
if isinstance(creators, list):
|
||||||
|
for creator in creators:
|
||||||
|
if not isinstance(creator, dict):
|
||||||
|
continue
|
||||||
|
creator_name = creator.get("username")
|
||||||
|
if not isinstance(creator_name, str):
|
||||||
|
continue
|
||||||
|
if creator_name.lower() != cache_key:
|
||||||
|
continue
|
||||||
|
model_count = creator.get("modelCount")
|
||||||
|
if isinstance(model_count, (int, float)) and not isinstance(
|
||||||
|
model_count, bool
|
||||||
|
):
|
||||||
|
count = int(model_count)
|
||||||
|
break
|
||||||
|
except Exception as exc: # best-effort only, never propagate
|
||||||
|
logger.debug(
|
||||||
|
"Failed to fetch creator model count for %s: %s", username, exc
|
||||||
|
)
|
||||||
|
|
||||||
|
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
|
||||||
|
return count
|
||||||
|
|||||||
+21
-12
@@ -1,6 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
import logging
|
import logging
|
||||||
|
import random
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
@@ -38,8 +39,8 @@ class ModelCache:
|
|||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self._lock = asyncio.Lock()
|
self._lock = asyncio.Lock()
|
||||||
# Cache for last sort: (sort_key, order) -> sorted list
|
# Cache for last sort: (sort_key, order, seed) -> sorted list
|
||||||
self._last_sort: Tuple[str, str] = (None, None)
|
self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
|
||||||
self._last_sorted_data: List[Dict] = []
|
self._last_sorted_data: List[Dict] = []
|
||||||
self._normalize_raw_data()
|
self._normalize_raw_data()
|
||||||
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
|
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
|
||||||
@@ -203,9 +204,9 @@ class ModelCache:
|
|||||||
async def resort(self):
|
async def resort(self):
|
||||||
"""Resort cached data according to last sort mode if set"""
|
"""Resort cached data according to last sort mode if set"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
if self._last_sort != (None, None):
|
if self._last_sort[0] is not None:
|
||||||
sort_key, order = self._last_sort
|
sort_key, order, seed = self._last_sort
|
||||||
sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||||
self._last_sorted_data = sorted_data
|
self._last_sorted_data = sorted_data
|
||||||
# Update folder list
|
# Update folder list
|
||||||
# else: do nothing
|
# else: do nothing
|
||||||
@@ -218,7 +219,7 @@ class ModelCache:
|
|||||||
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
||||||
self.rebuild_version_index()
|
self.rebuild_version_index()
|
||||||
|
|
||||||
def _sort_data(self, data: List[Dict], sort_key: str, order: str) -> List[Dict]:
|
def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]:
|
||||||
"""Sort data by sort_key and order"""
|
"""Sort data by sort_key and order"""
|
||||||
start_time = time.perf_counter()
|
start_time = time.perf_counter()
|
||||||
reverse = (order == 'desc')
|
reverse = (order == 'desc')
|
||||||
@@ -265,6 +266,13 @@ class ModelCache:
|
|||||||
),
|
),
|
||||||
reverse=reverse
|
reverse=reverse
|
||||||
)
|
)
|
||||||
|
elif sort_key == 'random':
|
||||||
|
# Random shuffle seeded for stable pagination: the same seed
|
||||||
|
# always yields the same order, so successive page requests
|
||||||
|
# stay consistent while browsing.
|
||||||
|
rng = random.Random(seed or 'random')
|
||||||
|
result = list(data)
|
||||||
|
rng.shuffle(result)
|
||||||
elif sort_key == 'versions_count':
|
elif sort_key == 'versions_count':
|
||||||
# Pre-dedup sort: fall back to name sort.
|
# Pre-dedup sort: fall back to name sort.
|
||||||
# Actual re-sort by version_count happens in get_paginated_data after dedup.
|
# Actual re-sort by version_count happens in get_paginated_data after dedup.
|
||||||
@@ -285,15 +293,16 @@ class ModelCache:
|
|||||||
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
|
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc') -> List[Dict]:
|
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]:
|
||||||
"""Get sorted data by sort_key and order, using cache if possible"""
|
"""Get sorted data by sort_key and order, using cache if possible"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
if (sort_key, order) == self._last_sort:
|
cache_key = (sort_key, order, seed)
|
||||||
|
if cache_key == self._last_sort:
|
||||||
return self._last_sorted_data
|
return self._last_sorted_data
|
||||||
|
|
||||||
start_time = time.perf_counter()
|
start_time = time.perf_counter()
|
||||||
sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||||
self._last_sort = (sort_key, order)
|
self._last_sort = cache_key
|
||||||
self._last_sorted_data = sorted_data
|
self._last_sorted_data = sorted_data
|
||||||
|
|
||||||
duration = time.perf_counter() - start_time
|
duration = time.perf_counter() - start_time
|
||||||
@@ -313,8 +322,8 @@ class ModelCache:
|
|||||||
self.name_display_mode = normalized
|
self.name_display_mode = normalized
|
||||||
|
|
||||||
if self._last_sort[0] == 'name':
|
if self._last_sort[0] == 'name':
|
||||||
sort_key, order = self._last_sort
|
sort_key, order, seed = self._last_sort
|
||||||
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||||
|
|
||||||
async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
|
async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
|
||||||
"""Update preview_url for a specific model in all cached data
|
"""Update preview_url for a specific model in all cached data
|
||||||
|
|||||||
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
"""Fetch models owned by the specified user"""
|
"""Fetch one page of models owned by the specified user.
|
||||||
|
|
||||||
|
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||||
|
or None when unsupported/failed. ``cursor`` continues a previous page.
|
||||||
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
"""Published model count for the user; None when unsupported."""
|
||||||
|
return None
|
||||||
|
|
||||||
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses Civitai API for metadata"""
|
"""Provider that uses Civitai API for metadata"""
|
||||||
|
|
||||||
@@ -175,8 +183,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
|||||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||||
return await self.client.get_model_version_info(version_id)
|
return await self.client.get_model_version_info(version_id)
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
return await self.client.get_user_models(username)
|
return await self.client.get_user_models(username, cursor)
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
return await self.client.get_creator_model_count(username)
|
||||||
|
|
||||||
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses CivArchive API for metadata"""
|
"""Provider that uses CivArchive API for metadata"""
|
||||||
@@ -196,7 +207,7 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
|||||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||||
return await self.client.get_model_version_info(version_id)
|
return await self.client.get_model_version_info(version_id)
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
"""Not supported by CivArchive provider"""
|
"""Not supported by CivArchive provider"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
|
|||||||
version_data = await self._get_version_with_model_data(db, model_id, version_id)
|
version_data = await self._get_version_with_model_data(db, model_id, version_id)
|
||||||
return version_data, None
|
return version_data, None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
"""Listing models by username is not supported for archive database"""
|
"""Listing models by username is not supported for archive database"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
continue
|
continue
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
try:
|
try:
|
||||||
result = await self._call_with_rate_limit(
|
result = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
provider.get_user_models,
|
provider.get_user_models,
|
||||||
username,
|
username,
|
||||||
|
cursor=cursor,
|
||||||
)
|
)
|
||||||
if result is not None:
|
if result is not None:
|
||||||
return result
|
return result
|
||||||
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
continue
|
continue
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
for provider, label in self._iter_providers():
|
||||||
|
try:
|
||||||
|
result = await provider.get_creator_model_count(username)
|
||||||
|
if result is not None:
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(
|
||||||
|
"Provider %s failed for get_creator_model_count: %s", label, e
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
return None
|
||||||
|
|
||||||
def _iter_providers(self):
|
def _iter_providers(self):
|
||||||
return zip(self.providers, self._provider_labels)
|
return zip(self.providers, self._provider_labels)
|
||||||
|
|
||||||
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
|
|||||||
version_id,
|
version_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||||
return await self._rate_limit_helper.run(
|
return await self._rate_limit_helper.run(
|
||||||
self._label,
|
self._label,
|
||||||
self._provider.get_user_models,
|
self._provider.get_user_models,
|
||||||
username,
|
username,
|
||||||
|
cursor=cursor,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
return await self._provider.get_creator_model_count(username)
|
||||||
|
|
||||||
class ModelMetadataProviderManager:
|
class ModelMetadataProviderManager:
|
||||||
"""Manager for selecting and using model metadata providers"""
|
"""Manager for selecting and using model metadata providers"""
|
||||||
|
|
||||||
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
|
|||||||
except NotImplementedError:
|
except NotImplementedError:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]:
|
async def get_user_models(
|
||||||
"""Fetch models owned by the specified user"""
|
self,
|
||||||
|
username: str,
|
||||||
|
provider_name: str = None,
|
||||||
|
cursor: Optional[str] = None,
|
||||||
|
) -> Optional[Dict]:
|
||||||
|
"""Fetch one page of models owned by the specified user"""
|
||||||
provider = self._get_provider(provider_name)
|
provider = self._get_provider(provider_name)
|
||||||
return await provider.get_user_models(username)
|
return await provider.get_user_models(username, cursor)
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]:
|
||||||
|
"""Best-effort published model count for the specified user"""
|
||||||
|
provider = self._get_provider(provider_name)
|
||||||
|
return await provider.get_creator_model_count(username)
|
||||||
|
|
||||||
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
|
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
|
||||||
"""Get provider by name or default provider"""
|
"""Get provider by name or default provider"""
|
||||||
|
|||||||
@@ -85,6 +85,7 @@ class SortParams:
|
|||||||
|
|
||||||
key: str
|
key: str
|
||||||
order: str
|
order: str
|
||||||
|
seed: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -116,7 +117,7 @@ class ModelCacheRepository:
|
|||||||
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
|
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
|
||||||
"""Fetch cached data pre-sorted according to ``params``."""
|
"""Fetch cached data pre-sorted according to ``params``."""
|
||||||
cache = await self.get_cache()
|
cache = await self.get_cache()
|
||||||
return await cache.get_sorted_data(params.key, params.order)
|
return await cache.get_sorted_data(params.key, params.order, params.seed)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def parse_sort(sort_by: str) -> SortParams:
|
def parse_sort(sort_by: str) -> SortParams:
|
||||||
@@ -132,10 +133,17 @@ class ModelCacheRepository:
|
|||||||
sort_key = sort_by.strip().lower() or "name"
|
sort_key = sort_by.strip().lower() or "name"
|
||||||
order = "asc"
|
order = "asc"
|
||||||
|
|
||||||
if order not in ("asc", "desc"):
|
seed = None
|
||||||
|
if sort_key == "random":
|
||||||
|
# Random sort: the portion after ':' is the shuffle seed.
|
||||||
|
# A stable seed keeps paginated requests consistent; order is
|
||||||
|
# meaningless for a random shuffle.
|
||||||
|
seed = order if order and order not in ("asc", "desc") else None
|
||||||
|
order = "asc"
|
||||||
|
elif order not in ("asc", "desc"):
|
||||||
order = "asc"
|
order = "asc"
|
||||||
|
|
||||||
return SortParams(key=sort_key, order=order)
|
return SortParams(key=sort_key, order=order, seed=seed)
|
||||||
|
|
||||||
|
|
||||||
class ModelFilterSet:
|
class ModelFilterSet:
|
||||||
|
|||||||
@@ -1752,7 +1752,7 @@ class ModelScanner:
|
|||||||
# ---- Conditional resort (only when sort-key fields changed) ----
|
# ---- Conditional resort (only when sort-key fields changed) ----
|
||||||
need_resort = False
|
need_resort = False
|
||||||
_last = cache._last_sort
|
_last = cache._last_sort
|
||||||
sort_key: Optional[str] = _last[0] if _last != (None, None) else None
|
sort_key: Optional[str] = _last[0] if _last[0] is not None else None
|
||||||
if sort_key == "name":
|
if sort_key == "name":
|
||||||
if (
|
if (
|
||||||
old_model_name != desired_entry.get("model_name", "")
|
old_model_name != desired_entry.get("model_name", "")
|
||||||
|
|||||||
@@ -14,11 +14,16 @@ from ..services.service_registry import ServiceRegistry
|
|||||||
from ..utils.example_images_paths import (
|
from ..utils.example_images_paths import (
|
||||||
ExampleImagePathResolver,
|
ExampleImagePathResolver,
|
||||||
ensure_library_root_exists,
|
ensure_library_root_exists,
|
||||||
|
get_example_images_root,
|
||||||
|
is_hash_folder,
|
||||||
uses_library_scoped_folders,
|
uses_library_scoped_folders,
|
||||||
)
|
)
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from .example_images_processor import ExampleImagesProcessor
|
from .example_images_processor import ExampleImagesProcessor
|
||||||
from .example_images_metadata import MetadataUpdater
|
from .example_images_metadata import (
|
||||||
|
MetadataUpdater,
|
||||||
|
update_cache_from_metadata,
|
||||||
|
)
|
||||||
from ..services.downloader import get_downloader
|
from ..services.downloader import get_downloader
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
|
|
||||||
@@ -87,6 +92,13 @@ class _DownloadProgress(dict):
|
|||||||
return snapshot
|
return snapshot
|
||||||
|
|
||||||
|
|
||||||
|
# When fewer candidates than this remain in check_pending_models, probe each
|
||||||
|
# model folder directly (preserving legacy-folder migration semantics). Above
|
||||||
|
# it, build a folder index with a single directory scan so libraries with
|
||||||
|
# 100k+ models do not pay one syscall per candidate.
|
||||||
|
_BULK_LOOKUP_THRESHOLD = 1000
|
||||||
|
|
||||||
|
|
||||||
def _model_directory_has_files(path: str) -> bool:
|
def _model_directory_has_files(path: str) -> bool:
|
||||||
"""Return True when the provided directory exists and contains entries."""
|
"""Return True when the provided directory exists and contains entries."""
|
||||||
|
|
||||||
@@ -103,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _build_example_folder_index(output_dir: str) -> dict[str, bool]:
|
||||||
|
"""Build a ``{hash: has_files}`` index for a library's example-image folders.
|
||||||
|
|
||||||
|
A single directory scan over the library root replaces ``O(candidates)``
|
||||||
|
per-folder ``os.scandir`` calls, which is required for libraries with
|
||||||
|
100k+ models. Each hash folder is classified by whether it contains any
|
||||||
|
entries, matching the semantics of ``_model_directory_has_files``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
index: dict[str, bool] = {}
|
||||||
|
if not output_dir or not os.path.isdir(output_dir):
|
||||||
|
return index
|
||||||
|
|
||||||
|
try:
|
||||||
|
with os.scandir(output_dir) as entries:
|
||||||
|
for entry in entries:
|
||||||
|
name = entry.name
|
||||||
|
if not entry.is_dir() or not is_hash_folder(name):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
with os.scandir(entry.path) as subentries:
|
||||||
|
index[name.lower()] = any(subentries)
|
||||||
|
except OSError:
|
||||||
|
index[name.lower()] = False
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return index
|
||||||
|
|
||||||
|
|
||||||
class DownloadManager:
|
class DownloadManager:
|
||||||
"""Manages downloading example images for models."""
|
"""Manages downloading example images for models."""
|
||||||
|
|
||||||
@@ -410,14 +452,49 @@ class DownloadManager:
|
|||||||
# Calculate pending count: check which models actually need processing.
|
# Calculate pending count: check which models actually need processing.
|
||||||
# A model is pending if it has a hash, is not already processed or known-failed,
|
# A model is pending if it has a hash, is not already processed or known-failed,
|
||||||
# and its folder doesn't exist or is empty.
|
# and its folder doesn't exist or is empty.
|
||||||
pending_hashes = set()
|
candidate_hashes = [
|
||||||
for model_hash, model_name in all_models_with_hash:
|
model_hash
|
||||||
if model_hash not in processed_models and model_hash not in failed_models:
|
for model_hash, _ in all_models_with_hash
|
||||||
|
if model_hash not in processed_models
|
||||||
|
and model_hash not in failed_models
|
||||||
|
]
|
||||||
|
|
||||||
|
pending_hashes: set[str] = set()
|
||||||
|
# For small candidate counts the existing per-folder check is fine
|
||||||
|
# and handles legacy folder migration.
|
||||||
|
# For large libraries, scan the library root once and do set lookups.
|
||||||
|
if len(candidate_hashes) <= _BULK_LOOKUP_THRESHOLD or not output_dir:
|
||||||
|
for model_hash in candidate_hashes:
|
||||||
model_dir = ExampleImagePathResolver.get_model_folder(
|
model_dir = ExampleImagePathResolver.get_model_folder(
|
||||||
model_hash, active_library
|
model_hash, active_library
|
||||||
)
|
)
|
||||||
if not _model_directory_has_files(model_dir):
|
if not _model_directory_has_files(model_dir):
|
||||||
pending_hashes.add(model_hash)
|
pending_hashes.add(model_hash)
|
||||||
|
else:
|
||||||
|
folder_index = await asyncio.get_event_loop().run_in_executor(
|
||||||
|
None, _build_example_folder_index, output_dir
|
||||||
|
)
|
||||||
|
# In multi-library mode, folders that have not been consolidated
|
||||||
|
# into the library root yet (startup migration skipped, failed
|
||||||
|
# move, or created at the legacy path afterwards) still live at
|
||||||
|
# the legacy root/<hash> location. Only scan that root when at
|
||||||
|
# least one candidate is missing from the library-root index, so
|
||||||
|
# the fully-consolidated case does not pay an extra directory
|
||||||
|
# pass on every call.
|
||||||
|
if uses_library_scoped_folders() and any(
|
||||||
|
not folder_index.get(model_hash, False)
|
||||||
|
for model_hash in candidate_hashes
|
||||||
|
):
|
||||||
|
legacy_root = get_example_images_root()
|
||||||
|
if legacy_root and legacy_root != output_dir:
|
||||||
|
legacy_index = await asyncio.get_event_loop().run_in_executor(
|
||||||
|
None, _build_example_folder_index, legacy_root
|
||||||
|
)
|
||||||
|
for hash_key, has_files in legacy_index.items():
|
||||||
|
folder_index.setdefault(hash_key, has_files)
|
||||||
|
for model_hash in candidate_hashes:
|
||||||
|
if not folder_index.get(model_hash, False):
|
||||||
|
pending_hashes.add(model_hash)
|
||||||
|
|
||||||
pending_count = len(pending_hashes)
|
pending_count = len(pending_hashes)
|
||||||
|
|
||||||
@@ -1343,8 +1420,8 @@ class DownloadManager:
|
|||||||
await MetadataManager.save_metadata(file_path, model_copy)
|
await MetadataManager.save_metadata(file_path, model_copy)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await scanner.update_single_model_cache(
|
await update_cache_from_metadata(
|
||||||
file_path, file_path, model_data
|
scanner, file_path, model_copy
|
||||||
)
|
)
|
||||||
except AttributeError:
|
except AttributeError:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
@@ -28,6 +29,31 @@ if TYPE_CHECKING: # pragma: no cover - import for type checkers only
|
|||||||
from ..services.settings_manager import SettingsManager
|
from ..services.settings_manager import SettingsManager
|
||||||
|
|
||||||
|
|
||||||
|
async def update_cache_from_metadata(
|
||||||
|
scanner: Any, file_path: str, metadata: Dict[str, Any]
|
||||||
|
) -> bool:
|
||||||
|
"""Update the scanner cache from a metadata dict using the in-place sync path.
|
||||||
|
|
||||||
|
``sync_cache_from_metadata`` patches the existing cache entry incrementally
|
||||||
|
(tag/hash/version indexes, targeted single-row SQL update) and only resorts
|
||||||
|
when a sort-key field changed. This avoids the ``O(n)`` full-list resort and
|
||||||
|
full cache rewrite that ``update_single_model_cache`` performs on every call,
|
||||||
|
which is critical for libraries with 100k+ models.
|
||||||
|
|
||||||
|
Falls back to the legacy full update when the scanner does not expose an
|
||||||
|
async ``sync_cache_from_metadata`` method.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``True`` if the cache entry was updated, ``False`` otherwise.
|
||||||
|
"""
|
||||||
|
|
||||||
|
sync_method = getattr(scanner, "sync_cache_from_metadata", None)
|
||||||
|
if inspect.iscoroutinefunction(sync_method):
|
||||||
|
return await sync_method(file_path, metadata)
|
||||||
|
|
||||||
|
return await scanner.update_single_model_cache(file_path, file_path, metadata)
|
||||||
|
|
||||||
|
|
||||||
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
|
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
|
||||||
"""Construct a metadata sync service bound to the provided settings."""
|
"""Construct a metadata sync service bound to the provided settings."""
|
||||||
|
|
||||||
@@ -103,8 +129,8 @@ class MetadataUpdater:
|
|||||||
progress['refreshed_models'].add(model_hash)
|
progress['refreshed_models'].add(model_hash)
|
||||||
|
|
||||||
async def update_cache_func(old_path, new_path, metadata):
|
async def update_cache_func(old_path, new_path, metadata):
|
||||||
return await scanner.update_single_model_cache(old_path, new_path, metadata)
|
return await update_cache_from_metadata(scanner, new_path, metadata)
|
||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model_data)
|
await MetadataManager.hydrate_model_data(model_data)
|
||||||
success, error = await _get_metadata_sync_service().fetch_and_update_model(
|
success, error = await _get_metadata_sync_service().fetch_and_update_model(
|
||||||
sha256=model_hash,
|
sha256=model_hash,
|
||||||
@@ -234,6 +260,7 @@ class MetadataUpdater:
|
|||||||
|
|
||||||
# Save metadata to .metadata.json file
|
# Save metadata to .metadata.json file
|
||||||
file_path = model.get('file_path')
|
file_path = model.get('file_path')
|
||||||
|
model_copy: Optional[Dict[str, Any]] = None
|
||||||
try:
|
try:
|
||||||
model_copy = model.copy()
|
model_copy = model.copy()
|
||||||
model_copy.pop('folder', None)
|
model_copy.pop('folder', None)
|
||||||
@@ -241,14 +268,18 @@ class MetadataUpdater:
|
|||||||
logger.info(f"Saved metadata for {model.get('model_name')}")
|
logger.info(f"Saved metadata for {model.get('model_name')}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
|
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
|
||||||
|
|
||||||
# Save updated metadata to scanner cache
|
# Save updated metadata to scanner cache. sync_cache_from_metadata
|
||||||
success = await scanner.update_single_model_cache(file_path, file_path, model)
|
# returns False both for "already in sync" and for actual failures,
|
||||||
if success:
|
# so the cache sync result is deliberately not treated as an error;
|
||||||
|
# the return value reflects whether the metadata was persisted.
|
||||||
|
if file_path and model_copy is not None:
|
||||||
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
|
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
|
||||||
return True
|
return True
|
||||||
else:
|
|
||||||
logger.warning(f"Failed to update metadata for {model.get('model_name')}")
|
logger.warning(f"Failed to update metadata for {model.get('model_name')}")
|
||||||
|
return False
|
||||||
|
|
||||||
return False
|
return False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -336,6 +367,7 @@ class MetadataUpdater:
|
|||||||
|
|
||||||
# Save metadata to .metadata.json file
|
# Save metadata to .metadata.json file
|
||||||
file_path = model_data.get('file_path')
|
file_path = model_data.get('file_path')
|
||||||
|
model_copy: Optional[Dict[str, Any]] = None
|
||||||
if file_path:
|
if file_path:
|
||||||
try:
|
try:
|
||||||
model_copy = model_data.copy()
|
model_copy = model_data.copy()
|
||||||
@@ -344,11 +376,11 @@ class MetadataUpdater:
|
|||||||
logger.info(f"Saved metadata for {model_data.get('model_name')}")
|
logger.info(f"Saved metadata for {model_data.get('model_name')}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to save metadata: {str(e)}")
|
logger.error(f"Failed to save metadata: {str(e)}")
|
||||||
|
|
||||||
# Save updated metadata to scanner cache
|
# Save updated metadata to scanner cache
|
||||||
if file_path:
|
if file_path and model_copy is not None:
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
|
|
||||||
# Get regular images array (might be None)
|
# Get regular images array (might be None)
|
||||||
regular_images = civitai_data.get('images', [])
|
regular_images = civitai_data.get('images', [])
|
||||||
|
|
||||||
@@ -475,13 +507,19 @@ class MetadataUpdater:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
model_folder = get_model_folder(model_hash)
|
model_folder = get_model_folder(model_hash)
|
||||||
if not model_folder:
|
if not model_folder or not os.path.isdir(model_folder):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
civitai = getattr(metadata, "civitai", None)
|
civitai = getattr(metadata, "civitai", None)
|
||||||
if not isinstance(civitai, dict):
|
if not isinstance(civitai, dict):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
# Read the directory listing once so every image entry reuses it.
|
||||||
|
try:
|
||||||
|
dir_entries = os.listdir(model_folder)
|
||||||
|
except OSError:
|
||||||
|
dir_entries = []
|
||||||
|
|
||||||
has_changes = False
|
has_changes = False
|
||||||
|
|
||||||
custom_images = civitai.get("customImages")
|
custom_images = civitai.get("customImages")
|
||||||
@@ -493,24 +531,15 @@ class MetadataUpdater:
|
|||||||
if not img_id:
|
if not img_id:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if not os.path.isdir(model_folder):
|
prefix = f"custom_{img_id}"
|
||||||
|
found = any(
|
||||||
|
f.startswith(prefix) and os.path.isfile(
|
||||||
|
os.path.join(model_folder, f)
|
||||||
|
)
|
||||||
|
for f in dir_entries
|
||||||
|
)
|
||||||
|
if not found:
|
||||||
stale.append(idx)
|
stale.append(idx)
|
||||||
else:
|
|
||||||
found = False
|
|
||||||
try:
|
|
||||||
prefix = f"custom_{img_id}"
|
|
||||||
for fname in os.listdir(model_folder):
|
|
||||||
if fname.startswith(prefix) and os.path.isfile(
|
|
||||||
os.path.join(model_folder, fname)
|
|
||||||
):
|
|
||||||
found = True
|
|
||||||
break
|
|
||||||
except OSError:
|
|
||||||
stale.append(idx)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if not found:
|
|
||||||
stale.append(idx)
|
|
||||||
|
|
||||||
if stale:
|
if stale:
|
||||||
for idx in reversed(stale):
|
for idx in reversed(stale):
|
||||||
@@ -532,22 +561,9 @@ class MetadataUpdater:
|
|||||||
# is gone.
|
# is gone.
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if not os.path.isdir(model_folder):
|
prefix = f"image_{idx}."
|
||||||
|
if not any(f.startswith(prefix) for f in dir_entries):
|
||||||
stale.append(idx)
|
stale.append(idx)
|
||||||
else:
|
|
||||||
found = False
|
|
||||||
try:
|
|
||||||
prefix = f"image_{idx}."
|
|
||||||
for fname in os.listdir(model_folder):
|
|
||||||
if fname.startswith(prefix):
|
|
||||||
found = True
|
|
||||||
break
|
|
||||||
except OSError:
|
|
||||||
stale.append(idx)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if not found:
|
|
||||||
stale.append(idx)
|
|
||||||
|
|
||||||
if stale:
|
if stale:
|
||||||
for idx in reversed(stale):
|
for idx in reversed(stale):
|
||||||
|
|||||||
@@ -3,11 +3,19 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import json
|
import json
|
||||||
|
import shutil
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..utils.example_images_paths import iter_library_roots
|
from ..utils.example_images_paths import (
|
||||||
|
get_example_images_root,
|
||||||
|
is_hash_folder,
|
||||||
|
iter_library_roots,
|
||||||
|
uses_library_scoped_folders,
|
||||||
|
_library_folder_has_only_hash_dirs,
|
||||||
|
)
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from ..utils.example_images_processor import ExampleImagesProcessor
|
from ..utils.example_images_processor import ExampleImagesProcessor
|
||||||
|
from ..utils.example_images_metadata import update_cache_from_metadata
|
||||||
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -36,6 +44,90 @@ settings = _SettingsProxy()
|
|||||||
class ExampleImagesMigration:
|
class ExampleImagesMigration:
|
||||||
"""Handles migrations for example images naming conventions"""
|
"""Handles migrations for example images naming conventions"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _consolidate_library_folders():
|
||||||
|
"""Move hash folders from library-named subdirectories back to root.
|
||||||
|
|
||||||
|
When a user switches from multi-library mode back to single-library
|
||||||
|
mode, example images previously stored under e.g.
|
||||||
|
``<root>/default/<hash>/`` need to be moved back to
|
||||||
|
``<root>/<hash>/``. Running this once at startup removes the need
|
||||||
|
for ``get_model_folder()`` to perform directory scans on every
|
||||||
|
request.
|
||||||
|
"""
|
||||||
|
if uses_library_scoped_folders():
|
||||||
|
return
|
||||||
|
|
||||||
|
root = get_example_images_root()
|
||||||
|
if not root or not os.path.isdir(root):
|
||||||
|
return
|
||||||
|
|
||||||
|
moved: list[str] = []
|
||||||
|
cleaned: list[str] = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
for entry in os.listdir(root):
|
||||||
|
# Fast regex checks first — no filesystem I/O.
|
||||||
|
if is_hash_folder(entry) or entry == "_deleted":
|
||||||
|
continue
|
||||||
|
|
||||||
|
entry_path = os.path.join(root, entry)
|
||||||
|
if not os.path.isdir(entry_path):
|
||||||
|
continue
|
||||||
|
if not _library_folder_has_only_hash_dirs(entry_path):
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
for hash_entry in os.listdir(entry_path):
|
||||||
|
hash_path = os.path.join(entry_path, hash_entry)
|
||||||
|
if not os.path.isdir(hash_path) or not is_hash_folder(hash_entry):
|
||||||
|
continue
|
||||||
|
target = os.path.join(root, hash_entry)
|
||||||
|
if not os.path.exists(target):
|
||||||
|
try:
|
||||||
|
shutil.move(hash_path, target)
|
||||||
|
moved.append(hash_entry)
|
||||||
|
except (OSError, shutil.Error) as exc:
|
||||||
|
logger.error(
|
||||||
|
"Failed to move '%s' → '%s': %s",
|
||||||
|
hash_path, target, exc,
|
||||||
|
)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.error(
|
||||||
|
"Failed to list library subdirectory '%s': %s",
|
||||||
|
entry_path, exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
remaining = os.listdir(entry_path)
|
||||||
|
except OSError:
|
||||||
|
remaining = []
|
||||||
|
if not remaining:
|
||||||
|
try:
|
||||||
|
os.rmdir(entry_path)
|
||||||
|
cleaned.append(entry)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.debug(
|
||||||
|
"Could not remove empty library dir '%s': %s",
|
||||||
|
entry_path, exc,
|
||||||
|
)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.error(
|
||||||
|
"Failed to list example images root during consolidation: %s",
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
if moved:
|
||||||
|
logger.info(
|
||||||
|
"Consolidated %d example image folder(s) to root",
|
||||||
|
len(moved),
|
||||||
|
)
|
||||||
|
if cleaned:
|
||||||
|
logger.info(
|
||||||
|
"Removed %d empty library directories",
|
||||||
|
len(cleaned),
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def check_and_run_migrations():
|
async def check_and_run_migrations():
|
||||||
"""Check if migrations are needed and run them in background"""
|
"""Check if migrations are needed and run them in background"""
|
||||||
@@ -44,6 +136,10 @@ class ExampleImagesMigration:
|
|||||||
logger.debug("No example images path configured or path doesn't exist, skipping migrations")
|
logger.debug("No example images path configured or path doesn't exist, skipping migrations")
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Run library-to-root consolidation once at startup so the hot
|
||||||
|
# path (get_model_folder) stays a pure-path computation.
|
||||||
|
ExampleImagesMigration._consolidate_library_folders()
|
||||||
|
|
||||||
for library_name, library_path in iter_library_roots():
|
for library_name, library_path in iter_library_roots():
|
||||||
if not library_path or not os.path.exists(library_path):
|
if not library_path or not os.path.exists(library_path):
|
||||||
continue
|
continue
|
||||||
@@ -326,7 +422,7 @@ class ExampleImagesMigration:
|
|||||||
await MetadataManager.save_metadata(file_path, model_copy)
|
await MetadataManager.save_metadata(file_path, model_copy)
|
||||||
|
|
||||||
# Update scanner cache
|
# Update scanner cache
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_metadata)
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
|
|
||||||
updated_models += 1
|
updated_models += 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -83,7 +83,12 @@ def ensure_library_root_exists(library_name: Optional[str] = None) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str:
|
def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str:
|
||||||
"""Return the folder path for a model's example images."""
|
"""Return the folder path for a model's example images.
|
||||||
|
|
||||||
|
Multi-library ↔ single-library consolidation is handled once at startup by
|
||||||
|
``ExampleImagesMigration._consolidate_library_folders`` — this function is a
|
||||||
|
pure path computation on the hot path (no directory scans).
|
||||||
|
"""
|
||||||
|
|
||||||
if not model_hash:
|
if not model_hash:
|
||||||
return ""
|
return ""
|
||||||
@@ -113,35 +118,6 @@ def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str
|
|||||||
exc,
|
exc,
|
||||||
)
|
)
|
||||||
return legacy_folder
|
return legacy_folder
|
||||||
elif not os.path.exists(resolved_folder):
|
|
||||||
# Reverse migration: when consolidating from multi-library to
|
|
||||||
# single-library mode (e.g. after "default" was cleaned up), look
|
|
||||||
# for existing example images inside library-named subdirectories
|
|
||||||
# and bring them back to the root level.
|
|
||||||
root = get_example_images_root()
|
|
||||||
if root:
|
|
||||||
try:
|
|
||||||
for entry in os.listdir(root):
|
|
||||||
entry_path = os.path.join(root, entry)
|
|
||||||
if not os.path.isdir(entry_path):
|
|
||||||
continue
|
|
||||||
if is_hash_folder(entry) or entry == "_deleted":
|
|
||||||
continue
|
|
||||||
if not _library_folder_has_only_hash_dirs(entry_path):
|
|
||||||
continue
|
|
||||||
legacy = os.path.join(entry_path, normalized_hash)
|
|
||||||
if os.path.exists(legacy):
|
|
||||||
shutil.move(legacy, resolved_folder)
|
|
||||||
logger.info(
|
|
||||||
"Consolidated example images from '%s' to '%s'",
|
|
||||||
legacy, resolved_folder,
|
|
||||||
)
|
|
||||||
break
|
|
||||||
except OSError as exc:
|
|
||||||
logger.error(
|
|
||||||
"Failed to consolidate example images during "
|
|
||||||
"library merge: %s", exc,
|
|
||||||
)
|
|
||||||
|
|
||||||
return resolved_folder
|
return resolved_folder
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
|||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
from ..utils.example_images_paths import get_model_folder, get_model_relative_path
|
from ..utils.example_images_paths import get_model_folder, get_model_relative_path
|
||||||
from .example_images_metadata import MetadataUpdater
|
from .example_images_metadata import MetadataUpdater, update_cache_from_metadata
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -644,7 +644,7 @@ class ExampleImagesProcessor:
|
|||||||
}, status=500)
|
}, status=500)
|
||||||
|
|
||||||
# Update cache
|
# Update cache
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
await update_cache_from_metadata(scanner, file_path, model_data)
|
||||||
|
|
||||||
# Get regular images array (might be None)
|
# Get regular images array (might be None)
|
||||||
regular_images = civitai_data.get('images', [])
|
regular_images = civitai_data.get('images', [])
|
||||||
@@ -759,7 +759,7 @@ class ExampleImagesProcessor:
|
|||||||
model_copy = model_data.copy()
|
model_copy = model_data.copy()
|
||||||
model_copy.pop('folder', None)
|
model_copy.pop('folder', None)
|
||||||
await MetadataManager.save_metadata(file_path, model_copy)
|
await MetadataManager.save_metadata(file_path, model_copy)
|
||||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||||
|
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': True,
|
'success': True,
|
||||||
|
|||||||
+48
-1
@@ -1,7 +1,7 @@
|
|||||||
from difflib import SequenceMatcher
|
from difflib import SequenceMatcher
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from typing import Dict
|
from typing import Any, Dict, List, Optional
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..config import config
|
from ..config import config
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
@@ -294,6 +294,53 @@ def _format_model_name_for_comfyui(file_path: str, model_roots: list) -> str:
|
|||||||
return os.path.basename(file_path)
|
return os.path.basename(file_path)
|
||||||
|
|
||||||
|
|
||||||
|
def model_patcher_to_name(model_patcher: Any) -> Optional[str]:
|
||||||
|
"""Extract a ComfyUI-style model name from a MODEL (ModelPatcher) object.
|
||||||
|
|
||||||
|
Core ComfyUI loaders record the absolute weight file path on the patcher's
|
||||||
|
``cached_patcher_init`` attribute:
|
||||||
|
- load_checkpoint_guess_config -> (fn, (ckpt_path, ...), index)
|
||||||
|
- load_diffusion_model -> (fn, (unet_path, model_options))
|
||||||
|
Patcher clones (LoRA loaders, model merges, ...) preserve the attribute,
|
||||||
|
so the name is recoverable anywhere downstream of a core loader — including
|
||||||
|
from LoRA Manager's own loaders (CheckpointLoaderLM / UNETLoaderLM), which
|
||||||
|
call the same core load functions.
|
||||||
|
|
||||||
|
The absolute path is converted to the ComfyUI-style relative name used by
|
||||||
|
the metadata pipeline (covering standard ComfyUI roots and LoRA Manager
|
||||||
|
extra folder paths).
|
||||||
|
|
||||||
|
Returns None when the path cannot be recovered (e.g. third-party loaders
|
||||||
|
that never set ``cached_patcher_init``).
|
||||||
|
"""
|
||||||
|
init = getattr(model_patcher, "cached_patcher_init", None)
|
||||||
|
if not isinstance(init, (tuple, list)) or len(init) < 2:
|
||||||
|
return None
|
||||||
|
args = init[1]
|
||||||
|
abs_path = args[0] if args else None
|
||||||
|
if not isinstance(abs_path, str) or not abs_path:
|
||||||
|
return None
|
||||||
|
return _abs_model_path_to_name(abs_path)
|
||||||
|
|
||||||
|
|
||||||
|
def _abs_model_path_to_name(abs_path: str) -> str:
|
||||||
|
"""Convert an absolute model path to a ComfyUI-style relative name.
|
||||||
|
|
||||||
|
Tries standard ComfyUI model roots plus LoRA Manager extra folder paths;
|
||||||
|
falls back to the bare filename.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
roots: List[str] = list(config.base_models_roots or [])
|
||||||
|
roots.extend(config.extra_checkpoints_roots or [])
|
||||||
|
roots.extend(config.extra_unet_roots or [])
|
||||||
|
formatted = _format_model_name_for_comfyui(abs_path, roots)
|
||||||
|
if formatted:
|
||||||
|
return formatted
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return os.path.basename(abs_path)
|
||||||
|
|
||||||
|
|
||||||
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
||||||
"""
|
"""
|
||||||
Check if text matches pattern using fuzzy matching.
|
Check if text matches pattern using fuzzy matching.
|
||||||
|
|||||||
+1
-1
@@ -1,7 +1,7 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "comfyui-lora-manager"
|
name = "comfyui-lora-manager"
|
||||||
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
|
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
|
||||||
version = "1.1.9"
|
version = "1.2.0"
|
||||||
license = {file = "LICENSE"}
|
license = {file = "LICENSE"}
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"aiohttp",
|
"aiohttp",
|
||||||
|
|||||||
@@ -108,10 +108,20 @@ export class PageControls {
|
|||||||
const sortSelect = document.getElementById('sortSelect');
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
if (sortSelect) {
|
if (sortSelect) {
|
||||||
initSortDropdown(sortSelect);
|
initSortDropdown(sortSelect);
|
||||||
sortSelect.value = this.pageState.sortBy;
|
this.applySortToSelect(this.pageState.sortBy);
|
||||||
sortSelect.addEventListener('change', async (e) => {
|
sortSelect.addEventListener('change', async (e) => {
|
||||||
this.pageState.sortBy = e.target.value;
|
let value = e.target.value;
|
||||||
this.saveSortPreference(e.target.value);
|
if (value.startsWith('random')) {
|
||||||
|
// Every pick of Random reshuffles the list: generate a
|
||||||
|
// fresh seed so the backend keeps a stable order across
|
||||||
|
// paginated requests.
|
||||||
|
value = this._randomizeSortValue();
|
||||||
|
}
|
||||||
|
this.pageState.sortBy = value;
|
||||||
|
this.saveSortPreference(value);
|
||||||
|
// Reset the seeded Random option when switching away from
|
||||||
|
// Random, or re-apply the fresh seed when picking it again.
|
||||||
|
this.applySortToSelect(value);
|
||||||
await this.resetAndReload();
|
await this.resetAndReload();
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
@@ -312,6 +322,44 @@ export class PageControls {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Apply a sort value to the native sort <select>, keeping the Random
|
||||||
|
* option's value in sync when the persisted value carries a seed
|
||||||
|
* (e.g. "random:abc123"). Must be used instead of assigning
|
||||||
|
* sortSelect.value directly whenever the value may be a seeded random
|
||||||
|
* sort, otherwise the native select has no matching option.
|
||||||
|
* @param {string} sortValue - Sort value like "name:asc" or "random:<seed>"
|
||||||
|
*/
|
||||||
|
applySortToSelect(sortValue) {
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
if (!sortSelect) return;
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
|
||||||
|
if (randomOpt) {
|
||||||
|
randomOpt.value = String(sortValue).startsWith('random') ? sortValue : 'random';
|
||||||
|
}
|
||||||
|
sortSelect.value = sortValue;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Generate a fresh seeded random sort value ("random:<seed>") and keep
|
||||||
|
* the native <select> in sync so its value matches the persisted sort
|
||||||
|
* string and the dropdown shows the selected label.
|
||||||
|
* @returns {string} The new sort value, e.g. "random:abc123xyz"
|
||||||
|
*/
|
||||||
|
_randomizeSortValue() {
|
||||||
|
const seed = Math.random().toString(36).slice(2, 12);
|
||||||
|
const value = `random:${seed}`;
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
if (sortSelect) {
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
|
||||||
|
if (randomOpt) {
|
||||||
|
randomOpt.value = value;
|
||||||
|
}
|
||||||
|
sortSelect.value = value;
|
||||||
|
}
|
||||||
|
return value;
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Load sort preference from storage
|
* Load sort preference from storage
|
||||||
*/
|
*/
|
||||||
@@ -326,10 +374,7 @@ export class PageControls {
|
|||||||
// Handle legacy format conversion
|
// Handle legacy format conversion
|
||||||
const convertedSort = this.convertLegacySortFormat(savedSort);
|
const convertedSort = this.convertLegacySortFormat(savedSort);
|
||||||
this.pageState.sortBy = convertedSort;
|
this.pageState.sortBy = convertedSort;
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
this.applySortToSelect(convertedSort);
|
||||||
if (sortSelect) {
|
|
||||||
sortSelect.value = convertedSort;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -523,9 +568,9 @@ export class PageControls {
|
|||||||
this.pageState.sortBy = restoredSort;
|
this.pageState.sortBy = restoredSort;
|
||||||
this.saveSortPreference(restoredSort);
|
this.saveSortPreference(restoredSort);
|
||||||
this._removeVlmSortOption();
|
this._removeVlmSortOption();
|
||||||
|
this.applySortToSelect(restoredSort);
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
if (sortSelect) {
|
if (sortSelect) {
|
||||||
sortSelect.value = restoredSort;
|
|
||||||
sortSelect.disabled = false;
|
sortSelect.disabled = false;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -575,10 +620,7 @@ export class PageControls {
|
|||||||
const savedGroupedSort = getStorageItem(groupedKey);
|
const savedGroupedSort = getStorageItem(groupedKey);
|
||||||
if (savedGroupedSort) {
|
if (savedGroupedSort) {
|
||||||
this.pageState.sortBy = savedGroupedSort;
|
this.pageState.sortBy = savedGroupedSort;
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
this.applySortToSelect(savedGroupedSort);
|
||||||
if (sortSelect) {
|
|
||||||
sortSelect.value = savedGroupedSort;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// Leaving group mode: persist current sort for next time, restore non-group sort
|
// Leaving group mode: persist current sort for next time, restore non-group sort
|
||||||
@@ -586,10 +628,7 @@ export class PageControls {
|
|||||||
const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
|
const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
|
||||||
if (savedNormalSort) {
|
if (savedNormalSort) {
|
||||||
this.pageState.sortBy = savedNormalSort;
|
this.pageState.sortBy = savedNormalSort;
|
||||||
const sortSelect = document.getElementById('sortSelect');
|
this.applySortToSelect(savedNormalSort);
|
||||||
if (sortSelect) {
|
|
||||||
sortSelect.value = savedNormalSort;
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -874,7 +913,7 @@ export class PageControls {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if (sortSelect) {
|
if (sortSelect) {
|
||||||
sortSelect.value = this.pageState.sortBy;
|
this.applySortToSelect(this.pageState.sortBy);
|
||||||
}
|
}
|
||||||
if (searchInput) {
|
if (searchInput) {
|
||||||
searchInput.value = this.pageState.filters?.search || '';
|
searchInput.value = this.pageState.filters?.search || '';
|
||||||
|
|||||||
@@ -96,7 +96,16 @@ export function initSortDropdown(select) {
|
|||||||
};
|
};
|
||||||
|
|
||||||
const choose = (value) => {
|
const choose = (value) => {
|
||||||
if (select.value === value) return;
|
if (select.value === value) {
|
||||||
|
// Re-picking the already-selected option is normally a no-op,
|
||||||
|
// matching native <select> behavior. The seeded Random sort is
|
||||||
|
// the exception: clicking it again should reshuffle, so let the
|
||||||
|
// change handler (PageControls) generate a fresh seed.
|
||||||
|
if (String(value).startsWith('random')) {
|
||||||
|
select.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
select.value = value;
|
select.value = value;
|
||||||
select.dispatchEvent(new Event('change', { bubbles: true }));
|
select.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
};
|
};
|
||||||
@@ -277,9 +286,10 @@ export function initSortDropdown(select) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Rebuild the menu when <option>s change (VLM adds/removes a temporary
|
// Rebuild the menu when <option>s change (VLM adds/removes a temporary
|
||||||
// option at runtime).
|
// option at runtime, and the seeded Random sort option gets a new value
|
||||||
|
// attribute each time it is picked).
|
||||||
const observer = new MutationObserver(() => buildMenu());
|
const observer = new MutationObserver(() => buildMenu());
|
||||||
observer.observe(select, { childList: true });
|
observer.observe(select, { childList: true, subtree: true, attributes: true, attributeFilter: ['value'] });
|
||||||
|
|
||||||
buildMenu();
|
buildMenu();
|
||||||
group.dataset.sortReady = '1';
|
group.dataset.sortReady = '1';
|
||||||
|
|||||||
@@ -27,6 +27,8 @@ export class BulkManager {
|
|||||||
|
|
||||||
// Drag detection properties
|
// Drag detection properties
|
||||||
this.dragThreshold = 5; // Pixels to move before considering it a drag
|
this.dragThreshold = 5; // Pixels to move before considering it a drag
|
||||||
|
this.dragDelayMs = 100; // Minimum hold time before a drag is treated as a marquee
|
||||||
|
this.minMarqueeSize = 10; // Minimum drag box (px) before a marquee counts as a selection
|
||||||
this.mouseDownTime = 0;
|
this.mouseDownTime = 0;
|
||||||
this.mouseDownPosition = { x: 0, y: 0 };
|
this.mouseDownPosition = { x: 0, y: 0 };
|
||||||
|
|
||||||
@@ -88,7 +90,7 @@ export class BulkManager {
|
|||||||
moveAll: true,
|
moveAll: true,
|
||||||
autoOrganize: false,
|
autoOrganize: false,
|
||||||
deleteAll: true,
|
deleteAll: true,
|
||||||
setContentRating: false,
|
setContentRating: true,
|
||||||
skipMetadataRefresh: false,
|
skipMetadataRefresh: false,
|
||||||
setFavorite: true,
|
setFavorite: true,
|
||||||
unfavorite: true,
|
unfavorite: true,
|
||||||
@@ -173,6 +175,19 @@ export class BulkManager {
|
|||||||
});
|
});
|
||||||
|
|
||||||
eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => {
|
eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => {
|
||||||
|
// Only track marquee/drag while the left button is physically held.
|
||||||
|
// mouseup can be missed (release outside the window, focus loss, driver quirks),
|
||||||
|
// so mousemove must verify the button state itself instead of relying on it.
|
||||||
|
if (!(e.buttons & 1)) {
|
||||||
|
if (this.isMarqueeActive) {
|
||||||
|
this.endMarqueeSelection(e);
|
||||||
|
} else {
|
||||||
|
this.mouseDownTime = 0;
|
||||||
|
this.isDragging = false;
|
||||||
|
}
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
|
||||||
if (this.isMarqueeActive) {
|
if (this.isMarqueeActive) {
|
||||||
this.lastClientX = e.clientX;
|
this.lastClientX = e.clientX;
|
||||||
this.lastClientY = e.clientY;
|
this.lastClientY = e.clientY;
|
||||||
@@ -184,7 +199,10 @@ export class BulkManager {
|
|||||||
const dy = e.clientY - this.mouseDownPosition.y;
|
const dy = e.clientY - this.mouseDownPosition.y;
|
||||||
const distance = Math.sqrt(dx * dx + dy * dy);
|
const distance = Math.sqrt(dx * dx + dy * dy);
|
||||||
|
|
||||||
if (distance >= this.dragThreshold) {
|
// Require both enough movement AND enough hold time so quick
|
||||||
|
// click jitter from micro-movement input devices is not a marquee.
|
||||||
|
const heldTime = Date.now() - this.mouseDownTime;
|
||||||
|
if (heldTime >= this.dragDelayMs && distance >= this.dragThreshold) {
|
||||||
this.isDragging = true;
|
this.isDragging = true;
|
||||||
this.startMarqueeSelection(e, true);
|
this.startMarqueeSelection(e, true);
|
||||||
}
|
}
|
||||||
@@ -1510,14 +1528,18 @@ export class BulkManager {
|
|||||||
let failureCount = 0;
|
let failureCount = 0;
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const isRecipesPage = state.currentPageType === 'recipes';
|
||||||
for (const filePath of targets) {
|
for (const filePath of targets) {
|
||||||
if (cancelled) {
|
if (cancelled) {
|
||||||
showToast('toast.api.operationCancelled', {}, 'info');
|
showToast('toast.api.operationCancelled', {}, 'info');
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
try {
|
try {
|
||||||
await apiClient.saveModelMetadata(filePath, { preview_nsfw_level: level });
|
if (isRecipesPage) {
|
||||||
|
await updateRecipeMetadata(filePath, { preview_nsfw_level: level });
|
||||||
|
} else {
|
||||||
|
await getModelApiClient().saveModelMetadata(filePath, { preview_nsfw_level: level });
|
||||||
|
}
|
||||||
successCount++;
|
successCount++;
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
failureCount++;
|
failureCount++;
|
||||||
@@ -1958,9 +1980,31 @@ export class BulkManager {
|
|||||||
// Remove visual feedback class
|
// Remove visual feedback class
|
||||||
document.body.classList.remove('marquee-selecting');
|
document.body.classList.remove('marquee-selecting');
|
||||||
|
|
||||||
|
// Compute the actual drag box size in document coordinates, matching how
|
||||||
|
// updateMarqueeSelectionFromPosition tracks the rectangle. Client-space
|
||||||
|
// size would wrongly flag auto-scroll marquees (tiny pointer movement,
|
||||||
|
// large document-space box) as accidental clicks.
|
||||||
|
const container = document.querySelector('.page-content');
|
||||||
|
const scrollX = container?.scrollLeft || 0;
|
||||||
|
const scrollY = container?.scrollTop || 0;
|
||||||
|
const dragWidth = Math.abs((e.clientX + scrollX) - this.marqueeStartDoc.x);
|
||||||
|
const dragHeight = Math.abs((e.clientY + scrollY) - this.marqueeStartDoc.y);
|
||||||
|
const isTinyMarquee = dragWidth < this.minMarqueeSize && dragHeight < this.minMarqueeSize;
|
||||||
|
|
||||||
// Get selection count
|
// Get selection count
|
||||||
const selectionCount = state.selectedModels.size;
|
const selectionCount = state.selectedModels.size;
|
||||||
|
|
||||||
|
// A tiny box (e.g. click jitter that happened to graze a card) is treated
|
||||||
|
// as an accidental click: undo any selection and leave bulk mode.
|
||||||
|
if (isTinyMarquee) {
|
||||||
|
this.clearSelection();
|
||||||
|
if (state.bulkMode) {
|
||||||
|
this.toggleBulkMode();
|
||||||
|
}
|
||||||
|
this.initialSelectedModels.clear();
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
// If no models were selected, exit bulk mode
|
// If no models were selected, exit bulk mode
|
||||||
if (selectionCount === 0) {
|
if (selectionCount === 0) {
|
||||||
if (state.bulkMode) {
|
if (state.bulkMode) {
|
||||||
|
|||||||
@@ -1517,11 +1517,20 @@ export class SettingsManager {
|
|||||||
return data;
|
return data;
|
||||||
}
|
}
|
||||||
|
|
||||||
async loadLoraRoots() {
|
showNoRootsPlaceholder(select) {
|
||||||
try {
|
select.innerHTML = '';
|
||||||
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot');
|
const option = document.createElement('option');
|
||||||
if (!defaultLoraRootSelect) return;
|
option.value = '';
|
||||||
|
option.textContent = translate('settings.folderSettings.noDefault', {}, 'No Default');
|
||||||
|
select.appendChild(option);
|
||||||
|
select.disabled = true;
|
||||||
|
}
|
||||||
|
|
||||||
|
async loadLoraRoots() {
|
||||||
|
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot');
|
||||||
|
if (!defaultLoraRootSelect) return;
|
||||||
|
|
||||||
|
try {
|
||||||
// Fetch lora roots
|
// Fetch lora roots
|
||||||
const response = await fetch('/api/lm/loras/roots');
|
const response = await fetch('/api/lm/loras/roots');
|
||||||
if (!response.ok) {
|
if (!response.ok) {
|
||||||
@@ -1530,10 +1539,12 @@ export class SettingsManager {
|
|||||||
|
|
||||||
const data = await response.json();
|
const data = await response.json();
|
||||||
if (!data.roots || data.roots.length === 0) {
|
if (!data.roots || data.roots.length === 0) {
|
||||||
throw new Error('No LoRA roots found');
|
this.showNoRootsPlaceholder(defaultLoraRootSelect);
|
||||||
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
defaultLoraRootSelect.innerHTML = '';
|
defaultLoraRootSelect.innerHTML = '';
|
||||||
|
defaultLoraRootSelect.disabled = false;
|
||||||
|
|
||||||
// Add options for each root
|
// Add options for each root
|
||||||
data.roots.forEach(root => {
|
data.roots.forEach(root => {
|
||||||
@@ -1548,15 +1559,16 @@ export class SettingsManager {
|
|||||||
|
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error loading LoRA roots:', error);
|
console.error('Error loading LoRA roots:', error);
|
||||||
|
this.showNoRootsPlaceholder(defaultLoraRootSelect);
|
||||||
showToast('toast.settings.loraRootsFailed', { message: error.message }, 'error');
|
showToast('toast.settings.loraRootsFailed', { message: error.message }, 'error');
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async loadCheckpointRoots() {
|
async loadCheckpointRoots() {
|
||||||
try {
|
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
|
||||||
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
|
if (!defaultCheckpointRootSelect) return;
|
||||||
if (!defaultCheckpointRootSelect) return;
|
|
||||||
|
|
||||||
|
try {
|
||||||
// Fetch checkpoint roots (checkpoint paths only, not unet)
|
// Fetch checkpoint roots (checkpoint paths only, not unet)
|
||||||
const response = await fetch('/api/lm/checkpoints/checkpoints_roots');
|
const response = await fetch('/api/lm/checkpoints/checkpoints_roots');
|
||||||
if (!response.ok) {
|
if (!response.ok) {
|
||||||
@@ -1565,10 +1577,12 @@ export class SettingsManager {
|
|||||||
|
|
||||||
const data = await response.json();
|
const data = await response.json();
|
||||||
if (!data.roots || data.roots.length === 0) {
|
if (!data.roots || data.roots.length === 0) {
|
||||||
throw new Error('No checkpoint roots found');
|
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
|
||||||
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
defaultCheckpointRootSelect.innerHTML = '';
|
defaultCheckpointRootSelect.innerHTML = '';
|
||||||
|
defaultCheckpointRootSelect.disabled = false;
|
||||||
|
|
||||||
// Add options for each root
|
// Add options for each root
|
||||||
data.roots.forEach(root => {
|
data.roots.forEach(root => {
|
||||||
@@ -1583,15 +1597,16 @@ export class SettingsManager {
|
|||||||
|
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error loading checkpoint roots:', error);
|
console.error('Error loading checkpoint roots:', error);
|
||||||
|
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
|
||||||
showToast('toast.settings.checkpointRootsFailed', { message: error.message }, 'error');
|
showToast('toast.settings.checkpointRootsFailed', { message: error.message }, 'error');
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async loadUnetRoots() {
|
async loadUnetRoots() {
|
||||||
try {
|
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
|
||||||
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
|
if (!defaultUnetRootSelect) return;
|
||||||
if (!defaultUnetRootSelect) return;
|
|
||||||
|
|
||||||
|
try {
|
||||||
// Fetch unet roots (diffusion model paths only)
|
// Fetch unet roots (diffusion model paths only)
|
||||||
const response = await fetch('/api/lm/checkpoints/unet_roots');
|
const response = await fetch('/api/lm/checkpoints/unet_roots');
|
||||||
if (!response.ok) {
|
if (!response.ok) {
|
||||||
@@ -1600,10 +1615,12 @@ export class SettingsManager {
|
|||||||
|
|
||||||
const data = await response.json();
|
const data = await response.json();
|
||||||
if (!data.roots || data.roots.length === 0) {
|
if (!data.roots || data.roots.length === 0) {
|
||||||
throw new Error('No diffusion model roots found');
|
this.showNoRootsPlaceholder(defaultUnetRootSelect);
|
||||||
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
defaultUnetRootSelect.innerHTML = '';
|
defaultUnetRootSelect.innerHTML = '';
|
||||||
|
defaultUnetRootSelect.disabled = false;
|
||||||
|
|
||||||
// Add options for each root
|
// Add options for each root
|
||||||
data.roots.forEach(root => {
|
data.roots.forEach(root => {
|
||||||
@@ -1618,15 +1635,16 @@ export class SettingsManager {
|
|||||||
|
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error loading diffusion model roots:', error);
|
console.error('Error loading diffusion model roots:', error);
|
||||||
|
this.showNoRootsPlaceholder(defaultUnetRootSelect);
|
||||||
showToast('toast.settings.unetRootsFailed', { message: error.message }, 'error');
|
showToast('toast.settings.unetRootsFailed', { message: error.message }, 'error');
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async loadEmbeddingRoots() {
|
async loadEmbeddingRoots() {
|
||||||
try {
|
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
|
||||||
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
|
if (!defaultEmbeddingRootSelect) return;
|
||||||
if (!defaultEmbeddingRootSelect) return;
|
|
||||||
|
|
||||||
|
try {
|
||||||
// Fetch embedding roots
|
// Fetch embedding roots
|
||||||
const response = await fetch('/api/lm/embeddings/roots');
|
const response = await fetch('/api/lm/embeddings/roots');
|
||||||
if (!response.ok) {
|
if (!response.ok) {
|
||||||
@@ -1635,10 +1653,12 @@ export class SettingsManager {
|
|||||||
|
|
||||||
const data = await response.json();
|
const data = await response.json();
|
||||||
if (!data.roots || data.roots.length === 0) {
|
if (!data.roots || data.roots.length === 0) {
|
||||||
throw new Error('No embedding roots found');
|
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
|
||||||
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
defaultEmbeddingRootSelect.innerHTML = '';
|
defaultEmbeddingRootSelect.innerHTML = '';
|
||||||
|
defaultEmbeddingRootSelect.disabled = false;
|
||||||
|
|
||||||
// Add options for each root
|
// Add options for each root
|
||||||
data.roots.forEach(root => {
|
data.roots.forEach(root => {
|
||||||
@@ -1653,6 +1673,7 @@ export class SettingsManager {
|
|||||||
|
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error loading embedding roots:', error);
|
console.error('Error loading embedding roots:', error);
|
||||||
|
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
|
||||||
showToast('toast.settings.embeddingRootsFailed', { message: error.message }, 'error');
|
showToast('toast.settings.embeddingRootsFailed', { message: error.message }, 'error');
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,11 +1,12 @@
|
|||||||
import { modalManager } from './ModalManager.js';
|
import { modalManager } from './ModalManager.js';
|
||||||
import {
|
import {
|
||||||
getStorageItem,
|
getStorageItem,
|
||||||
setStorageItem,
|
setStorageItem,
|
||||||
getStoredVersionInfo,
|
getStoredVersionInfo,
|
||||||
setStoredVersionInfo,
|
setStoredVersionInfo,
|
||||||
isVersionMatch
|
isVersionMatch
|
||||||
} from '../utils/storageHelpers.js';
|
} from '../utils/storageHelpers.js';
|
||||||
|
import { state } from '../state/index.js';
|
||||||
import { bannerService } from './BannerService.js';
|
import { bannerService } from './BannerService.js';
|
||||||
import { translate } from '../utils/i18nHelpers.js';
|
import { translate } from '../utils/i18nHelpers.js';
|
||||||
|
|
||||||
@@ -26,6 +27,8 @@ export class UpdateService {
|
|||||||
this.isUpdating = false;
|
this.isUpdating = false;
|
||||||
this.channelMode = null;
|
this.channelMode = null;
|
||||||
this.hasGit = false;
|
this.hasGit = false;
|
||||||
|
this.nightlyNotifyDate = getStorageItem('nightly_notify_date', '');
|
||||||
|
this.nightlyBadgeShown = false;
|
||||||
this.progressKeepVisible = false;
|
this.progressKeepVisible = false;
|
||||||
this.currentVersionInfo = null;
|
this.currentVersionInfo = null;
|
||||||
this.versionMismatch = false;
|
this.versionMismatch = false;
|
||||||
@@ -59,9 +62,6 @@ export class UpdateService {
|
|||||||
|
|
||||||
// Perform update check if needed
|
// Perform update check if needed
|
||||||
this.checkVersionInfo().then(() => {
|
this.checkVersionInfo().then(() => {
|
||||||
if (this.channelMode === null) {
|
|
||||||
this.channelMode = this.hasGit ? 'nightly' : 'release';
|
|
||||||
}
|
|
||||||
this.checkForUpdates().then(() => {
|
this.checkForUpdates().then(() => {
|
||||||
this.updateBadgeVisibility();
|
this.updateBadgeVisibility();
|
||||||
});
|
});
|
||||||
@@ -118,6 +118,14 @@ export class UpdateService {
|
|||||||
|
|
||||||
if (data.success) {
|
if (data.success) {
|
||||||
this.channelMode = channel;
|
this.channelMode = channel;
|
||||||
|
// Persist channel preference to settings.json
|
||||||
|
fetch('/api/lm/settings', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({ update_channel: channel })
|
||||||
|
}).then(r => {
|
||||||
|
if (!r.ok) console.warn('Failed to persist update channel:', r.status);
|
||||||
|
}).catch(e => console.warn('Failed to persist update channel:', e));
|
||||||
await this.checkForUpdates({ force: true });
|
await this.checkForUpdates({ force: true });
|
||||||
this.updateModalContent();
|
this.updateModalContent();
|
||||||
this.updateChannelUI();
|
this.updateChannelUI();
|
||||||
@@ -154,6 +162,20 @@ export class UpdateService {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
_resolveChannelFromSettings() {
|
||||||
|
const stored = state?.global?.settings?.update_channel;
|
||||||
|
if (stored === 'nightly' || stored === 'release') {
|
||||||
|
return stored;
|
||||||
|
}
|
||||||
|
if (!this.hasGit) {
|
||||||
|
return 'release';
|
||||||
|
}
|
||||||
|
if (this.gitInfo?.branch === 'detached') {
|
||||||
|
return 'release';
|
||||||
|
}
|
||||||
|
return 'nightly';
|
||||||
|
}
|
||||||
|
|
||||||
async _confirmChannelSwitch(titleKey, messageKey) {
|
async _confirmChannelSwitch(titleKey, messageKey) {
|
||||||
return new Promise((resolve) => {
|
return new Promise((resolve) => {
|
||||||
const title = translate(titleKey);
|
const title = translate(titleKey);
|
||||||
@@ -475,6 +497,18 @@ export class UpdateService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async checkForUpdates({ force = false } = {}) {
|
async checkForUpdates({ force = false } = {}) {
|
||||||
|
let needsMigration = false;
|
||||||
|
if (this.channelMode === null) {
|
||||||
|
const stored = state?.global?.settings?.update_channel;
|
||||||
|
if (stored === 'nightly' || stored === 'release') {
|
||||||
|
this.channelMode = stored;
|
||||||
|
} else if (!this.hasGit) {
|
||||||
|
this.channelMode = 'release';
|
||||||
|
needsMigration = true;
|
||||||
|
}
|
||||||
|
// hasGit=true with no stored value: wait for gitInfo.branch
|
||||||
|
}
|
||||||
|
|
||||||
if (!force && !this.updateNotificationsEnabled) {
|
if (!force && !this.updateNotificationsEnabled) {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -493,7 +527,7 @@ export class UpdateService {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
// Call backend API to check for updates with nightly flag
|
// Call backend API to check for updates with nightly flag
|
||||||
const nightly = this.channelMode === 'nightly';
|
const nightly = (this.channelMode ?? (this.hasGit ? 'nightly' : 'release')) === 'nightly';
|
||||||
const response = await fetch(`/api/lm/check-updates?nightly=${nightly}`);
|
const response = await fetch(`/api/lm/check-updates?nightly=${nightly}`);
|
||||||
const data = await response.json();
|
const data = await response.json();
|
||||||
|
|
||||||
@@ -503,12 +537,28 @@ export class UpdateService {
|
|||||||
this.updateInfo = data;
|
this.updateInfo = data;
|
||||||
this.gitInfo = data.git_info || this.gitInfo;
|
this.gitInfo = data.git_info || this.gitInfo;
|
||||||
this.hasGit = data.has_git || false;
|
this.hasGit = data.has_git || false;
|
||||||
if (this.channelMode === null) {
|
|
||||||
this.channelMode = this.hasGit ? 'nightly' : 'release';
|
if (needsMigration || this.channelMode === null) {
|
||||||
|
this.channelMode = this._resolveChannelFromSettings();
|
||||||
|
if (state?.global?.settings) {
|
||||||
|
state.global.settings.update_channel = this.channelMode;
|
||||||
|
}
|
||||||
|
fetch('/api/lm/settings', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: { 'Content-Type': 'application/json' },
|
||||||
|
body: JSON.stringify({ update_channel: this.channelMode })
|
||||||
|
}).then(r => {
|
||||||
|
if (!r.ok) console.warn('Failed to persist update channel:', r.status);
|
||||||
|
}).catch(e => console.warn('Failed to persist update channel:', e));
|
||||||
}
|
}
|
||||||
|
|
||||||
this.updateAvailable = data.update_available;
|
this.updateAvailable = data.update_available;
|
||||||
|
|
||||||
|
// Nightly channel: surface the update badge at most once per calendar day.
|
||||||
|
if (this.updateAvailable && this.channelMode === 'nightly' && this.nightlyNotifyDate !== this._getTodayKey()) {
|
||||||
|
this._markNightlyNotified();
|
||||||
|
}
|
||||||
|
|
||||||
this.lastCheckTime = now;
|
this.lastCheckTime = now;
|
||||||
setStorageItem('last_update_check', now.toString());
|
setStorageItem('last_update_check', now.toString());
|
||||||
|
|
||||||
@@ -558,6 +608,28 @@ export class UpdateService {
|
|||||||
|
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
_getTodayKey() {
|
||||||
|
const now = new Date();
|
||||||
|
const month = String(now.getMonth() + 1).padStart(2, '0');
|
||||||
|
const day = String(now.getDate()).padStart(2, '0');
|
||||||
|
return `${now.getFullYear()}-${month}-${day}`;
|
||||||
|
}
|
||||||
|
|
||||||
|
_isNightlyBadgeAllowed() {
|
||||||
|
if (this.channelMode !== 'nightly') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
// Keep the badge visible for the rest of the session once shown, but do
|
||||||
|
// not show it again on later sessions within the same calendar day.
|
||||||
|
return this.nightlyNotifyDate !== this._getTodayKey() || this.nightlyBadgeShown;
|
||||||
|
}
|
||||||
|
|
||||||
|
_markNightlyNotified() {
|
||||||
|
this.nightlyNotifyDate = this._getTodayKey();
|
||||||
|
this.nightlyBadgeShown = true;
|
||||||
|
setStorageItem('nightly_notify_date', this.nightlyNotifyDate);
|
||||||
|
}
|
||||||
|
|
||||||
updateBadgeVisibility() {
|
updateBadgeVisibility() {
|
||||||
const updateToggle = document.querySelector('.update-toggle');
|
const updateToggle = document.querySelector('.update-toggle');
|
||||||
@@ -566,9 +638,12 @@ export class UpdateService {
|
|||||||
? bannerService.getUnreadBannerCount()
|
? bannerService.getUnreadBannerCount()
|
||||||
: 0;
|
: 0;
|
||||||
|
|
||||||
|
// Force updating badges visibility based on current state
|
||||||
|
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable && this._isNightlyBadgeAllowed();
|
||||||
|
|
||||||
if (updateToggle) {
|
if (updateToggle) {
|
||||||
let tooltipKey = 'header.actions.notifications';
|
let tooltipKey = 'header.actions.notifications';
|
||||||
if (this.updateNotificationsEnabled && this.updateAvailable) {
|
if (shouldShowUpdate) {
|
||||||
tooltipKey = 'update.updateAvailable';
|
tooltipKey = 'update.updateAvailable';
|
||||||
} else if (unreadBanners > 0) {
|
} else if (unreadBanners > 0) {
|
||||||
tooltipKey = 'update.tabs.messages';
|
tooltipKey = 'update.tabs.messages';
|
||||||
@@ -576,8 +651,6 @@ export class UpdateService {
|
|||||||
updateToggle.title = translate(tooltipKey);
|
updateToggle.title = translate(tooltipKey);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Force updating badges visibility based on current state
|
|
||||||
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable;
|
|
||||||
const shouldShow = shouldShowUpdate || unreadBanners > 0;
|
const shouldShow = shouldShowUpdate || unreadBanners > 0;
|
||||||
|
|
||||||
if (updateBadge) {
|
if (updateBadge) {
|
||||||
|
|||||||
@@ -48,6 +48,11 @@
|
|||||||
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
|
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
|
||||||
</optgroup>
|
</optgroup>
|
||||||
{% endif %}
|
{% endif %}
|
||||||
|
{% if page_id != 'recipes' %}
|
||||||
|
<optgroup label="{{ t('loras.controls.sort.random', default='Random') }}">
|
||||||
|
<option value="random">{{ t('loras.controls.sort.randomAction', default='Randomize (shuffle)') }}</option>
|
||||||
|
</optgroup>
|
||||||
|
{% endif %}
|
||||||
{% if page_id == 'recipes' %}
|
{% if page_id == 'recipes' %}
|
||||||
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
|
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
|
||||||
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
|
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
|
||||||
|
|||||||
@@ -0,0 +1,221 @@
|
|||||||
|
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
|
||||||
|
|
||||||
|
const resetAndReloadMock = vi.fn();
|
||||||
|
const getModelApiClientMock = vi.fn();
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||||
|
getModelApiClient: getModelApiClientMock,
|
||||||
|
resetAndReload: resetAndReloadMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: vi.fn(),
|
||||||
|
openCivitaiByMetadata: vi.fn(),
|
||||||
|
updatePanelPositions: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/managers/DownloadManager.js', () => ({
|
||||||
|
downloadManager: { showDownloadModal: vi.fn() },
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/SidebarManager.js', () => ({
|
||||||
|
sidebarManager: {
|
||||||
|
setHostPageControls: vi.fn(),
|
||||||
|
initialize: vi.fn(async () => {}),
|
||||||
|
refresh: vi.fn(async () => {}),
|
||||||
|
cleanup: vi.fn(),
|
||||||
|
isInitialized: false,
|
||||||
|
},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/alphabet/index.js', () => ({
|
||||||
|
createAlphabetBar: vi.fn(() => ({ destroy: vi.fn() })),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/updateCheckHelpers.js', () => ({
|
||||||
|
performModelUpdateCheck: vi.fn(async () => ({ status: 'success', displayName: 'LoRA', records: [] })),
|
||||||
|
}));
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.resetModules();
|
||||||
|
vi.clearAllMocks();
|
||||||
|
localStorage.clear();
|
||||||
|
sessionStorage.clear();
|
||||||
|
|
||||||
|
resetAndReloadMock.mockResolvedValue(undefined);
|
||||||
|
getModelApiClientMock.mockReturnValue({});
|
||||||
|
|
||||||
|
global.fetch = vi.fn().mockResolvedValue({
|
||||||
|
ok: true,
|
||||||
|
json: async () => ({ success: true, base_models: [] }),
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
delete window.bulkManager;
|
||||||
|
delete window.modelDuplicatesManager;
|
||||||
|
delete global.fetch;
|
||||||
|
});
|
||||||
|
|
||||||
|
function renderControlsDom(pageKey) {
|
||||||
|
document.body.dataset.page = pageKey;
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div class="controls">
|
||||||
|
<div id="excludedViewBanner" class="excluded-view-banner hidden">
|
||||||
|
<button id="excludedViewBackBtn">Back</button>
|
||||||
|
</div>
|
||||||
|
<div class="actions">
|
||||||
|
<div class="action-buttons">
|
||||||
|
<div class="control-group">
|
||||||
|
<select id="sortSelect">
|
||||||
|
<option value="name:asc">Name Asc</option>
|
||||||
|
<option value="name:desc">Name Desc</option>
|
||||||
|
<option value="random">Randomize (shuffle)</option>
|
||||||
|
</select>
|
||||||
|
</div>
|
||||||
|
<div class="control-group dropdown-group">
|
||||||
|
<button data-action="refresh" class="dropdown-main"></button>
|
||||||
|
<button class="dropdown-toggle"></button>
|
||||||
|
<div class="dropdown-menu">
|
||||||
|
<div class="dropdown-item" data-action="full-rebuild"></div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="fetch"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="download"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="bulk"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button data-action="find-duplicates"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group">
|
||||||
|
<button id="favoriteFilterBtn" class="favorite-filter"></button>
|
||||||
|
</div>
|
||||||
|
<div class="control-group dropdown-group update-filter-group">
|
||||||
|
<button id="updateFilterBtn" class="dropdown-main update-filter" aria-busy="false">
|
||||||
|
<span>Updates</span>
|
||||||
|
</button>
|
||||||
|
<button id="updateFilterMenuToggle" class="dropdown-toggle"></button>
|
||||||
|
<div class="dropdown-menu">
|
||||||
|
<div id="checkUpdatesMenuItem" class="dropdown-item" data-action="check-updates">
|
||||||
|
<span>Check updates</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div id="customFilterIndicator" class="control-group hidden">
|
||||||
|
<div class="filter-active">
|
||||||
|
<span class="customFilterText" title=""></span>
|
||||||
|
<i class="fas fa-times-circle clear-filter"></i>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div id="breadcrumbContainer"></div>
|
||||||
|
<div id="duplicatesBanner" style="display: none;"></div>
|
||||||
|
<div class="alphabet-bar-container"></div>
|
||||||
|
`;
|
||||||
|
}
|
||||||
|
|
||||||
|
async function createControls() {
|
||||||
|
const stateModule = await import('../../../static/js/state/index.js');
|
||||||
|
stateModule.initPageState('loras');
|
||||||
|
const { LorasControls } = await import('../../../static/js/components/controls/LorasControls.js');
|
||||||
|
return { stateModule, controls: new LorasControls() };
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('Random sort option', () => {
|
||||||
|
it('generates a seeded sort value when Random is picked', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
|
||||||
|
expect(localStorage.getItem('lora_manager_loras_sort')).toBe(controls.pageState.sortBy);
|
||||||
|
expect(randomOpt.value).toBe(controls.pageState.sortBy);
|
||||||
|
expect(sortSelect.value).toBe(controls.pageState.sortBy);
|
||||||
|
expect(resetAndReloadMock).toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('reshuffles with a fresh seed every time Random is picked again', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
// First pick
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
const firstSeed = controls.pageState.sortBy;
|
||||||
|
|
||||||
|
// Second pick: the option now carries the seeded value, like a menu click
|
||||||
|
sortSelect.value = randomOpt.value;
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
|
||||||
|
expect(controls.pageState.sortBy).not.toBe(firstSeed);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('restores a persisted seeded random sort on load', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const savedSort = 'random:persistedseed';
|
||||||
|
localStorage.setItem('lora_manager_loras_sort', savedSort);
|
||||||
|
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toBe(savedSort);
|
||||||
|
expect(sortSelect.value).toBe(savedSort);
|
||||||
|
expect(sortSelect.querySelector('option[value="random:persistedseed"]')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('applies a non-random sort back to the plain random option', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
// Seed a random sort, then switch to a normal sort
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
controls.applySortToSelect('name:desc');
|
||||||
|
|
||||||
|
expect(sortSelect.value).toBe('name:desc');
|
||||||
|
expect(randomOpt.value).toBe('random');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('resets the seeded option when switching away from Random via the dropdown change handler', async () => {
|
||||||
|
renderControlsDom('loras');
|
||||||
|
const { controls } = await createControls();
|
||||||
|
const sortSelect = document.getElementById('sortSelect');
|
||||||
|
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||||
|
|
||||||
|
// Pick Random: the option is now seeded
|
||||||
|
sortSelect.value = 'random';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
expect(randomOpt.value).toMatch(/^random:[a-z0-9]+$/);
|
||||||
|
|
||||||
|
// Switch to a non-random sort through the change handler (as a menu
|
||||||
|
// click does); the option must go back to the plain "random" value
|
||||||
|
sortSelect.value = 'name:desc';
|
||||||
|
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
expect(controls.pageState.sortBy).toBe('name:desc');
|
||||||
|
expect(sortSelect.value).toBe('name:desc');
|
||||||
|
expect(randomOpt.value).toBe('random');
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,68 @@
|
|||||||
|
import { describe, it, beforeEach, expect } from 'vitest';
|
||||||
|
import { initSortDropdown } from '../../../static/js/components/controls/SortDropdown.js';
|
||||||
|
|
||||||
|
function renderSortDropdownDom() {
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div class="sort-dropdown-group">
|
||||||
|
<select id="sortSelect">
|
||||||
|
<option value="name:asc">Name Asc</option>
|
||||||
|
<option value="name:desc">Name Desc</option>
|
||||||
|
<option value="random" selected>Randomize (shuffle)</option>
|
||||||
|
</select>
|
||||||
|
<button class="sort-trigger" type="button">
|
||||||
|
<span class="sort-trigger__label"></span>
|
||||||
|
</button>
|
||||||
|
<div class="sort-dropdown-menu"></div>
|
||||||
|
</div>
|
||||||
|
`;
|
||||||
|
return {
|
||||||
|
select: document.getElementById('sortSelect'),
|
||||||
|
menu: document.querySelector('.sort-dropdown-menu'),
|
||||||
|
label: document.querySelector('.sort-trigger__label'),
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('SortDropdown menu sync', () => {
|
||||||
|
let select;
|
||||||
|
let menu;
|
||||||
|
let label;
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
({ select, menu, label } = renderSortDropdownDom());
|
||||||
|
initSortDropdown(select);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('rebuilds the menu and highlights the selected item when an option value attribute changes', async () => {
|
||||||
|
// The seeded Random option gets a new value each time it is picked.
|
||||||
|
// The select's value getter follows the selected option's new value.
|
||||||
|
const randomOpt = select.querySelector('option[value="random"]');
|
||||||
|
randomOpt.value = 'random:abc123';
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const items = [...menu.querySelectorAll('.sort-option')];
|
||||||
|
expect(items.map((el) => el.dataset.value)).toContain('random:abc123');
|
||||||
|
const seededItem = items.find((el) => el.dataset.value === 'random:abc123');
|
||||||
|
expect(seededItem.classList.contains('is-selected')).toBe(true);
|
||||||
|
expect(label.textContent).toBe('Randomize (shuffle)');
|
||||||
|
});
|
||||||
|
|
||||||
|
it('drops the stale seeded item and re-selects the plain random item when the option is reset', async () => {
|
||||||
|
const randomOpt = select.querySelector('option[value="random"]');
|
||||||
|
randomOpt.value = 'random:abc123';
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
// The rebuild must have happened: the seeded item is in the menu
|
||||||
|
const seededItems = [...menu.querySelectorAll('.sort-option')]
|
||||||
|
.filter((el) => el.dataset.value === 'random:abc123');
|
||||||
|
expect(seededItems).toHaveLength(1);
|
||||||
|
|
||||||
|
// PageControls resets the option to "random" when switching away
|
||||||
|
randomOpt.value = 'random';
|
||||||
|
await Promise.resolve();
|
||||||
|
|
||||||
|
const items = [...menu.querySelectorAll('.sort-option')];
|
||||||
|
expect(items.map((el) => el.dataset.value)).not.toContain('random:abc123');
|
||||||
|
const randomItem = items.find((el) => el.dataset.value === 'random');
|
||||||
|
expect(randomItem.classList.contains('is-selected')).toBe(true);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,195 @@
|
|||||||
|
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||||
|
|
||||||
|
const { APP_MODULE, EXTENSION_MODULE, appMock, registeredExtensions } =
|
||||||
|
vi.hoisted(() => {
|
||||||
|
const registeredExtensions = [];
|
||||||
|
const appMock = {
|
||||||
|
configuringGraph: false,
|
||||||
|
registerExtension: (ext) => registeredExtensions.push(ext),
|
||||||
|
};
|
||||||
|
return {
|
||||||
|
APP_MODULE: new URL("../../../scripts/app.js", import.meta.url).pathname,
|
||||||
|
EXTENSION_MODULE: new URL(
|
||||||
|
"../../../web/comfyui/lora_stack_dynamic_inputs.js",
|
||||||
|
import.meta.url
|
||||||
|
).pathname,
|
||||||
|
appMock,
|
||||||
|
registeredExtensions,
|
||||||
|
};
|
||||||
|
});
|
||||||
|
|
||||||
|
vi.mock(APP_MODULE, () => ({
|
||||||
|
app: appMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
describe("Lora Stack Combiner dynamic inputs", () => {
|
||||||
|
let extension;
|
||||||
|
|
||||||
|
beforeEach(async () => {
|
||||||
|
vi.resetModules();
|
||||||
|
registeredExtensions.length = 0;
|
||||||
|
appMock.configuringGraph = false;
|
||||||
|
await import(EXTENSION_MODULE);
|
||||||
|
extension = registeredExtensions.find(
|
||||||
|
(ext) => ext.name === "Comfy.LoraManager.LoraStackCombiner"
|
||||||
|
);
|
||||||
|
expect(extension).toBeDefined();
|
||||||
|
});
|
||||||
|
|
||||||
|
function createNodeType() {
|
||||||
|
const nodeType = { prototype: {} };
|
||||||
|
extension.beforeRegisterNodeDef(
|
||||||
|
nodeType,
|
||||||
|
{ name: "Lora Stack Combiner (LoraManager)" },
|
||||||
|
appMock
|
||||||
|
);
|
||||||
|
return nodeType;
|
||||||
|
}
|
||||||
|
|
||||||
|
function createNode(inputs = []) {
|
||||||
|
const node = {
|
||||||
|
comfyClass: "Lora Stack Combiner (LoraManager)",
|
||||||
|
inputs: inputs.map((name) => ({ name, type: "LORA_STACK" })),
|
||||||
|
addInput: vi.fn(function (name, type, opts) {
|
||||||
|
this.inputs.push({ name, type, ...opts });
|
||||||
|
}),
|
||||||
|
removeInput: vi.fn(function (index) {
|
||||||
|
this.inputs.splice(index, 1);
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
return node;
|
||||||
|
}
|
||||||
|
|
||||||
|
function makeLinkInfo() {
|
||||||
|
return { id: 999, origin_id: 1, target_id: 2 };
|
||||||
|
}
|
||||||
|
|
||||||
|
it("adds a third input when the last slot gets connected", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
"lora_stack3",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("does not add an input when a non-last slot gets connected", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 0, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
"lora_stack3",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("removes a disconnected middle slot and renumbers", () => {
|
||||||
|
// Simulates a real LiteGraph disconnect event: it fires only for slots that
|
||||||
|
// had a link, and input.link has already been cleared before the event fires.
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||||
|
node.inputs[0].link = 11;
|
||||||
|
node.inputs[1].link = null; // slot 2 was just disconnected
|
||||||
|
node.inputs[2].link = 13;
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, false, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("keeps the last slot when it is disconnected", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||||
|
node.inputs[0].link = 11;
|
||||||
|
node.inputs[1].link = 12;
|
||||||
|
node.inputs[2].link = null; // last slot was just disconnected
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 2, false, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
"lora_stack3",
|
||||||
|
]);
|
||||||
|
expect(node.removeInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("keeps at least two inputs when disconnecting", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||||
|
node.inputs[0].link = 11;
|
||||||
|
node.inputs[1].link = null; // slot 2 was just disconnected
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, false, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
expect(node.removeInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("does nothing while the graph is being configured", () => {
|
||||||
|
appMock.configuringGraph = true;
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 1, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
expect(node.addInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("leaves legacy lora_stack_a/b inputs untouched", () => {
|
||||||
|
const nodeType = createNodeType();
|
||||||
|
const node = createNode(["lora_stack_a", "lora_stack_b"]);
|
||||||
|
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
node.onConnectionsChange(1, 0, true, makeLinkInfo());
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack_a",
|
||||||
|
"lora_stack_b",
|
||||||
|
]);
|
||||||
|
expect(node.addInput).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
|
||||||
|
it("ensures two numbered inputs exist on creation", () => {
|
||||||
|
const node = createNode([]);
|
||||||
|
extension.nodeCreated(node, {});
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack1",
|
||||||
|
"lora_stack2",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
|
||||||
|
it("does not add numbered inputs to legacy workflows", () => {
|
||||||
|
const node = createNode(["lora_stack_a", "lora_stack_b"]);
|
||||||
|
extension.nodeCreated(node, {});
|
||||||
|
|
||||||
|
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||||
|
"lora_stack_a",
|
||||||
|
"lora_stack_b",
|
||||||
|
]);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
import { describe, it, beforeEach, expect, vi } from 'vitest';
|
||||||
|
|
||||||
|
const showToastMock = vi.fn();
|
||||||
|
const translateMock = vi.fn((key, params, fallback) => (typeof fallback === 'string' ? fallback : key));
|
||||||
|
const getNSFWLevelNameMock = vi.fn((level) => {
|
||||||
|
if (level >= 16) return 'XXX';
|
||||||
|
if (level >= 8) return 'X';
|
||||||
|
if (level >= 4) return 'R';
|
||||||
|
if (level >= 2) return 'PG13';
|
||||||
|
if (level >= 1) return 'PG';
|
||||||
|
return 'Unknown';
|
||||||
|
});
|
||||||
|
|
||||||
|
const loadingManagerStub = {
|
||||||
|
showSimpleLoading: vi.fn(),
|
||||||
|
showCancelButton: vi.fn(),
|
||||||
|
hide: vi.fn(),
|
||||||
|
};
|
||||||
|
|
||||||
|
const stateStub = {
|
||||||
|
currentPageType: 'recipes',
|
||||||
|
bulkMode: false,
|
||||||
|
selectedModels: new Set(),
|
||||||
|
loadingManager: loadingManagerStub,
|
||||||
|
virtualScroller: { updateSingleItem: vi.fn() },
|
||||||
|
global: { settings: {} },
|
||||||
|
};
|
||||||
|
|
||||||
|
const saveModelMetadataMock = vi.fn();
|
||||||
|
const getModelApiClientMock = vi.fn(() => ({ saveModelMetadata: saveModelMetadataMock }));
|
||||||
|
const updateRecipeMetadataMock = vi.fn(() => Promise.resolve({ success: true }));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/state/index.js', () => ({
|
||||||
|
state: stateStub,
|
||||||
|
getCurrentPageState: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
copyToClipboard: vi.fn(),
|
||||||
|
sendLoraToWorkflow: vi.fn(),
|
||||||
|
sendEmbeddingToWorkflow: vi.fn(),
|
||||||
|
buildLoraSyntax: vi.fn(),
|
||||||
|
getNSFWLevelName: getNSFWLevelNameMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||||
|
getModelApiClient: getModelApiClientMock,
|
||||||
|
resetAndReload: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
||||||
|
RecipeSidebarApiClient: class {},
|
||||||
|
updateRecipeMetadata: updateRecipeMetadataMock,
|
||||||
|
extractRecipeId: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/apiConfig.js', () => ({
|
||||||
|
MODEL_TYPES: { LORA: 'loras', CHECKPOINT: 'checkpoints', EMBEDDING: 'embeddings' },
|
||||||
|
MODEL_CONFIG: {},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/managers/ModalManager.js', () => ({
|
||||||
|
modalManager: { showModal: vi.fn(), closeModal: vi.fn() },
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/ModelCard.js', () => ({
|
||||||
|
updateCardsForBulkMode: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/i18nHelpers.js', () => ({
|
||||||
|
translate: translateMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/priorityTagHelpers.js', () => ({
|
||||||
|
getPriorityTagSuggestions: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/NsfwLevelSelector.js', () => ({
|
||||||
|
getNsfwLevelSelector: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
describe('BulkManager bulk content rating', () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks();
|
||||||
|
stateStub.currentPageType = 'recipes';
|
||||||
|
stateStub.bulkMode = false;
|
||||||
|
stateStub.selectedModels.clear();
|
||||||
|
saveModelMetadataMock.mockResolvedValue(undefined);
|
||||||
|
updateRecipeMetadataMock.mockResolvedValue({ success: true });
|
||||||
|
});
|
||||||
|
|
||||||
|
async function createBulkManager() {
|
||||||
|
const { BulkManager } = await import('../../../static/js/managers/BulkManager.js');
|
||||||
|
return new BulkManager();
|
||||||
|
}
|
||||||
|
|
||||||
|
it('exposes the content rating action on the recipes page action config', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
expect(bulk.actionConfig.recipes.setContentRating).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('persists the rating through the recipe API when on the recipes page', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
stateStub.currentPageType = 'recipes';
|
||||||
|
stateStub.selectedModels.add('/recipes/test.webp');
|
||||||
|
|
||||||
|
const ok = await bulk.setBulkContentRating(4, ['/recipes/test.webp']);
|
||||||
|
|
||||||
|
expect(ok).toBe(true);
|
||||||
|
expect(updateRecipeMetadataMock).toHaveBeenCalledWith('/recipes/test.webp', { preview_nsfw_level: 4 });
|
||||||
|
expect(updateRecipeMetadataMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(saveModelMetadataMock).not.toHaveBeenCalled();
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith(
|
||||||
|
'toast.models.bulkContentRatingSet',
|
||||||
|
{ count: 1, level: 'R' },
|
||||||
|
'success'
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('persists the rating through the model API on model pages', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
stateStub.currentPageType = 'loras';
|
||||||
|
stateStub.selectedModels.add('/models/test.safetensors');
|
||||||
|
|
||||||
|
const ok = await bulk.setBulkContentRating(8, ['/models/test.safetensors']);
|
||||||
|
|
||||||
|
expect(ok).toBe(true);
|
||||||
|
expect(saveModelMetadataMock).toHaveBeenCalledWith('/models/test.safetensors', { preview_nsfw_level: 8 });
|
||||||
|
expect(saveModelMetadataMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(updateRecipeMetadataMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -0,0 +1,186 @@
|
|||||||
|
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
|
||||||
|
import { state } from '../../../static/js/state/index.js';
|
||||||
|
import { MODEL_TYPES } from '../../../static/js/api/apiConfig.js';
|
||||||
|
import { eventManager } from '../../../static/js/utils/EventManager.js';
|
||||||
|
import { BulkManager } from '../../../static/js/managers/BulkManager.js';
|
||||||
|
|
||||||
|
function fire(type, init = {}) {
|
||||||
|
return new MouseEvent(type, { bubbles: true, cancelable: true, ...init });
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('BulkManager marquee guards', () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.useFakeTimers();
|
||||||
|
// jsdom may not provide requestAnimationFrame; stub it so the auto-scroll loop is a no-op.
|
||||||
|
window.requestAnimationFrame = vi.fn();
|
||||||
|
window.cancelAnimationFrame = vi.fn();
|
||||||
|
|
||||||
|
eventManager.cleanup();
|
||||||
|
state.currentPageType = MODEL_TYPES.LORA;
|
||||||
|
state.bulkMode = false;
|
||||||
|
state.selectedModels.clear();
|
||||||
|
|
||||||
|
document.body.innerHTML = '<div class="page-content"></div>';
|
||||||
|
const pageContent = document.querySelector('.page-content');
|
||||||
|
pageContent.getBoundingClientRect = () => ({
|
||||||
|
top: 0,
|
||||||
|
left: 0,
|
||||||
|
right: 1000,
|
||||||
|
bottom: 1000,
|
||||||
|
width: 1000,
|
||||||
|
height: 1000,
|
||||||
|
x: 0,
|
||||||
|
y: 0,
|
||||||
|
toJSON: () => ({}),
|
||||||
|
});
|
||||||
|
pageContent.scrollBy = vi.fn();
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
eventManager.cleanup();
|
||||||
|
vi.useRealTimers();
|
||||||
|
document.body.innerHTML = '';
|
||||||
|
});
|
||||||
|
|
||||||
|
function createBulkManager() {
|
||||||
|
const bulk = new BulkManager();
|
||||||
|
bulk.initialize();
|
||||||
|
return bulk;
|
||||||
|
}
|
||||||
|
|
||||||
|
it('never starts a marquee when the left button is not held', () => {
|
||||||
|
const bulk = createBulkManager();
|
||||||
|
const pageContent = document.querySelector('.page-content');
|
||||||
|
|
||||||
|
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 10, clientY: 10 }));
|
||||||
|
document.dispatchEvent(fire('mousemove', { buttons: 0, clientX: 50, clientY: 50 }));
|
||||||
|
|
||||||
|
expect(bulk.mouseDownTime).toBe(0);
|
||||||
|
expect(bulk.isMarqueeActive).toBe(false);
|
||||||
|
expect(state.bulkMode).toBe(false);
|
||||||
|
expect(document.querySelector('.marquee-selection')).toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('requires holding the left button for the drag delay before starting a marquee', () => {
|
||||||
|
const bulk = createBulkManager();
|
||||||
|
const pageContent = document.querySelector('.page-content');
|
||||||
|
|
||||||
|
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 10, clientY: 10 }));
|
||||||
|
|
||||||
|
// Fast movement: far enough, but too soon after mousedown.
|
||||||
|
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 30, clientY: 10 }));
|
||||||
|
expect(state.bulkMode).toBe(false);
|
||||||
|
expect(bulk.isMarqueeActive).toBe(false);
|
||||||
|
|
||||||
|
// Once the hold time has elapsed, the same drag qualifies.
|
||||||
|
vi.advanceTimersByTime(100);
|
||||||
|
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 35, clientY: 12 }));
|
||||||
|
expect(state.bulkMode).toBe(true);
|
||||||
|
expect(bulk.isMarqueeActive).toBe(true);
|
||||||
|
expect(document.querySelector('.marquee-selection')).not.toBeNull();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('ends an active marquee if the left button is released without a mouseup event', () => {
|
||||||
|
const bulk = createBulkManager();
|
||||||
|
bulk.mouseDownPosition = { x: 10, y: 10 };
|
||||||
|
bulk.startMarqueeSelection({}, true);
|
||||||
|
expect(state.bulkMode).toBe(true);
|
||||||
|
expect(document.querySelector('.marquee-selection')).not.toBeNull();
|
||||||
|
|
||||||
|
// No mouseup was dispatched; a plain move with the button released finalizes it.
|
||||||
|
document.dispatchEvent(fire('mousemove', { buttons: 0, clientX: 50, clientY: 50 }));
|
||||||
|
|
||||||
|
expect(bulk.isMarqueeActive).toBe(false);
|
||||||
|
expect(document.querySelector('.marquee-selection')).toBeNull();
|
||||||
|
expect(state.bulkMode).toBe(false); // zero selected -> auto-exit
|
||||||
|
});
|
||||||
|
|
||||||
|
it('treats a tiny marquee as an accidental click: clears selection and exits bulk mode', () => {
|
||||||
|
const bulk = createBulkManager();
|
||||||
|
const card = document.createElement('div');
|
||||||
|
card.className = 'model-card selected';
|
||||||
|
card.dataset.filepath = '/models/test.safetensors';
|
||||||
|
document.body.appendChild(card);
|
||||||
|
state.selectedModels.add('/models/test.safetensors');
|
||||||
|
|
||||||
|
bulk.mouseDownPosition = { x: 100, y: 100 };
|
||||||
|
bulk.startMarqueeSelection({}, true);
|
||||||
|
expect(state.bulkMode).toBe(true);
|
||||||
|
|
||||||
|
bulk.endMarqueeSelection({ clientX: 103, clientY: 104 });
|
||||||
|
|
||||||
|
expect(state.bulkMode).toBe(false);
|
||||||
|
expect(state.selectedModels.size).toBe(0);
|
||||||
|
expect(card.classList.contains('selected')).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('keeps selection and bulk mode when the marquee is large enough', () => {
|
||||||
|
const bulk = createBulkManager();
|
||||||
|
const card = document.createElement('div');
|
||||||
|
card.className = 'model-card selected';
|
||||||
|
card.dataset.filepath = '/models/test.safetensors';
|
||||||
|
document.body.appendChild(card);
|
||||||
|
state.selectedModels.add('/models/test.safetensors');
|
||||||
|
|
||||||
|
bulk.mouseDownPosition = { x: 100, y: 100 };
|
||||||
|
bulk.startMarqueeSelection({}, true);
|
||||||
|
|
||||||
|
bulk.endMarqueeSelection({ clientX: 130, clientY: 140 });
|
||||||
|
|
||||||
|
expect(state.bulkMode).toBe(true);
|
||||||
|
expect(state.selectedModels.has('/models/test.safetensors')).toBe(true);
|
||||||
|
expect(card.classList.contains('selected')).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('keeps auto-scroll marquee selections when the pointer only moved a few pixels', () => {
|
||||||
|
const bulk = createBulkManager();
|
||||||
|
const pageContent = document.querySelector('.page-content');
|
||||||
|
|
||||||
|
// Card just below the press point in document coordinates.
|
||||||
|
const card = document.createElement('div');
|
||||||
|
card.className = 'model-card';
|
||||||
|
card.dataset.filepath = '/models/off-screen.safetensors';
|
||||||
|
card.getBoundingClientRect = () => ({
|
||||||
|
top: 950,
|
||||||
|
left: 400,
|
||||||
|
right: 600,
|
||||||
|
bottom: 1050,
|
||||||
|
width: 200,
|
||||||
|
height: 100,
|
||||||
|
x: 400,
|
||||||
|
y: 950,
|
||||||
|
toJSON: () => ({}),
|
||||||
|
});
|
||||||
|
document.body.appendChild(card);
|
||||||
|
|
||||||
|
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 500, clientY: 900 }));
|
||||||
|
vi.advanceTimersByTime(100);
|
||||||
|
|
||||||
|
// Small pointer move: enough to start the marquee, but under minMarqueeSize.
|
||||||
|
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 506, clientY: 906 }));
|
||||||
|
expect(bulk.isMarqueeActive).toBe(true);
|
||||||
|
|
||||||
|
// Auto-scroll grows the document-space box while the pointer stays nearly still.
|
||||||
|
pageContent.scrollTop = 200;
|
||||||
|
card.getBoundingClientRect = () => ({
|
||||||
|
top: 750,
|
||||||
|
left: 400,
|
||||||
|
right: 600,
|
||||||
|
bottom: 850,
|
||||||
|
width: 200,
|
||||||
|
height: 100,
|
||||||
|
x: 400,
|
||||||
|
y: 750,
|
||||||
|
toJSON: () => ({}),
|
||||||
|
});
|
||||||
|
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 506, clientY: 906 }));
|
||||||
|
|
||||||
|
expect(state.selectedModels.has('/models/off-screen.safetensors')).toBe(true);
|
||||||
|
|
||||||
|
// Release: the client-space box is tiny, but the document-space box is not.
|
||||||
|
document.dispatchEvent(fire('mouseup', { button: 0, clientX: 506, clientY: 906 }));
|
||||||
|
|
||||||
|
expect(state.selectedModels.has('/models/off-screen.safetensors')).toBe(true);
|
||||||
|
expect(state.bulkMode).toBe(true);
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -106,6 +106,118 @@ afterEach(() => {
|
|||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
|
describe('SettingsManager root selects', () => {
|
||||||
|
const rootCases = [
|
||||||
|
{
|
||||||
|
method: 'loadLoraRoots',
|
||||||
|
selectId: 'defaultLoraRoot',
|
||||||
|
endpoint: '/api/lm/loras/roots',
|
||||||
|
errorKey: 'toast.settings.loraRootsFailed',
|
||||||
|
},
|
||||||
|
{
|
||||||
|
method: 'loadCheckpointRoots',
|
||||||
|
selectId: 'defaultCheckpointRoot',
|
||||||
|
endpoint: '/api/lm/checkpoints/checkpoints_roots',
|
||||||
|
errorKey: 'toast.settings.checkpointRootsFailed',
|
||||||
|
},
|
||||||
|
{
|
||||||
|
method: 'loadUnetRoots',
|
||||||
|
selectId: 'defaultUnetRoot',
|
||||||
|
endpoint: '/api/lm/checkpoints/unet_roots',
|
||||||
|
errorKey: 'toast.settings.unetRootsFailed',
|
||||||
|
},
|
||||||
|
{
|
||||||
|
method: 'loadEmbeddingRoots',
|
||||||
|
selectId: 'defaultEmbeddingRoot',
|
||||||
|
endpoint: '/api/lm/embeddings/roots',
|
||||||
|
errorKey: 'toast.settings.embeddingRootsFailed',
|
||||||
|
},
|
||||||
|
];
|
||||||
|
|
||||||
|
const appendRootSelect = (id) => {
|
||||||
|
const select = document.createElement('select');
|
||||||
|
select.id = id;
|
||||||
|
document.body.appendChild(select);
|
||||||
|
return select;
|
||||||
|
};
|
||||||
|
|
||||||
|
it.each(rootCases)(
|
||||||
|
'populates the $method select with roots and keeps it enabled',
|
||||||
|
async ({ method, selectId, endpoint }) => {
|
||||||
|
const manager = createManager();
|
||||||
|
const select = appendRootSelect(selectId);
|
||||||
|
select.disabled = true;
|
||||||
|
|
||||||
|
global.fetch = vi.fn().mockResolvedValue({
|
||||||
|
ok: true,
|
||||||
|
json: async () => ({
|
||||||
|
success: true,
|
||||||
|
roots: ['/models/root-a', '/models/root-b'],
|
||||||
|
}),
|
||||||
|
});
|
||||||
|
|
||||||
|
await manager[method]();
|
||||||
|
|
||||||
|
expect(global.fetch).toHaveBeenCalledWith(endpoint);
|
||||||
|
expect(Array.from(select.options).map(option => option.value)).toEqual([
|
||||||
|
'/models/root-a',
|
||||||
|
'/models/root-b',
|
||||||
|
]);
|
||||||
|
expect(select.disabled).toBe(false);
|
||||||
|
expect(showToast).not.toHaveBeenCalled();
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
it.each(rootCases)(
|
||||||
|
'shows a placeholder and no error toast when $method has empty roots',
|
||||||
|
async ({ method, selectId, endpoint }) => {
|
||||||
|
const manager = createManager();
|
||||||
|
const select = appendRootSelect(selectId);
|
||||||
|
|
||||||
|
global.fetch = vi.fn().mockResolvedValue({
|
||||||
|
ok: true,
|
||||||
|
json: async () => ({
|
||||||
|
success: true,
|
||||||
|
roots: [],
|
||||||
|
}),
|
||||||
|
});
|
||||||
|
|
||||||
|
await manager[method]();
|
||||||
|
|
||||||
|
expect(global.fetch).toHaveBeenCalledWith(endpoint);
|
||||||
|
expect(select.options).toHaveLength(1);
|
||||||
|
expect(select.options[0].value).toBe('');
|
||||||
|
expect(select.options[0].textContent).toBe('No Default');
|
||||||
|
expect(select.disabled).toBe(true);
|
||||||
|
expect(showToast).not.toHaveBeenCalled();
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
it.each(rootCases)(
|
||||||
|
'shows an error toast when the $method roots request fails',
|
||||||
|
async ({ method, selectId, errorKey }) => {
|
||||||
|
const manager = createManager();
|
||||||
|
const select = appendRootSelect(selectId);
|
||||||
|
|
||||||
|
global.fetch = vi.fn().mockResolvedValue({
|
||||||
|
ok: false,
|
||||||
|
status: 500,
|
||||||
|
});
|
||||||
|
|
||||||
|
await manager[method]();
|
||||||
|
|
||||||
|
expect(select.options).toHaveLength(1);
|
||||||
|
expect(select.options[0].value).toBe('');
|
||||||
|
expect(select.disabled).toBe(true);
|
||||||
|
expect(showToast).toHaveBeenCalledWith(
|
||||||
|
errorKey,
|
||||||
|
expect.objectContaining({ message: expect.any(String) }),
|
||||||
|
'error',
|
||||||
|
);
|
||||||
|
}
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
describe('SettingsManager library controls', () => {
|
describe('SettingsManager library controls', () => {
|
||||||
it('loads libraries and populates the select', async () => {
|
it('loads libraries and populates the select', async () => {
|
||||||
const manager = createManager();
|
const manager = createManager();
|
||||||
|
|||||||
@@ -1,12 +1,26 @@
|
|||||||
import { describe, beforeEach, afterEach, expect, it, vi } from 'vitest';
|
import { describe, beforeEach, afterEach, expect, it, vi } from 'vitest';
|
||||||
import { UpdateService } from '../../../static/js/managers/UpdateService.js';
|
import { UpdateService } from '../../../static/js/managers/UpdateService.js';
|
||||||
|
import { state } from '../../../static/js/state/index.js';
|
||||||
|
|
||||||
function createFetchResponse(payload) {
|
function createFetchResponse(payload) {
|
||||||
return {
|
return {
|
||||||
json: vi.fn().mockResolvedValue(payload)
|
json: vi.fn().mockResolvedValue(payload),
|
||||||
|
ok: true,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function stubSettingsUpdateChannel(channel) {
|
||||||
|
state.global = state.global || {};
|
||||||
|
state.global.settings = state.global.settings || {};
|
||||||
|
state.global.settings.update_channel = channel;
|
||||||
|
}
|
||||||
|
|
||||||
|
function clearSettingsUpdateChannel() {
|
||||||
|
if (state.global?.settings) {
|
||||||
|
delete state.global.settings.update_channel;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
describe('UpdateService passive checks', () => {
|
describe('UpdateService passive checks', () => {
|
||||||
let service;
|
let service;
|
||||||
let fetchMock;
|
let fetchMock;
|
||||||
@@ -16,10 +30,13 @@ describe('UpdateService passive checks', () => {
|
|||||||
success: true,
|
success: true,
|
||||||
current_version: 'v1.0.0',
|
current_version: 'v1.0.0',
|
||||||
latest_version: 'v1.0.0',
|
latest_version: 'v1.0.0',
|
||||||
git_info: { short_hash: 'abc123' }
|
git_info: { short_hash: 'abc123' },
|
||||||
|
has_git: true,
|
||||||
}));
|
}));
|
||||||
global.fetch = fetchMock;
|
global.fetch = fetchMock;
|
||||||
|
|
||||||
|
stubSettingsUpdateChannel('release');
|
||||||
|
|
||||||
service = new UpdateService();
|
service = new UpdateService();
|
||||||
service.updateNotificationsEnabled = false;
|
service.updateNotificationsEnabled = false;
|
||||||
service.lastCheckTime = 0;
|
service.lastCheckTime = 0;
|
||||||
@@ -28,6 +45,7 @@ describe('UpdateService passive checks', () => {
|
|||||||
|
|
||||||
afterEach(() => {
|
afterEach(() => {
|
||||||
delete global.fetch;
|
delete global.fetch;
|
||||||
|
clearSettingsUpdateChannel();
|
||||||
});
|
});
|
||||||
|
|
||||||
it('skips passive update checks when notifications are disabled', async () => {
|
it('skips passive update checks when notifications are disabled', async () => {
|
||||||
@@ -43,3 +61,106 @@ describe('UpdateService passive checks', () => {
|
|||||||
expect(fetchMock).toHaveBeenCalledWith('/api/lm/check-updates?nightly=false');
|
expect(fetchMock).toHaveBeenCalledWith('/api/lm/check-updates?nightly=false');
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
|
describe('UpdateService nightly notification throttling', () => {
|
||||||
|
let fetchMock;
|
||||||
|
let updateToggle;
|
||||||
|
let updateBadge;
|
||||||
|
|
||||||
|
function stubUpdateBadgeDom() {
|
||||||
|
updateToggle = document.createElement('div');
|
||||||
|
updateToggle.className = 'update-toggle';
|
||||||
|
updateBadge = document.createElement('span');
|
||||||
|
updateBadge.className = 'update-badge';
|
||||||
|
updateToggle.appendChild(updateBadge);
|
||||||
|
document.body.appendChild(updateToggle);
|
||||||
|
|
||||||
|
vi.spyOn(document, 'querySelector').mockImplementation((selector) => {
|
||||||
|
if (selector === '.update-toggle') return updateToggle;
|
||||||
|
if (selector === '.update-toggle .update-badge') return updateBadge;
|
||||||
|
return null;
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
function makeUpdateResponse(channel) {
|
||||||
|
return {
|
||||||
|
success: true,
|
||||||
|
current_version: 'v1.0.0',
|
||||||
|
latest_version: channel === 'nightly' ? 'main-abc1234' : 'v1.1.0',
|
||||||
|
update_available: true,
|
||||||
|
git_info: { short_hash: 'abc123' },
|
||||||
|
has_git: true,
|
||||||
|
nightly: channel === 'nightly',
|
||||||
|
changelog: ['test: change'],
|
||||||
|
releases: [],
|
||||||
|
behind_by: 3,
|
||||||
|
commit_date: '2026-07-31',
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
beforeEach(() => {
|
||||||
|
fetchMock = vi.fn().mockResolvedValue(createFetchResponse(makeUpdateResponse('release')));
|
||||||
|
global.fetch = fetchMock;
|
||||||
|
stubUpdateBadgeDom();
|
||||||
|
});
|
||||||
|
|
||||||
|
afterEach(() => {
|
||||||
|
vi.restoreAllMocks();
|
||||||
|
delete global.fetch;
|
||||||
|
});
|
||||||
|
|
||||||
|
it('shows the nightly badge once and keeps it visible for the session', async () => {
|
||||||
|
stubSettingsUpdateChannel('nightly');
|
||||||
|
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('nightly')));
|
||||||
|
|
||||||
|
const service = new UpdateService();
|
||||||
|
service.updateNotificationsEnabled = true;
|
||||||
|
|
||||||
|
await service.checkForUpdates({ force: true });
|
||||||
|
|
||||||
|
expect(service.updateAvailable).toBe(true);
|
||||||
|
expect(service.nightlyBadgeShown).toBe(true);
|
||||||
|
expect(service.nightlyNotifyDate).toBe(service._getTodayKey());
|
||||||
|
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||||
|
|
||||||
|
// A repeated check within the same session keeps the badge visible.
|
||||||
|
await service.checkForUpdates({ force: true });
|
||||||
|
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('suppresses the nightly badge on a later session in the same day', async () => {
|
||||||
|
stubSettingsUpdateChannel('nightly');
|
||||||
|
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('nightly')));
|
||||||
|
|
||||||
|
const firstService = new UpdateService();
|
||||||
|
firstService.updateNotificationsEnabled = true;
|
||||||
|
await firstService.checkForUpdates({ force: true });
|
||||||
|
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||||
|
|
||||||
|
// Simulate a fresh page session on the same calendar day.
|
||||||
|
const secondService = new UpdateService();
|
||||||
|
secondService.updateNotificationsEnabled = true;
|
||||||
|
await secondService.checkForUpdates({ force: true });
|
||||||
|
|
||||||
|
expect(secondService.updateAvailable).toBe(true);
|
||||||
|
expect(secondService.nightlyBadgeShown).toBe(false);
|
||||||
|
expect(updateBadge.classList.contains('visible')).toBe(false);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('is not affected by the daily limit on the release channel', async () => {
|
||||||
|
stubSettingsUpdateChannel('release');
|
||||||
|
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('release')));
|
||||||
|
|
||||||
|
const firstService = new UpdateService();
|
||||||
|
firstService.updateNotificationsEnabled = true;
|
||||||
|
await firstService.checkForUpdates({ force: true });
|
||||||
|
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||||
|
|
||||||
|
const secondService = new UpdateService();
|
||||||
|
secondService.updateNotificationsEnabled = true;
|
||||||
|
await secondService.checkForUpdates({ force: true });
|
||||||
|
|
||||||
|
expect(secondService.updateAvailable).toBe(true);
|
||||||
|
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|||||||
@@ -1,4 +1,11 @@
|
|||||||
from py.nodes.lora_stack_combiner import LoraStackCombinerLM
|
import types
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from py.nodes.lora_stack_combiner import (
|
||||||
|
LoraStackCombinerLM,
|
||||||
|
_LoraStackOptionalInputs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_combine_stacks_preserves_order():
|
def test_combine_stacks_preserves_order():
|
||||||
@@ -49,3 +56,109 @@ def test_combine_stacks_allows_duplicate_entries():
|
|||||||
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
|
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
|
||||||
|
|
||||||
assert combined_stack == [duplicate_entry, duplicate_entry]
|
assert combined_stack == [duplicate_entry, duplicate_entry]
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_returns_empty_when_both_unconnected():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks()
|
||||||
|
|
||||||
|
assert combined_stack == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_returns_other_when_one_unconnected():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
|
||||||
|
(combined_stack_a,) = node.combine_stacks(lora_stack1=stack_a)
|
||||||
|
(combined_stack_b,) = node.combine_stacks(lora_stack2=stack_a)
|
||||||
|
|
||||||
|
assert combined_stack_a == stack_a
|
||||||
|
assert combined_stack_b == stack_a
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_with_dynamic_third_slot():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||||
|
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(
|
||||||
|
lora_stack1=stack_a, lora_stack2=stack_b, lora_stack3=stack_c
|
||||||
|
)
|
||||||
|
|
||||||
|
assert combined_stack == stack_a + stack_b + stack_c
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_orders_by_slot_number_not_call_order():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||||
|
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(
|
||||||
|
lora_stack3=stack_c, lora_stack2=stack_b, lora_stack1=stack_a
|
||||||
|
)
|
||||||
|
|
||||||
|
assert combined_stack == stack_a + stack_b + stack_c
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_accepts_only_dynamic_slot():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(lora_stack3=stack_c)
|
||||||
|
|
||||||
|
assert combined_stack == stack_c
|
||||||
|
|
||||||
|
|
||||||
|
def test_combine_stacks_handles_legacy_input_names():
|
||||||
|
node = LoraStackCombinerLM()
|
||||||
|
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||||
|
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||||
|
|
||||||
|
(combined_stack,) = node.combine_stacks(lora_stack_a=stack_a, lora_stack_b=stack_b)
|
||||||
|
|
||||||
|
assert combined_stack == stack_a + stack_b
|
||||||
|
|
||||||
|
|
||||||
|
def test_input_types_exposes_two_default_slots():
|
||||||
|
input_types = LoraStackCombinerLM.INPUT_TYPES()
|
||||||
|
|
||||||
|
assert set(input_types["optional"]) == {"lora_stack1", "lora_stack2"}
|
||||||
|
assert input_types["optional"]["lora_stack1"][0] == "LORA_STACK"
|
||||||
|
assert input_types["optional"]["lora_stack2"][0] == "LORA_STACK"
|
||||||
|
|
||||||
|
|
||||||
|
def test_input_types_recognizes_dynamic_slots_from_get_input_info(monkeypatch):
|
||||||
|
frames = [None, None, types.SimpleNamespace(function="get_input_info")]
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"py.nodes.lora_stack_combiner.inspect.stack", lambda: frames
|
||||||
|
)
|
||||||
|
|
||||||
|
input_types = LoraStackCombinerLM.INPUT_TYPES()
|
||||||
|
optional = input_types["optional"]
|
||||||
|
|
||||||
|
assert "lora_stack3" in optional
|
||||||
|
assert optional["lora_stack3"][0] == "LORA_STACK"
|
||||||
|
assert "lora_stack25" in optional
|
||||||
|
assert optional["lora_stack25"][0] == "LORA_STACK"
|
||||||
|
|
||||||
|
|
||||||
|
def test_lora_stack_optional_inputs_proxy():
|
||||||
|
proxy = _LoraStackOptionalInputs({"lora_stack1": ("LORA_STACK", {})})
|
||||||
|
|
||||||
|
assert "lora_stack1" in proxy
|
||||||
|
assert "lora_stack2" in proxy
|
||||||
|
assert "lora_stack10" in proxy
|
||||||
|
assert "lora_stack_a" in proxy
|
||||||
|
assert "lora_stack" not in proxy
|
||||||
|
assert "lora_stacka" not in proxy
|
||||||
|
assert "lora_stack_1" not in proxy
|
||||||
|
assert "text" not in proxy
|
||||||
|
|
||||||
|
assert proxy["lora_stack1"][0] == "LORA_STACK"
|
||||||
|
assert proxy["lora_stack5"][0] == "LORA_STACK"
|
||||||
|
|
||||||
|
with pytest.raises(KeyError):
|
||||||
|
proxy["not_a_stack"]
|
||||||
|
|||||||
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
|
|||||||
async def get_model_versions(self, _model_id):
|
async def get_model_versions(self, _model_id):
|
||||||
return {"modelVersions": [], "name": "", "type": "lora"}
|
return {"modelVersions": [], "name": "", "type": "lora"}
|
||||||
|
|
||||||
async def get_user_models(self, _username):
|
async def get_user_models(self, _username, cursor=None):
|
||||||
return []
|
return {"items": [], "nextCursor": None}
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, _username):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
class FakeUserModelsProvider(FakeMetadataProvider):
|
class FakeUserModelsProvider(FakeMetadataProvider):
|
||||||
def __init__(self, models):
|
def __init__(self, models, next_cursor=None, estimated_total=None):
|
||||||
self.models = models
|
self.models = models
|
||||||
|
self.next_cursor = next_cursor
|
||||||
|
self.estimated_total = estimated_total
|
||||||
self.received_usernames: list[str] = []
|
self.received_usernames: list[str] = []
|
||||||
|
self.received_cursors: list = []
|
||||||
|
|
||||||
async def get_user_models(self, username):
|
async def get_user_models(self, username, cursor=None):
|
||||||
self.received_usernames.append(username)
|
self.received_usernames.append(username)
|
||||||
return self.models
|
self.received_cursors.append(cursor)
|
||||||
|
return {"items": self.models, "nextCursor": self.next_cursor}
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, _username):
|
||||||
|
return self.estimated_total
|
||||||
|
|
||||||
|
|
||||||
async def fake_metadata_provider_factory():
|
async def fake_metadata_provider_factory():
|
||||||
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
|
|||||||
assert "username" in payload["error"].lower()
|
assert "username" in payload["error"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_civitai_user_models_returns_pagination_fields():
|
||||||
|
models = [
|
||||||
|
{
|
||||||
|
"id": 1,
|
||||||
|
"name": "Model A",
|
||||||
|
"type": "LORA",
|
||||||
|
"tags": [],
|
||||||
|
"modelVersions": [
|
||||||
|
{"id": 100, "name": "v1", "images": [{"url": "http://example.com/a.jpg"}]},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 2,
|
||||||
|
"name": "Unsupported",
|
||||||
|
"type": "Other",
|
||||||
|
"modelVersions": [{"id": 200, "name": "v1"}],
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
|
||||||
|
|
||||||
|
async def provider_factory():
|
||||||
|
return provider
|
||||||
|
|
||||||
|
handler = ModelLibraryHandler(
|
||||||
|
ServiceRegistryAdapter(
|
||||||
|
get_lora_scanner=fake_scanner_factory,
|
||||||
|
get_checkpoint_scanner=fake_scanner_factory,
|
||||||
|
get_embedding_scanner=fake_scanner_factory,
|
||||||
|
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||||
|
),
|
||||||
|
metadata_provider_factory=provider_factory,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await handler.get_civitai_user_models(
|
||||||
|
FakeRequest(query={"username": "pixel"})
|
||||||
|
)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
assert payload["success"] is True
|
||||||
|
# modelCount only counts models surviving the type filter
|
||||||
|
assert payload["modelCount"] == 1
|
||||||
|
assert payload["nextCursor"] == "cursor-token"
|
||||||
|
assert payload["hasMore"] is True
|
||||||
|
# first page includes the estimated total
|
||||||
|
assert payload["estimatedTotal"] == 2140
|
||||||
|
assert provider.received_cursors == [None]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
|
||||||
|
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
|
||||||
|
|
||||||
|
async def provider_factory():
|
||||||
|
return provider
|
||||||
|
|
||||||
|
handler = ModelLibraryHandler(
|
||||||
|
ServiceRegistryAdapter(
|
||||||
|
get_lora_scanner=fake_scanner_factory,
|
||||||
|
get_checkpoint_scanner=fake_scanner_factory,
|
||||||
|
get_embedding_scanner=fake_scanner_factory,
|
||||||
|
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||||
|
),
|
||||||
|
metadata_provider_factory=provider_factory,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await handler.get_civitai_user_models(
|
||||||
|
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
|
||||||
|
)
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
|
||||||
|
assert response.status == 200
|
||||||
|
assert payload["success"] is True
|
||||||
|
assert payload["nextCursor"] is None
|
||||||
|
assert payload["hasMore"] is False
|
||||||
|
# cursor requests must not include the estimated total
|
||||||
|
assert payload["estimatedTotal"] is None
|
||||||
|
assert provider.received_cursors == ["opaque-token"]
|
||||||
|
|
||||||
|
|
||||||
def test_ensure_handler_mapping_caches_result():
|
def test_ensure_handler_mapping_caches_result():
|
||||||
call_records = []
|
call_records = []
|
||||||
|
|
||||||
|
|||||||
@@ -344,7 +344,7 @@ async def test_switch_channel_to_nightly_with_git_calls_git_update(monkeypatch,
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_switch_channel_to_release_with_git_downloads_zip(monkeypatch, tmp_path):
|
async def test_switch_channel_to_release_with_git_calls_git_update(monkeypatch, tmp_path):
|
||||||
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
||||||
routes_file.parent.mkdir(parents=True)
|
routes_file.parent.mkdir(parents=True)
|
||||||
routes_file.write_text("")
|
routes_file.write_text("")
|
||||||
@@ -353,11 +353,11 @@ async def test_switch_channel_to_release_with_git_downloads_zip(monkeypatch, tmp
|
|||||||
|
|
||||||
(tmp_path / ".git").mkdir()
|
(tmp_path / ".git").mkdir()
|
||||||
|
|
||||||
async def _fake_zip(*args, **kwargs):
|
async def _fake_git_update(*args, **kwargs):
|
||||||
return True, "v9.9.9"
|
return True, "v9.9.9"
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
update_routes.UpdateRoutes, "_download_and_replace_zip", _fake_zip
|
update_routes.UpdateRoutes, "_perform_git_update", _fake_git_update
|
||||||
)
|
)
|
||||||
|
|
||||||
req = _fake_request({"channel": "release"})
|
req = _fake_request({"channel": "release"})
|
||||||
|
|||||||
@@ -183,7 +183,7 @@ class FakeCache:
|
|||||||
def __init__(self, items):
|
def __init__(self, items):
|
||||||
self.items = list(items)
|
self.items = list(items)
|
||||||
|
|
||||||
async def get_sorted_data(self, sort_key, order):
|
async def get_sorted_data(self, sort_key, order, seed=None):
|
||||||
if sort_key == "name":
|
if sort_key == "name":
|
||||||
data = sorted(self.items, key=lambda x: x["model_name"].lower())
|
data = sorted(self.items, key=lambda x: x["model_name"].lower())
|
||||||
if order == "desc":
|
if order == "desc":
|
||||||
|
|||||||
@@ -363,6 +363,148 @@ async def test_check_pending_models_handles_corrupted_progress_file(
|
|||||||
assert result["pending_count"] == 1
|
assert result["pending_count"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.usefixtures("tmp_path")
|
||||||
|
async def test_check_pending_models_uses_bulk_folder_index_for_large_libraries(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
"""For >1000 candidates the pre-check scans the library root once instead of
|
||||||
|
probing every folder individually."""
|
||||||
|
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||||
|
|
||||||
|
# 1500 unprocessed models triggers the bulk lookup path
|
||||||
|
models = [
|
||||||
|
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||||
|
for i in range(1500)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Create folders with files for the first 500 models
|
||||||
|
for i in range(500):
|
||||||
|
model_dir = tmp_path / f"{i:064x}"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "image_0.png").write_text("data")
|
||||||
|
|
||||||
|
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||||
|
|
||||||
|
per_model_checks = 0
|
||||||
|
|
||||||
|
def counting_model_directory_has_files(path: str) -> bool:
|
||||||
|
nonlocal per_model_checks
|
||||||
|
per_model_checks += 1
|
||||||
|
return False
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module,
|
||||||
|
"_model_directory_has_files",
|
||||||
|
counting_model_directory_has_files,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await manager.check_pending_models(["lora"])
|
||||||
|
|
||||||
|
assert result["success"] is True
|
||||||
|
assert result["total_models"] == 1500
|
||||||
|
assert result["pending_count"] == 1000
|
||||||
|
assert result["needs_download"] is True
|
||||||
|
# The per-folder check should not be used once we cross the threshold.
|
||||||
|
assert per_model_checks == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.usefixtures("tmp_path")
|
||||||
|
async def test_check_pending_models_uses_per_folder_check_for_small_candidate_sets(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
"""For <=1000 candidates the pre-check keeps the accurate per-folder path."""
|
||||||
|
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||||
|
|
||||||
|
models = [
|
||||||
|
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||||
|
for i in range(500)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Create folders with files for the first 200 models
|
||||||
|
for i in range(200):
|
||||||
|
model_dir = tmp_path / f"{i:064x}"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "image_0.png").write_text("data")
|
||||||
|
|
||||||
|
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||||
|
|
||||||
|
per_model_checks = 0
|
||||||
|
original_has_files = download_module._model_directory_has_files
|
||||||
|
|
||||||
|
def counting_model_directory_has_files(path: str) -> bool:
|
||||||
|
nonlocal per_model_checks
|
||||||
|
per_model_checks += 1
|
||||||
|
return original_has_files(path)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module,
|
||||||
|
"_model_directory_has_files",
|
||||||
|
counting_model_directory_has_files,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await manager.check_pending_models(["lora"])
|
||||||
|
|
||||||
|
assert result["success"] is True
|
||||||
|
assert result["total_models"] == 500
|
||||||
|
assert result["pending_count"] == 300
|
||||||
|
assert result["needs_download"] is True
|
||||||
|
# Per-folder path should run once per candidate.
|
||||||
|
assert per_model_checks == 500
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.usefixtures("tmp_path")
|
||||||
|
async def test_check_pending_models_bulk_index_includes_legacy_folders(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
"""In multi-library mode the bulk index also scans the legacy root so models
|
||||||
|
whose folders have not been consolidated yet are not reported pending."""
|
||||||
|
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "libraries", {"default": {}, "extra": {}})
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "active_library", "extra")
|
||||||
|
|
||||||
|
# 1500 unprocessed models triggers the bulk lookup path
|
||||||
|
models = [
|
||||||
|
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||||
|
for i in range(1500)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Folders live at the LEGACY root/<hash> path (not yet consolidated)
|
||||||
|
for i in range(500):
|
||||||
|
model_dir = tmp_path / f"{i:064x}"
|
||||||
|
model_dir.mkdir()
|
||||||
|
(model_dir / "image_0.png").write_text("data")
|
||||||
|
|
||||||
|
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||||
|
|
||||||
|
result = await manager.check_pending_models(["lora"])
|
||||||
|
|
||||||
|
assert result["success"] is True
|
||||||
|
assert result["total_models"] == 1500
|
||||||
|
assert result["pending_count"] == 1000
|
||||||
|
assert result["needs_download"] is True
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def settings_manager():
|
def settings_manager():
|
||||||
return get_settings_manager()
|
return get_settings_manager()
|
||||||
|
|||||||
@@ -35,9 +35,11 @@ class DummyDownloader:
|
|||||||
def reset_singletons():
|
def reset_singletons():
|
||||||
CivitaiClient._instance = None
|
CivitaiClient._instance = None
|
||||||
ModelMetadataProviderManager._instance = None
|
ModelMetadataProviderManager._instance = None
|
||||||
|
civitai_client_module._creator_model_count_cache.clear()
|
||||||
yield
|
yield
|
||||||
CivitaiClient._instance = None
|
CivitaiClient._instance = None
|
||||||
ModelMetadataProviderManager._instance = None
|
ModelMetadataProviderManager._instance = None
|
||||||
|
civitai_client_module._creator_model_count_cache.clear()
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
|
|||||||
|
|
||||||
assert result is None
|
assert result is None
|
||||||
assert "Invalid image ID format" in caplog.text
|
assert "Invalid image ID format" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_requests_first_page_with_stable_params(downloader):
|
||||||
|
request_calls = []
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
request_calls.append({"method": method, "url": url, "kwargs": kwargs})
|
||||||
|
return True, {
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"id": 1,
|
||||||
|
"modelVersions": [
|
||||||
|
{"id": 100, "images": [{"meta": {"comfy": {"x": 1}}}]}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"metadata": {"nextCursor": "next-token"},
|
||||||
|
}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel")
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result["nextCursor"] == "next-token"
|
||||||
|
assert len(result["items"]) == 1
|
||||||
|
# comfy metadata is still stripped
|
||||||
|
assert "comfy" not in result["items"][0]["modelVersions"][0]["images"][0]["meta"]
|
||||||
|
|
||||||
|
call = request_calls[0]
|
||||||
|
assert call["method"] == "GET"
|
||||||
|
assert call["url"] == "https://civitai.red/api/v1/models"
|
||||||
|
params = call["kwargs"]["params"]
|
||||||
|
assert params["username"] == "pixel"
|
||||||
|
assert params["nsfw"] == "true"
|
||||||
|
assert params["limit"] == 100
|
||||||
|
assert params["sort"] == "Newest"
|
||||||
|
assert params["period"] == "AllTime"
|
||||||
|
assert "cursor" not in params
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_passes_cursor_and_stringifies_next_cursor(downloader):
|
||||||
|
request_calls = []
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
request_calls.append(kwargs)
|
||||||
|
return True, {"items": [], "metadata": {"nextCursor": 12345}}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel", cursor="opaque-token")
|
||||||
|
|
||||||
|
assert request_calls[0]["params"]["cursor"] == "opaque-token"
|
||||||
|
assert result == {"items": [], "nextCursor": "12345"}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_without_next_cursor_returns_none_cursor(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return True, {"items": [{"id": 1, "modelVersions": []}], "metadata": {}}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel")
|
||||||
|
|
||||||
|
assert result == {"items": [{"id": 1, "modelVersions": []}], "nextCursor": None}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_user_models_failure_returns_none(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return False, "500 server error"
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
result = await client.get_user_models("pixel")
|
||||||
|
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_matches_exact_username(downloader):
|
||||||
|
request_calls = []
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
request_calls.append({"url": url, "kwargs": kwargs})
|
||||||
|
return True, {
|
||||||
|
"items": [
|
||||||
|
{"username": "pixelart", "modelCount": 5},
|
||||||
|
{"username": "Pixel", "modelCount": 2140},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
count = await client.get_creator_model_count("pixel")
|
||||||
|
|
||||||
|
assert count == 2140
|
||||||
|
assert request_calls[0]["url"] == "https://civitai.red/api/v1/creators"
|
||||||
|
assert request_calls[0]["kwargs"]["params"] == {"query": "pixel", "limit": 10}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_without_exact_match_returns_none(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return True, {"items": [{"username": "pixelart", "modelCount": 5}]}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
count = await client.get_creator_model_count("pixel")
|
||||||
|
|
||||||
|
assert count is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_caches_results(downloader):
|
||||||
|
request_count = 0
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
nonlocal request_count
|
||||||
|
request_count += 1
|
||||||
|
return True, {"items": [{"username": "pixel", "modelCount": 42}]}
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
|
||||||
|
assert await client.get_creator_model_count("pixel") == 42
|
||||||
|
# case-insensitive cache key, second call served from cache
|
||||||
|
assert await client.get_creator_model_count("Pixel") == 42
|
||||||
|
assert request_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_caches_failures(downloader):
|
||||||
|
request_count = 0
|
||||||
|
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
nonlocal request_count
|
||||||
|
request_count += 1
|
||||||
|
return False, "500 server error"
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
|
||||||
|
assert await client.get_creator_model_count("pixel") is None
|
||||||
|
assert await client.get_creator_model_count("pixel") is None
|
||||||
|
assert request_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_creator_model_count_never_raises(downloader):
|
||||||
|
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||||
|
return True, "unexpected non-dict payload"
|
||||||
|
|
||||||
|
downloader.make_request = fake_make_request
|
||||||
|
|
||||||
|
client = await CivitaiClient.get_instance()
|
||||||
|
assert await client.get_creator_model_count("pixel") is None
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ class StubScanner:
|
|||||||
|
|
||||||
def __init__(self, models: list[dict]) -> None:
|
def __init__(self, models: list[dict]) -> None:
|
||||||
self._cache = SimpleNamespace(raw_data=models)
|
self._cache = SimpleNamespace(raw_data=models)
|
||||||
|
self.sync_calls: list[tuple[str, dict]] = []
|
||||||
|
|
||||||
async def get_cached_data(self):
|
async def get_cached_data(self):
|
||||||
return self._cache
|
return self._cache
|
||||||
@@ -38,6 +39,14 @@ class StubScanner:
|
|||||||
break
|
break
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
|
||||||
|
self.sync_calls.append((file_path, metadata))
|
||||||
|
for index, model in enumerate(self._cache.raw_data):
|
||||||
|
if model.get("file_path") == metadata.get("file_path"):
|
||||||
|
self._cache.raw_data[index] = metadata
|
||||||
|
break
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
|
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
|
||||||
async def _get_lora_scanner(cls):
|
async def _get_lora_scanner(cls):
|
||||||
@@ -588,6 +597,9 @@ async def test_not_found_example_images_are_cleaned(
|
|||||||
assert missing_url in downloader.calls
|
assert missing_url in downloader.calls
|
||||||
assert manager._progress["failed_models"] == {model_hash}
|
assert manager._progress["failed_models"] == {model_hash}
|
||||||
assert model_hash in manager._progress["processed_models"]
|
assert model_hash in manager._progress["processed_models"]
|
||||||
|
assert scanner.sync_calls
|
||||||
|
assert len(scanner.sync_calls) == 1
|
||||||
|
assert scanner.sync_calls[0][0] == str(model_path)
|
||||||
|
|
||||||
remaining_images = model_metadata["civitai"]["images"]
|
remaining_images = model_metadata["civitai"]["images"]
|
||||||
assert remaining_images == [
|
assert remaining_images == [
|
||||||
|
|||||||
@@ -884,7 +884,7 @@ async def test_sync_cache_conditional_resort_skipped(tmp_path: Path, monkeypatch
|
|||||||
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
||||||
)
|
)
|
||||||
await scanner._cache.resort()
|
await scanner._cache.resort()
|
||||||
scanner._cache._last_sort = ("name", "asc") # name sort is active
|
scanner._cache._last_sort = ("name", "asc", None) # name sort is active
|
||||||
scanner._tags_count = {"alpha": 1}
|
scanner._tags_count = {"alpha": 1}
|
||||||
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
||||||
|
|
||||||
@@ -935,7 +935,7 @@ async def test_sync_cache_conditional_resort_triggered(tmp_path: Path, monkeypat
|
|||||||
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
||||||
)
|
)
|
||||||
await scanner._cache.resort()
|
await scanner._cache.resort()
|
||||||
scanner._cache._last_sort = ("name", "asc")
|
scanner._cache._last_sort = ("name", "asc", None)
|
||||||
scanner._tags_count = {"alpha": 1}
|
scanner._tags_count = {"alpha": 1}
|
||||||
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,97 @@
|
|||||||
|
"""Tests for sort parsing and the seeded random sort mode."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from py.services.model_cache import ModelCache
|
||||||
|
from py.services.model_query import ModelCacheRepository, SortParams
|
||||||
|
|
||||||
|
|
||||||
|
def _make_cache(items):
|
||||||
|
return ModelCache(
|
||||||
|
raw_data=[
|
||||||
|
{
|
||||||
|
"file_path": f"/models/{name}.safetensors",
|
||||||
|
"file_name": f"{name}.safetensors",
|
||||||
|
"model_name": name,
|
||||||
|
"folder": "",
|
||||||
|
"size": 100,
|
||||||
|
"modified": 0.0,
|
||||||
|
}
|
||||||
|
for name in items
|
||||||
|
],
|
||||||
|
folders=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestParseSort:
|
||||||
|
def test_random_with_seed(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("random:abc123")
|
||||||
|
assert params == SortParams(key="random", order="asc", seed="abc123")
|
||||||
|
|
||||||
|
def test_random_without_seed(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("random")
|
||||||
|
assert params == SortParams(key="random", order="asc", seed=None)
|
||||||
|
|
||||||
|
def test_random_empty_seed_falls_back_to_none(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("random:")
|
||||||
|
assert params.seed is None
|
||||||
|
|
||||||
|
def test_regular_sorts_unaffected(self):
|
||||||
|
params = ModelCacheRepository.parse_sort("name:desc")
|
||||||
|
assert params == SortParams(key="name", order="desc", seed=None)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRandomShuffle:
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_same_seed_yields_same_order(self):
|
||||||
|
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||||
|
await asyncio.sleep(0) # allow background resort task to run
|
||||||
|
|
||||||
|
first = await cache.get_sorted_data("random", "asc", "seed1")
|
||||||
|
second = await cache.get_sorted_data("random", "asc", "seed1")
|
||||||
|
|
||||||
|
assert [item["model_name"] for item in first] == [
|
||||||
|
item["model_name"] for item in second
|
||||||
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_different_seeds_yield_different_orders(self):
|
||||||
|
cache = _make_cache([f"m{i}" for i in range(20)])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
first = await cache.get_sorted_data("random", "asc", "seed-a")
|
||||||
|
second = await cache.get_sorted_data("random", "asc", "seed-b")
|
||||||
|
|
||||||
|
assert [item["model_name"] for item in first] != [
|
||||||
|
item["model_name"] for item in second
|
||||||
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_shuffle_is_a_permutation(self):
|
||||||
|
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
shuffled = await cache.get_sorted_data("random", "asc", "seed")
|
||||||
|
|
||||||
|
assert sorted(item["model_name"] for item in shuffled) == [
|
||||||
|
"a",
|
||||||
|
"b",
|
||||||
|
"c",
|
||||||
|
"d",
|
||||||
|
"e",
|
||||||
|
]
|
||||||
|
assert len({item["file_path"] for item in shuffled}) == 5
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_seed_is_stable(self):
|
||||||
|
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
first = await cache.get_sorted_data("random", "asc")
|
||||||
|
second = await cache.get_sorted_data("random", "asc")
|
||||||
|
|
||||||
|
assert [item["model_name"] for item in first] == [
|
||||||
|
item["model_name"] for item in second
|
||||||
|
]
|
||||||
@@ -15,6 +15,7 @@ class StubScanner:
|
|||||||
def __init__(self, cache_items: List[Dict[str, Any]]) -> None:
|
def __init__(self, cache_items: List[Dict[str, Any]]) -> None:
|
||||||
self.cache = SimpleNamespace(raw_data=cache_items)
|
self.cache = SimpleNamespace(raw_data=cache_items)
|
||||||
self.updates: List[Tuple[str, str, Dict[str, Any]]] = []
|
self.updates: List[Tuple[str, str, Dict[str, Any]]] = []
|
||||||
|
self.sync_updates: List[Tuple[str, Dict[str, Any]]] = []
|
||||||
|
|
||||||
async def get_cached_data(self):
|
async def get_cached_data(self):
|
||||||
return self.cache
|
return self.cache
|
||||||
@@ -23,6 +24,10 @@ class StubScanner:
|
|||||||
self.updates.append((old_path, new_path, metadata))
|
self.updates.append((old_path, new_path, metadata))
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def sync_cache_from_metadata(self, file_path: str, metadata: Dict[str, Any]) -> bool:
|
||||||
|
self.sync_updates.append((file_path, metadata))
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
@pytest.fixture(autouse=True)
|
||||||
def patch_metadata_manager(monkeypatch: pytest.MonkeyPatch):
|
def patch_metadata_manager(monkeypatch: pytest.MonkeyPatch):
|
||||||
@@ -83,7 +88,7 @@ async def test_update_metadata_after_import_enriches_entries(monkeypatch: pytest
|
|||||||
assert custom[0]["type"] == "image"
|
assert custom[0]["type"] == "image"
|
||||||
|
|
||||||
assert Path(patch_metadata_manager[0][0]) == model_file
|
assert Path(patch_metadata_manager[0][0]) == model_file
|
||||||
assert scanner.updates
|
assert scanner.sync_updates
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -151,8 +156,8 @@ async def test_update_metadata_after_import_preserves_existing_metadata(
|
|||||||
assert saved_payload["civitai"]["trainedWords"] == ["foo"]
|
assert saved_payload["civitai"]["trainedWords"] == ["foo"]
|
||||||
assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
||||||
|
|
||||||
assert scanner.updates
|
assert scanner.sync_updates
|
||||||
updated_metadata = scanner.updates[-1][2]
|
updated_metadata = scanner.sync_updates[-1][1]
|
||||||
assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"]
|
assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"]
|
||||||
assert {entry["id"] for entry in updated_metadata["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
assert {entry["id"] for entry in updated_metadata["civitai"]["customImages"]} == {"existing-id", "new-id"}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,127 @@
|
|||||||
|
import { app } from "../../scripts/app.js";
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Extension for LoraStackCombinerLM node to support dynamic lora_stack inputs.
|
||||||
|
* Defaults to two inputs; connecting the last slot adds a new empty one, and
|
||||||
|
* disconnecting a non-last slot removes it (at least two are always kept).
|
||||||
|
* Based on the dynamic input pattern from Impact Pack's Switch (Any) node.
|
||||||
|
*/
|
||||||
|
const STACK_INPUT_PATTERN = /^lora_stack\d+$/;
|
||||||
|
|
||||||
|
app.registerExtension({
|
||||||
|
name: "Comfy.LoraManager.LoraStackCombiner",
|
||||||
|
|
||||||
|
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||||
|
if (nodeData.name !== "Lora Stack Combiner (LoraManager)") {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||||
|
|
||||||
|
nodeType.prototype.onConnectionsChange = function(type, index, connected, link_info) {
|
||||||
|
// Skip while the graph is being (re)configured (load, paste, subgraph ops)
|
||||||
|
if (app.configuringGraph) {
|
||||||
|
return onConnectionsChange?.apply?.(this, arguments);
|
||||||
|
}
|
||||||
|
|
||||||
|
const stackTrace = new Error().stack;
|
||||||
|
|
||||||
|
// Skip during graph loading/pasting to avoid interference
|
||||||
|
if (stackTrace.includes('loadGraphData') || stackTrace.includes('pasteFromClipboard')) {
|
||||||
|
return onConnectionsChange?.apply?.(this, arguments);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip subgraph operations
|
||||||
|
if (stackTrace.includes('convertToSubgraph') || stackTrace.includes('Subgraph.configure')) {
|
||||||
|
return onConnectionsChange?.apply?.(this, arguments);
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!link_info) {
|
||||||
|
return onConnectionsChange?.apply?.(this, arguments);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle input connections (type === 1)
|
||||||
|
if (type === 1) {
|
||||||
|
const input = this.inputs[index];
|
||||||
|
|
||||||
|
// Only process numbered lora_stack inputs (legacy a/b slots are left untouched)
|
||||||
|
if (!input || !STACK_INPUT_PATTERN.test(input.name)) {
|
||||||
|
return onConnectionsChange?.apply?.(this, arguments);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Count existing numbered lora_stack inputs
|
||||||
|
let stackInputCount = 0;
|
||||||
|
for (const inp of this.inputs) {
|
||||||
|
if (STACK_INPUT_PATTERN.test(inp.name)) {
|
||||||
|
stackInputCount++;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Renumber all numbered lora_stack inputs sequentially
|
||||||
|
let slotIndex = 1;
|
||||||
|
for (const inp of this.inputs) {
|
||||||
|
if (STACK_INPUT_PATTERN.test(inp.name)) {
|
||||||
|
inp.name = `lora_stack${slotIndex}`;
|
||||||
|
slotIndex++;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add new input slot if connected and this was the last one
|
||||||
|
if (connected) {
|
||||||
|
const lastStackIndex = stackInputCount;
|
||||||
|
if (index === lastStackIndex || index === this.inputs.findIndex(i => i.name === `lora_stack${lastStackIndex}`)) {
|
||||||
|
this.addInput(`lora_stack${slotIndex}`, "LORA_STACK", {
|
||||||
|
tooltip: "A LoRA stack to combine. Connect to add more inputs."
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove disconnected input slots (but keep at least two).
|
||||||
|
// LiteGraph fires this event only for slots that had a link, and
|
||||||
|
// it has already cleared input.link by the time the event fires,
|
||||||
|
// so the disconnected slot is always empty at this point.
|
||||||
|
if (!connected && stackInputCount > 2) {
|
||||||
|
const disconnectedInput = this.inputs[index];
|
||||||
|
if (disconnectedInput && STACK_INPUT_PATTERN.test(disconnectedInput.name)) {
|
||||||
|
// Keep the last slot so there is always an empty slot to reconnect into
|
||||||
|
const isLastStackSlot = index === this.inputs.findLastIndex(i => STACK_INPUT_PATTERN.test(i.name));
|
||||||
|
if (!isLastStackSlot) {
|
||||||
|
this.removeInput(index);
|
||||||
|
|
||||||
|
// Renumber again after removal
|
||||||
|
let newSlotIndex = 1;
|
||||||
|
for (const inp of this.inputs) {
|
||||||
|
if (STACK_INPUT_PATTERN.test(inp.name)) {
|
||||||
|
inp.name = `lora_stack${newSlotIndex}`;
|
||||||
|
newSlotIndex++;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return onConnectionsChange?.apply?.(this, arguments);
|
||||||
|
};
|
||||||
|
},
|
||||||
|
|
||||||
|
nodeCreated(node, app) {
|
||||||
|
if (node.comfyClass !== "Lora Stack Combiner (LoraManager)") {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Leave legacy (a/b) workflows untouched
|
||||||
|
const hasLegacyInputs = node.inputs.some(inp => inp.name === "lora_stack_a" || inp.name === "lora_stack_b");
|
||||||
|
if (hasLegacyInputs) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure at least two numbered lora_stack inputs exist on creation
|
||||||
|
const stackInputCount = node.inputs.filter(inp => STACK_INPUT_PATTERN.test(inp.name)).length;
|
||||||
|
for (let i = stackInputCount + 1; i <= 2; i++) {
|
||||||
|
node.addInput(`lora_stack${i}`, "LORA_STACK", {
|
||||||
|
tooltip: "A LoRA stack to combine. Connect to add more inputs."
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
Reference in New Issue
Block a user