Compare commits

..
Author SHA1 Message Date
Will Miao 27027c4497 refactor(recipes): reuse shared local hash cache in create-from-example 2026-08-08 22:45:29 +08:00
Will Miao 86c85c08ec feat(recipes): pass local hash cache to remote and url recipe imports 2026-08-08 22:13:19 +08:00
Will Miao 196c8ffc3e feat(recipes): match civitai image hash sections against local hash cache 2026-08-08 22:12:47 +08:00
Will Miao cfc95ee02a feat(recipes): pass local hash cache through analysis recipe parsing 2026-08-08 22:11:50 +08:00
Will Miao 479fa36997 feat(recipes): add version-cached local hash cache builder 2026-08-08 22:04:09 +08:00
Will Miao 3e1216e9bc feat(recipes): add cache version counter to model scanners 2026-08-08 21:57:36 +08:00
Will Miao 007883b7d1 fix(recipes): backfill lora cache item by autov2/autov3 hash too 2026-08-08 21:51:34 +08:00
Will Miao dc9200a12c fix(recipes): match recipe-format lora cache item by autov2/autov3 hash 2026-08-08 21:50:33 +08:00
Will Miao d2f955266d fix(types): resolve pre-existing basedpyright errors in tests
Fix ~790 basedpyright errors across the test suite:
- Type stub subclasses of real production classes with super().__init__()
- Add missing generic type arguments and Dict[str, Any] annotations
- Add None guards before subscript/member access
- Adapt tests to production API changes (removed dead handlers,
  PersistentModelCache.get_default, _i18n_filter_added location)
2026-08-08 20:12:59 +08:00
Will Miao 8e724538bd fix(types): resolve pre-existing basedpyright errors in py/ and standalone.py
Fix ~950 basedpyright errors across the backend:
- Convert ineffective # type: ignore comments to # pyright: ignore[rule]
- Add missing generic type arguments (Dict[str, Any], list[Any], ...)
- Annotate dynamic dict literals and runtime-initialized attributes
- Widen CivitAI provider tuple signatures in recipe parsers
- Remove dead LoraRoutes handlers calling nonexistent LoraService methods
- Suppress unavoidable ServiceRegistry import cycles (basedpyright counts
  function-local imports as cycle edges)
2026-08-08 20:12:52 +08:00
Will Miao 6fcdeb799d feat(metadata): resolve AutoV3 at download time without waiting for backfill
- Read AutoV3 directly from the downloaded file's own file_info hashes
  (no SHA256 cross-matching against version_info.files, so the value is
  captured even when the API omits SHA256)
- Extract normalize_autov3() validation helper shared with the
  sha256-matching autov3_from_civitai_files path
- Fall back to the embedded safetensors header hash at download
  completion; mark '' (checked-unavailable) so the startup backfill
  query (autov3 IS NULL) never revisits the row
- Clear archive-level AutoV3 for zip-extracted models so per-file
  header resolution applies to every extracted model
2026-08-08 15:20:43 +08:00
Will Miao 97b9b1f62b feat(metadata): add CivitAI AutoV3 hash support across all storage layers
- Three-state autov3 field (not-checked / checked-unavailable / 12-hex value)
  in .metadata.json sidecars, in-memory ModelHashIndex, and SQLite
  (models.autov3 column + autov3_index table) with column-presence migration
- Background self-terminating backfill for legacy rows: per-model-type
  concurrency guard, executor-offloaded I/O, Civitai-first resolution
  (SHA256-matched version file) falling back to the embedded safetensors
  header hash
- Civitai-first propagation on metadata refresh, scan, and download paths;
  reject the empty-string SHA256 placeholder and strip OneTrainer 0x prefix
- List API hash filters and hash index lookups accept 12-char AutoV3
- Cap safetensors header reads at 64 MiB to prevent crafted-file allocation
- Prevent stale AutoV3 mappings on file replacement while preserving them on
  same-file re-registration (lazy-hash completion)
2026-08-08 14:30:34 +08:00
Will Miao 4bf9a4b640 refactor(download): rename locationStep id to downloadLocationStep
The download modal's step shared the 'locationStep' id with the import
modal, so getElementById('locationStep') could resolve to the wrong
element depending on template include order. The import flow relied on
an injected display:block !important rule to work around it.

Rename the download modal's step id and update all references so each
modal owns a unique step id.
2026-08-08 08:50:36 +08:00
Will Miao c5088772e8 fix(ui): pin import modal action buttons with sticky footer
Make the import modal a flex column with a scrollable step area so the
Back/Import buttons stay visible on short viewports (1080p / 150% zoom)
instead of being cut off at the bottom of the scroll flow.

Also reset step scroll positions via class since 'locationStep' has a
duplicate id in the download modal template.
2026-08-08 08:49:20 +08:00
Will Miao 56acefbd6c feat(autocomplete): search loras within active filters of LoRA Manager page
Add /af and /noaf toggle commands (plus /activefilters aliases) to the
loras autocomplete widget. When enabled (default off), suggestions are
matched within the active filters (folder, base model, tags, auto-tags,
license, tag logic) persisted by the LoRA Manager page in localStorage,
keeping the match pool consistent with the list endpoint, including the
global show_only_sfw setting.

Backend: /lm/{prefix}/relative-paths accepts the filter query params and
pre-filters the scanner cache with ModelFilterSet. The presence of the
recursive param signals the filter pipeline to run even without concrete
filters so global settings stay in parity with the list endpoint.
2026-08-07 20:07:34 +08:00
Will Miao 5ab06c4aae docs: rename "Standalone Web UI" to "LoRA Manager Web UI" in AGENTS.md 2026-08-07 17:57:01 +08:00
Will Miao c11f4b5c68 feat(ui): widen filter panel and preset name limit 2026-08-07 16:53:14 +08:00
Will Miao 86376284f4 fix(ui): clamp filter panel height to viewport 2026-08-07 16:53:04 +08:00
Will Miao 2b8a2fc7d8 feat(filters): remove preset count limit 2026-08-07 16:52:53 +08:00
Will Miao f26e1b41c8 fix(i18n): translate zh-TW api key placeholder 2026-08-07 16:25:40 +08:00
Will Miao c1671af99f feat(downloads): translate batch download summary strings 2026-08-07 16:23:55 +08:00
Will Miao ac7707d0f6 fix(cards): clear model-card min-width on the item element itself 2026-08-07 16:09:51 +08:00
Will Miao 381cd710a2 feat(recipes): translate recipes layout setting strings 2026-08-07 15:36:30 +08:00
Will Miao ad0d18cb79 chore: ignore .playwright-mcp working directory 2026-08-07 15:33:43 +08:00
Will Miao 7980ee77d0 perf(recipes): batch preview dimension reads via asyncio.gather 2026-08-07 15:31:13 +08:00
Will Miao 916b8bb327 fix(recipes): skip stale scroller re-enable on deferred layout switch 2026-08-07 14:03:13 +08:00
Will Miao 87e3d4dea9 feat(recipes): wire recipes layout switch event and rebuild 2026-08-07 12:49:30 +08:00
Will Miao 76a913f5e0 feat(recipes): complete MasonryScroller public API parity with VirtualScroller 2026-08-07 12:40:28 +08:00
Will Miao d8c192e647 feat(recipes): branch masonry scroller instantiation for recipes page 2026-08-07 12:38:17 +08:00
Will Miao c453437620 feat(recipes): add MasonryScroller with column-based virtual scrolling 2026-08-07 12:31:13 +08:00
Will Miao 720fa6d909 feat(recipes): expose preview width/height in recipe listing API 2026-08-07 11:59:26 +08:00
Will Miao b4f71089f4 feat(recipes): add recipes_layout setting (grid|masonry) with i18n 2026-08-07 11:49:12 +08:00
Will Miao 83e6657ead feat(recipes): add get_image_dimensions helper with LRU cache 2026-08-07 11:47:28 +08:00
Will Miao 7ea6df4111 feat(downloads): default to latest version when URL lacks modelVersionId
Auto-select the first (newest) version for URLs without an explicit
modelVersionId, matching the existing batch flow, so users can proceed
to location/download without manually picking a version.
2026-08-07 11:24:45 +08:00
Will Miao d9ab92602a feat(downloads): show failure summary modal for single downloads too 2026-08-07 10:49:14 +08:00
pixelpaws 5ffadaed31 Merge pull request #1054 from willmiao/feat/gemini-provider
feat(llm): add Gemini as a preset AI provider
2026-08-07 10:30:45 +08:00
248 changed files with 10738 additions and 1980 deletions
+1
View File
@@ -25,6 +25,7 @@ model_cache/
reasonix.toml
.reasonix/
.codegraph/
.playwright-mcp/
# Vue widgets development cache (but keep build output)
vue-widgets/node_modules/
+3 -3
View File
@@ -31,7 +31,7 @@ COVERAGE_FILE=coverage/backend/.coverage pytest \
--cov-report=xml:coverage/backend/coverage.xml
```
### Frontend Development (Standalone Web UI)
### Frontend Development (LoRA Manager Web UI)
```bash
npm install
@@ -154,9 +154,9 @@ npm run test:coverage # Generate coverage report
## Frontend UI Architecture
### 1. Standalone Web UI
### 1. LoRA Manager Web UI
- Location: `./static/` and `./templates/`
- Tech: Vanilla JS + CSS, served by standalone server
- Tech: Vanilla JS + CSS, served by the hosting server (ComfyUI app in plugin mode, `standalone.py` in standalone mode)
- Tests via npm in root directory
### 2. ComfyUI Custom Node Widgets
+4
View File
@@ -39,6 +39,7 @@ These fields are present in all model metadata files.
| `metadata_source` | string\|null | ❌ No | ✅ Yes | Last provider that supplied metadata (see below) |
| `last_checked_at` | float | ❌ No (default: `0`) | ✅ Yes | Unix timestamp of last metadata check |
| `hash_status` | string | ❌ No (default: `"completed"`) | ✅ Yes | Hash calculation status: `"pending"`, `"calculating"`, `"completed"`, `"failed"` |
| `autov3` | string\|null | ❌ No | ✅ Yes | CivitAI AutoV3 hash (first 12 chars, lowercase hex) sourced from the safetensors embedded metadata (`sshs_model_hash` / `modelspec.hash_sha256`). **Absent** = not yet checked (may be backfilled later); **`null`** = checked but unavailable (header has no recognized hash); **12-char hex string** = value |
---
@@ -287,6 +288,7 @@ These fields are automatically synchronized with the filesystem:
- `preview_url` — Updated if preview file is moved/removed
- `sha256` — Updated during hash calculation (when `hash_status="pending"`)
- `hash_status` — Updated during hash calculation
- `autov3` — Set when metadata is first created (from safetensors header); may be backfilled later for entries where it is absent
- `last_checked_at` — Timestamp of scan
- `metadata_source` — Set based on metadata provider
@@ -345,6 +347,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
| `metadata_source` | `null` |
| `last_checked_at` | `0` |
| `hash_status` | `"completed"` |
| `autov3` | absent (not checked) or `null` (checked, no value) |
| `usage_tips` | `"{}"` (LoRA only) |
| `model_type` | `"checkpoint"` or `"embedding"` (not present in LoRA models) |
@@ -354,6 +357,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
| Version | Date | Changes |
|---------|------|---------|
| 1.1 | 2026-08 | Added `autov3` field (CivitAI AutoV3 hash with three-state semantics) |
| 1.0 | 2026-03 | Initial schema documentation |
---
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "7 (1080p), 8 (2K), 10 (4K)"
},
"displayDensityWarning": "Warnung: Höhere Dichten können bei Systemen mit begrenzten Ressourcen zu Performance-Problemen führen.",
"recipesLayout": "Rezepte-Layout",
"recipesLayoutHelp": "Wählen Sie, wie Rezeptkarten angeordnet werden: ein einheitliches Raster oder ein Masonry-Layout (Pinterest-Stil), das das Seitenverhältnis jedes Bildes beibehält.",
"recipesLayoutOptions": {
"grid": "Raster",
"masonry": "Masonry"
},
"showFolderSidebar": "Ordner-Seitenleiste anzeigen",
"showFolderSidebarHelp": "Blenden Sie die Ordner-Navigationsleiste auf den Modellseiten ein oder aus. Wenn deaktiviert, bleiben Seitenleiste und Hoverbereich verborgen.",
"cardInfoDisplay": "Karten-Info-Anzeige",
@@ -1584,19 +1590,19 @@
"columnError": "Fehler"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "Zusammenfassung des Batch-Downloads",
"statSuccess": "Erfolgreich",
"statFailed": "Fehlgeschlagen",
"statTotal": "Gesamt",
"successMessage": "Alle {count} Modelle erfolgreich heruntergeladen",
"completedWithErrors": "Abgeschlossen, aber mit Fehlern",
"failed": "Download fehlgeschlagen",
"failedItems": "Fehlgeschlagene Elemente ({count})",
"columnName": "Modellname",
"columnError": "Fehler",
"close": "Schließen",
"copyReport": "Bericht kopieren",
"retryFailed": "Fehlgeschlagene erneut versuchen ({count})"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "Voreinstellungsname darf maximal {max} Zeichen haben",
"presetNameInvalidChars": "Voreinstellungsname enthält ungültige Zeichen",
"presetNameExists": "Eine Voreinstellung mit diesem Namen existiert bereits",
"maxPresetsReached": "Maximal {max} Voreinstellungen erlaubt. Löschen Sie eine, um weitere hinzuzufügen.",
"presetNotFound": "Voreinstellung nicht gefunden",
"invalidPreset": "Ungültige Voreinstellungsdaten",
"deletePresetFailed": "Fehler beim Löschen der Voreinstellung",
+6 -1
View File
@@ -449,6 +449,12 @@
"compact": "7 (1080p), 8 (2K), 10 (4K)"
},
"displayDensityWarning": "Warning: Higher densities may cause performance issues on systems with limited resources.",
"recipesLayout": "Recipes Layout",
"recipesLayoutHelp": "Choose how recipe cards are arranged: a uniform grid or a masonry (Pinterest-style) layout that preserves each image's aspect ratio.",
"recipesLayoutOptions": {
"grid": "Grid",
"masonry": "Masonry"
},
"showFolderSidebar": "Show Folder Sidebar",
"showFolderSidebarHelp": "Toggle the folder navigation sidebar on model pages. When disabled, the sidebar and hover area stay hidden.",
"cardInfoDisplay": "Card Info Display",
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "Preset name must be {max} characters or less",
"presetNameInvalidChars": "Preset name contains invalid characters",
"presetNameExists": "A preset with this name already exists",
"maxPresetsReached": "Maximum {max} presets allowed. Delete one to add more.",
"presetNotFound": "Preset not found",
"invalidPreset": "Invalid preset data",
"deletePresetFailed": "Failed to delete preset",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "7 (1080p), 8 (2K), 10 (4K)"
},
"displayDensityWarning": "Advertencia: Densidades más altas pueden causar problemas de rendimiento en sistemas con recursos limitados.",
"recipesLayout": "Diseño de recetas",
"recipesLayoutHelp": "Elige cómo se organizan las tarjetas de recetas: una cuadrícula uniforme o un diseño masonry (estilo Pinterest) que conserva la proporción de aspecto de cada imagen.",
"recipesLayoutOptions": {
"grid": "Cuadrícula",
"masonry": "Masonry"
},
"showFolderSidebar": "Mostrar barra lateral de carpetas",
"showFolderSidebarHelp": "Activa o desactiva la barra lateral de navegación de carpetas en las páginas de modelos. Cuando está desactivada, la barra lateral y el área de desplazamiento permanecen ocultas.",
"cardInfoDisplay": "Visualización de información de tarjeta",
@@ -1584,19 +1590,19 @@
"columnError": "Error"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "Resumen de descarga por lotes",
"statSuccess": "Correctos",
"statFailed": "Fallidos",
"statTotal": "Total",
"successMessage": "Todos los {count} modelos se descargaron correctamente",
"completedWithErrors": "Completado con errores",
"failed": "Descarga fallida",
"failedItems": "Elementos fallidos ({count})",
"columnName": "Nombre del modelo",
"columnError": "Error",
"close": "Cerrar",
"copyReport": "Copiar informe",
"retryFailed": "Reintentar fallidos ({count})"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "El nombre del preajuste debe tener {max} caracteres o menos",
"presetNameInvalidChars": "El nombre del preajuste contiene caracteres inválidos",
"presetNameExists": "Ya existe un preajuste con este nombre",
"maxPresetsReached": "Máximo {max} preajustes permitidos. Elimine uno para agregar más.",
"presetNotFound": "Preajuste no encontrado",
"invalidPreset": "Datos de preajuste inválidos",
"deletePresetFailed": "Error al eliminar el preajuste",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "7 (1080p), 8 (2K), 10 (4K)"
},
"displayDensityWarning": "Attention : Des densités plus élevées peuvent causer des problèmes de performance sur les systèmes avec des ressources limitées.",
"recipesLayout": "Disposition des recettes",
"recipesLayoutHelp": "Choisissez comment les cartes de recettes sont organisées : une grille uniforme ou une disposition masonry (style Pinterest) qui préserve le rapport d'aspect de chaque image.",
"recipesLayoutOptions": {
"grid": "Grille",
"masonry": "Masonry"
},
"showFolderSidebar": "Afficher la barre latérale des dossiers",
"showFolderSidebarHelp": "Activez ou désactivez la barre latérale de navigation des dossiers sur les pages de modèles. Lorsqu'elle est désactivée, la barre latérale et la zone de survol restent masquées.",
"cardInfoDisplay": "Affichage des informations de carte",
@@ -1584,19 +1590,19 @@
"columnError": "Erreur"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "Résumé du téléchargement groupé",
"statSuccess": "Réussis",
"statFailed": "Échoués",
"statTotal": "Total",
"successMessage": "Les {count} modèles ont été téléchargés avec succès",
"completedWithErrors": "Terminé avec des erreurs",
"failed": "Échec du téléchargement",
"failedItems": "Éléments échoués ({count})",
"columnName": "Nom du modèle",
"columnError": "Erreur",
"close": "Fermer",
"copyReport": "Copier le rapport",
"retryFailed": "Réessayer les échecs ({count})"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "Le nom du préréglage doit contenir au maximum {max} caractères",
"presetNameInvalidChars": "Le nom du préréglage contient des caractères invalides",
"presetNameExists": "Un préréglage avec ce nom existe déjà",
"maxPresetsReached": "Maximum {max} préréglages autorisés. Supprimez-en un pour en ajouter plus.",
"presetNotFound": "Préréglage non trouvé",
"invalidPreset": "Données de préréglage invalides",
"deletePresetFailed": "Échec de la suppression du préréglage",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "7 (1080p), 8 (2K), 10 (4K)"
},
"displayDensityWarning": "אזהרה: צפיפויות גבוהות יותר עלולות לגרום לבעיות ביצועים במערכות עם משאבים מוגבלים.",
"recipesLayout": "פריסת מתכונים",
"recipesLayoutHelp": "בחר כיצד יסודרו כרטיסי המתכונים: רשת אחידה או פריסת Masonry (בסגנון Pinterest) השומרת על יחס הגובה-רוחב של כל תמונה.",
"recipesLayoutOptions": {
"grid": "רשת",
"masonry": "Masonry"
},
"showFolderSidebar": "הצג סרגל צד תיקיות",
"showFolderSidebarHelp": "הפעל או כבה את סרגל הצד לניווט תיקיות בדפי המודל. כאשר הוא כבוי, סרגל הצד ואזור הריחוף נשארים מוסתרים.",
"cardInfoDisplay": "תצוגת מידע בכרטיס",
@@ -1584,19 +1590,19 @@
"columnError": "שגיאה"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "סיכום הורדה בכמות",
"statSuccess": "הצליחו",
"statFailed": "נכשלו",
"statTotal": "סה\"כ",
"successMessage": "כל {count} הדגמים הורדו בהצלחה",
"completedWithErrors": "הושלם עם שגיאות",
"failed": "ההורדה נכשלה",
"failedItems": "פריטים שנכשלו ({count})",
"columnName": "שם הדגם",
"columnError": "שגיאה",
"close": "סגור",
"copyReport": "העתק דוח",
"retryFailed": "נסה שוב ({count})"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "שם קביעה מראש חייב להיות {max} תווים או פחות",
"presetNameInvalidChars": "שם קביעה מראש מכיל תווים לא חוקיים",
"presetNameExists": "קביעה מראש עם שם זה כבר קיימת",
"maxPresetsReached": "מותר מקסימום {max} קביעות מראש. מחק אחת כדי להוסיף עוד.",
"presetNotFound": "קביעה מראש לא נמצאה",
"invalidPreset": "נתוני קביעה מראש לא חוקיים",
"deletePresetFailed": "מחיקת קביעה מראש נכשלה",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "71080p)、82K)、104K"
},
"displayDensityWarning": "警告:高密度設定は、リソースが限られたシステムでパフォーマンスの問題を引き起こす可能性があります。",
"recipesLayout": "レシピのレイアウト",
"recipesLayoutHelp": "レシピカードの配置方法を選択:均一なグリッド、または各画像のアスペクト比を保持するメイソンリー(Pinterest スタイル)レイアウト。",
"recipesLayoutOptions": {
"grid": "グリッド",
"masonry": "メイソンリー"
},
"showFolderSidebar": "フォルダサイドバーを表示",
"showFolderSidebarHelp": "モデルページのフォルダナビゲーションサイドバーを表示/非表示にします。無効にするとサイドバーとホバーエリアは表示されません。",
"cardInfoDisplay": "カード情報表示",
@@ -1584,19 +1590,19 @@
"columnError": "エラー"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "バッチダウンロードの概要",
"statSuccess": "成功",
"statFailed": "失敗",
"statTotal": "合計",
"successMessage": "{count} 個のモデルがすべて正常にダウンロードされました",
"completedWithErrors": "エラーありで完了",
"failed": "ダウンロードに失敗しました",
"failedItems": "失敗した項目({count}",
"columnName": "モデル名",
"columnError": "エラー",
"close": "閉じる",
"copyReport": "レポートをコピー",
"retryFailed": "失敗した項目を再試行({count}"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "プリセット名は{max}文字以内にしてください",
"presetNameInvalidChars": "プリセット名に使用できない文字が含まれています",
"presetNameExists": "同じ名前のプリセットが既に存在します",
"maxPresetsReached": "プリセットは最大{max}個までです。追加するには既存のものを削除してください。",
"presetNotFound": "プリセットが見つかりません",
"invalidPreset": "無効なプリセットデータです",
"deletePresetFailed": "プリセットの削除に失敗しました",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "7개 (1080p), 8개 (2K), 10개 (4K)"
},
"displayDensityWarning": "경고: 높은 밀도는 리소스가 제한된 시스템에서 성능 문제를 일으킬 수 있습니다.",
"recipesLayout": "레시피 레이아웃",
"recipesLayoutHelp": "레시피 카드의 배열 방식을 선택하세요: 균일한 그리드 또는 각 이미지의 종횡비를 유지하는 메이슨리(Pinterest 스타일) 레이아웃.",
"recipesLayoutOptions": {
"grid": "그리드",
"masonry": "메이슨리"
},
"showFolderSidebar": "폴더 사이드바 표시",
"showFolderSidebarHelp": "모델 페이지에서 폴더 탐색 사이드바를 켜거나 끕니다. 비활성화하면 사이드바와 호버 영역이 표시되지 않습니다.",
"cardInfoDisplay": "카드 정보 표시",
@@ -1584,19 +1590,19 @@
"columnError": "오류"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "일괄 다운로드 요약",
"statSuccess": "성공",
"statFailed": "실패",
"statTotal": "전체",
"successMessage": "{count}개 모델이 모두 성공적으로 다운로드되었습니다",
"completedWithErrors": "오류와 함께 완료됨",
"failed": "다운로드 실패",
"failedItems": "실패한 항목 ({count})",
"columnName": "모델 이름",
"columnError": "오류",
"close": "닫기",
"copyReport": "보고서 복사",
"retryFailed": "실패 항목 재시도 ({count})"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "프리셋 이름은 {max}자 이하여야 합니다",
"presetNameInvalidChars": "프리셋 이름에 유효하지 않은 문자가 포함되어 있습니다",
"presetNameExists": "동일한 이름의 프리셋이 이미 존재합니다",
"maxPresetsReached": "최대 {max}개의 프리셋만 허용됩니다. 더 추가하려면 기존 것을 삭제하세요.",
"presetNotFound": "프리셋을 찾을 수 없습니다",
"invalidPreset": "잘못된 프리셋 데이터입니다",
"deletePresetFailed": "프리셋 삭제에 실패했습니다",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "7 (1080p), 8 (2K), 10 (4K)"
},
"displayDensityWarning": "Предупреждение: Высокая плотность может вызвать проблемы с производительностью на системах с ограниченными ресурсами.",
"recipesLayout": "Макет рецептов",
"recipesLayoutHelp": "Выберите, как располагаются карточки рецептов: единая сетка или masonry-макет (в стиле Pinterest), сохраняющий пропорции каждого изображения.",
"recipesLayoutOptions": {
"grid": "Сетка",
"masonry": "Masonry"
},
"showFolderSidebar": "Показывать боковую панель папок",
"showFolderSidebarHelp": "Включает или выключает боковую панель навигации по папкам на страницах моделей. При отключении панель и область наведения скрыты.",
"cardInfoDisplay": "Отображение информации карточки",
@@ -1584,19 +1590,19 @@
"columnError": "Ошибка"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "Сводка пакетной загрузки",
"statSuccess": "Успешно",
"statFailed": "Ошибки",
"statTotal": "Всего",
"successMessage": "Все {count} моделей успешно загружены",
"completedWithErrors": "Завершено с ошибками",
"failed": "Не удалось загрузить",
"failedItems": "Неудачные элементы ({count})",
"columnName": "Имя модели",
"columnError": "Ошибка",
"close": "Закрыть",
"copyReport": "Скопировать отчёт",
"retryFailed": "Повторить неудачные ({count})"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "Имя пресета должно содержать не более {max} символов",
"presetNameInvalidChars": "Имя пресета содержит недопустимые символы",
"presetNameExists": "Пресет с таким именем уже существует",
"maxPresetsReached": "Допустимо максимум {max} пресетов. Удалите один, чтобы добавить больше.",
"presetNotFound": "Пресет не найден",
"invalidPreset": "Недопустимые данные пресета",
"deletePresetFailed": "Не удалось удалить пресет",
+19 -14
View File
@@ -449,6 +449,12 @@
"compact": "71080p),82K),104K"
},
"displayDensityWarning": "警告:高密度可能导致资源有限的系统性能下降。",
"recipesLayout": "配方布局",
"recipesLayoutHelp": "选择配方卡片的排列方式:统一网格,或保留每张图片原始宽高比的瀑布流(Pinterest 风格)布局。",
"recipesLayoutOptions": {
"grid": "网格",
"masonry": "瀑布流"
},
"showFolderSidebar": "显示文件夹侧边栏",
"showFolderSidebarHelp": "在模型页面启用或禁用文件夹导航侧边栏。关闭后,侧边栏和悬停区域将保持隐藏。",
"cardInfoDisplay": "卡片信息显示",
@@ -1584,19 +1590,19 @@
"columnError": "错误"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "批量下载摘要",
"statSuccess": "成功",
"statFailed": "失败",
"statTotal": "总数",
"successMessage": "全部 {count} 个模型下载成功",
"completedWithErrors": "已完成,但有错误",
"failed": "下载失败",
"failedItems": "失败项({count}",
"columnName": "模型名称",
"columnError": "错误",
"close": "关闭",
"copyReport": "复制报告",
"retryFailed": "重试失败项({count}"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "预设名称不能超过 {max} 个字符",
"presetNameInvalidChars": "预设名称包含无效字符",
"presetNameExists": "已存在同名预设",
"maxPresetsReached": "最多允许 {max} 个预设。删除一个以添加更多。",
"presetNotFound": "预设未找到",
"invalidPreset": "无效的预设数据",
"deletePresetFailed": "删除预设失败",
+20 -15
View File
@@ -449,6 +449,12 @@
"compact": "71080p)、82K)、104K"
},
"displayDensityWarning": "警告:較高密度可能導致資源有限的系統效能下降。",
"recipesLayout": "配方版面",
"recipesLayoutHelp": "選擇配方卡片的排列方式:統一網格,或保留每張圖片原始寬高比的瀑布流(Pinterest 風格)版面。",
"recipesLayoutOptions": {
"grid": "網格",
"masonry": "瀑布流"
},
"showFolderSidebar": "顯示資料夾側邊欄",
"showFolderSidebarHelp": "在模型頁面啟用或停用資料夾導覽側邊欄。停用後,側邊欄與滑鼠懸停區域將保持隱藏。",
"cardInfoDisplay": "卡片資訊顯示",
@@ -687,7 +693,7 @@
"apiBasePlaceholder": "https://api.openai.com/v1",
"apiKey": "API 金鑰",
"apiKeyHelp": "LLM 提供者的 API 金鑰。儲存在本地,除您選擇的 LLM 提供者外不會傳送到任何伺服器。",
"apiKeyPlaceholder": "[TODO: Translate] sk-...",
"apiKeyPlaceholder": "sk-...",
"apiKeyNotSet": "未設定",
"apiKeyConfigured": "已設定",
"apiKeySet": "設定",
@@ -1584,19 +1590,19 @@
"columnError": "錯誤"
},
"downloadBatchSummary": {
"title": "[TODO: Translate] Batch Download Summary",
"statSuccess": "[TODO: Translate] Success",
"statFailed": "[TODO: Translate] Failed",
"statTotal": "[TODO: Translate] Total",
"successMessage": "[TODO: Translate] All {count} models downloaded successfully",
"completedWithErrors": "[TODO: Translate] Completed with errors",
"failed": "[TODO: Translate] Download failed",
"failedItems": "[TODO: Translate] Failed Items ({count})",
"columnName": "[TODO: Translate] Model Name",
"columnError": "[TODO: Translate] Error",
"close": "[TODO: Translate] Close",
"copyReport": "[TODO: Translate] Copy Report",
"retryFailed": "[TODO: Translate] Retry Failed ({count})"
"title": "批次下載摘要",
"statSuccess": "成功",
"statFailed": "失敗",
"statTotal": "總數",
"successMessage": "全部 {count} 個模型下載成功",
"completedWithErrors": "已完成,但有錯誤",
"failed": "下載失敗",
"failedItems": "失敗項目({count}",
"columnName": "模型名稱",
"columnError": "錯誤",
"close": "關閉",
"copyReport": "複製報告",
"retryFailed": "重試失敗項目({count}"
}
},
"modelTags": {
@@ -2053,7 +2059,6 @@
"presetNameTooLong": "預設名稱不能超過 {max} 個字元",
"presetNameInvalidChars": "預設名稱包含無效字元",
"presetNameExists": "已存在同名預設",
"maxPresetsReached": "最多允許 {max} 個預設。刪除一個以新增更多。",
"presetNotFound": "預設未找到",
"invalidPreset": "無效的預設資料",
"deletePresetFailed": "刪除預設失敗",
+15 -10
View File
@@ -1,9 +1,13 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import os
import platform
import posixpath
import threading
from pathlib import Path
import folder_paths # type: ignore
import folder_paths # pyright: ignore[reportMissingImports]
from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple
import logging
import json
@@ -90,7 +94,7 @@ def _resolve_valid_default_root(
def _normalize_folder_paths_for_comparison(
folder_paths: Mapping[str, Iterable[str]],
folder_paths: Mapping[str, Any],
) -> Dict[str, Set[str]]:
"""Normalize folder paths for comparison across libraries."""
@@ -482,7 +486,7 @@ class Config:
import ctypes
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # type: ignore[attr-defined]
attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # pyright: ignore[reportAttributeAccessIssue]
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
except Exception as e:
logger.error(f"Error checking Windows reparse point: {e}")
@@ -491,7 +495,7 @@ class Config:
logger.error(f"Error checking link status for {path}: {e}")
return False
def _entry_is_symlink(self, entry: os.DirEntry) -> bool:
def _entry_is_symlink(self, entry: os.DirEntry[str]) -> bool:
"""Check if a directory entry is a symlink, including Windows junctions."""
if entry.is_symlink():
return True
@@ -500,7 +504,7 @@ class Config:
import ctypes
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # type: ignore[attr-defined]
attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # pyright: ignore[reportAttributeAccessIssue]
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
except Exception:
pass
@@ -1126,8 +1130,8 @@ class Config:
def _apply_library_paths(
self,
folder_paths: Mapping[str, Iterable[str]],
extra_folder_paths: Optional[Mapping[str, Iterable[str]]] = None,
folder_paths: Mapping[str, Any],
extra_folder_paths: Optional[Mapping[str, Any]] = None,
recipes_path: str = "",
) -> None:
self._path_mappings.clear()
@@ -1432,12 +1436,13 @@ class Config:
# ('_lm_config_cache') that is NEVER removed from sys.modules (its key does
# NOT start with 'py.'), so it survives re-imports of py.* modules.
_CONFIG_SENTINEL = "_lm_config_cache"
config: Config
if _CONFIG_SENTINEL in _sys.modules:
# Re-import: reuse the existing singleton from the sentinel.
config: Config = _sys.modules[_CONFIG_SENTINEL].config # type: ignore[valid-type]
config = _sys.modules[_CONFIG_SENTINEL].config
else:
config: Config = Config()
config = Config()
# Register the sentinel so re-imports of py.config find us.
_sentinel_mod = _types.ModuleType(_CONFIG_SENTINEL)
_sentinel_mod.config = config
setattr(_sentinel_mod, "config", config)
_sys.modules[_CONFIG_SENTINEL] = _sentinel_mod
+1 -1
View File
@@ -14,7 +14,7 @@ standalone_mode = (
if not standalone_mode:
setup_logging()
from server import PromptServer # type: ignore
from server import PromptServer # pyright: ignore[reportMissingImports]
from .config import config
from .services.model_service_factory import (
+2 -2
View File
@@ -22,7 +22,7 @@ if not standalone_mode:
logger.info("ComfyUI Metadata Collector initialized")
def get_metadata(prompt_id=None): # type: ignore[no-redef]
def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration]
"""Helper function to get metadata from the registry"""
registry = MetadataRegistry()
return registry.get_metadata(prompt_id)
@@ -31,6 +31,6 @@ else:
def init():
logger.info("ComfyUI Metadata Collector disabled in standalone mode")
def get_metadata(prompt_id=None): # type: ignore[no-redef]
def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration]
"""Dummy implementation for standalone mode"""
return {}
+1 -1
View File
@@ -16,7 +16,7 @@ class MetadataHook:
execution = None
try:
# Try direct import first
import execution # type: ignore
import execution # pyright: ignore[reportMissingImports]
except ImportError:
# Try to locate from system modules
for module_name in sys.modules:
+11 -1
View File
@@ -1,5 +1,6 @@
import time
from nodes import NODE_CLASS_MAPPINGS # type: ignore
from typing import Any
from nodes import NODE_CLASS_MAPPINGS # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
@@ -9,6 +10,15 @@ class MetadataRegistry:
_instance = None
current_prompt_id: Any = None
current_prompt: Any = None
metadata: dict[str, Any] = {}
prompt_metadata: dict[str, Any] = {}
executed_nodes: set[str] = set()
node_cache: dict[str, Any] = {}
max_prompt_history: int = 3
metadata_categories: list[str] = METADATA_CATEGORIES
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
+2 -2
View File
@@ -43,7 +43,7 @@ SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys())
async def _find_model_entry(
model_path: str,
) -> tuple[object, object, str | None] | tuple[None, None, None]:
) -> tuple[Any, object, str | None] | tuple[None, None, None]:
"""Iterate all scanners and return the first (scanner, entry, getter_name)
that owns *model_path*. Returns ``(None, None, None)`` when no scanner
claims it.
@@ -73,7 +73,7 @@ async def _find_model_entry(
async def _find_scanner_for_model(
model_path: str,
) -> tuple[object, object] | tuple[None, None]:
) -> tuple[Any, object] | tuple[None, None]:
"""Find the (scanner, cache_entry) responsible for *model_path*."""
scanner, entry, _ = await _find_model_entry(model_path)
return scanner, entry
+6 -6
View File
@@ -1,7 +1,7 @@
import logging
from typing import List, Tuple
import comfy.sd # type: ignore
import folder_paths # type: ignore
from typing import Any, List, Tuple
import comfy.sd # pyright: ignore[reportMissingImports]
import folder_paths # pyright: ignore[reportMissingImports]
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
logger = logging.getLogger(__name__)
@@ -18,9 +18,9 @@ class CheckpointLoaderLM:
CATEGORY = "Lora Manager/loaders"
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
# Get list of checkpoint names from scanner (includes extra folder paths)
checkpoint_names = s._get_checkpoint_names()
checkpoint_names = cls._get_checkpoint_names()
return {
"required": {
"ckpt_name": (
@@ -89,7 +89,7 @@ class CheckpointLoaderLM:
logger.error(f"Error getting checkpoint names: {e}")
return []
def load_checkpoint(self, ckpt_name: str) -> Tuple:
def load_checkpoint(self, ckpt_name: str) -> Tuple[Any, Any, Any]:
"""Load a checkpoint by name, supporting extra folder paths
Args:
+2 -2
View File
@@ -57,8 +57,8 @@ class CreateHookLoraLM:
del text # used by the frontend widget only
# Lazy imports: comfy is not available in CI/test environment at module level
import comfy.hooks # type: ignore # noqa: C0415
import comfy.utils # type: ignore # noqa: C0415
import comfy.hooks # pyright: ignore[reportMissingImports] # noqa: C0415
import comfy.utils # pyright: ignore[reportMissingImports] # noqa: C0415
prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks")
+2 -2
View File
@@ -1,8 +1,8 @@
import importlib
import logging
import comfy.sd # type: ignore
import comfy.utils # type: ignore
import comfy.sd # pyright: ignore[reportMissingImports]
import comfy.utils # pyright: ignore[reportMissingImports]
from ..utils.utils import get_lora_info_absolute
from .utils import (
+1 -1
View File
@@ -73,7 +73,7 @@ class LoraStackCombinerLM:
stack = inspect.stack()
if len(stack) > 2 and stack[2].function == "get_input_info":
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment]
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType]
return {
"required": {},
+12 -13
View File
@@ -15,15 +15,15 @@ import os
import re
from collections import defaultdict
from pathlib import Path
from typing import Dict, List, Optional, Tuple, Union
from typing import Any, Dict, List, Optional, Tuple, Union, cast
import comfy.utils # type: ignore
import folder_paths # type: ignore
import comfy.utils # pyright: ignore[reportMissingImports]
import folder_paths # pyright: ignore[reportMissingImports]
import torch
import torch.nn as nn
from safetensors import safe_open
from nunchaku.lora.flux.nunchaku_converter import (
from nunchaku.lora.flux.nunchaku_converter import ( # pyright: ignore[reportMissingTypeStubs]
pack_lowrank_weight,
unpack_lowrank_weight,
)
@@ -87,10 +87,6 @@ def _rename_layer_underscore_layer_name(old_name: str) -> str:
return new_name
def _is_indexable_module(module):
return isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple))
def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
if not name:
return model
@@ -100,7 +96,7 @@ def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
continue
if hasattr(module, part):
module = getattr(module, part)
elif part.isdigit() and _is_indexable_module(module):
elif part.isdigit() and isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple)):
try:
module = module[int(part)]
except (IndexError, TypeError):
@@ -267,7 +263,9 @@ def _handle_proj_out_split(lora_dict: Dict[str, Dict[str, torch.Tensor]], base_k
return result, consumed
def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: nn.Module) -> None:
def _apply_lora_to_module(module: Any, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: Any) -> None:
# These modules are dynamic torch containers; monkey-patched attributes
# below are set at runtime, so the module/model types are deliberately Any.
if not hasattr(module, "in_features") or not hasattr(module, "out_features"):
raise ValueError(f"{module_name}: unsupported module without in/out features")
if a_tensor.shape[1] != module.in_features or b_tensor.shape[0] != module.out_features:
@@ -336,7 +334,7 @@ def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: t
raise ValueError(f"{module_name}: unsupported module type {type(module)}")
def reset_lora_v2(model: nn.Module) -> None:
def reset_lora_v2(model: Any) -> None:
slots = getattr(model, "_lora_slots", None)
if not slots:
return
@@ -344,6 +342,7 @@ def reset_lora_v2(model: nn.Module) -> None:
module = _get_module_by_name(model, name)
if module is None:
continue
module = cast(Any, module)
module_type = info.get("type", "nunchaku")
if module_type == "nunchaku":
base_rank = info["base_rank"]
@@ -371,7 +370,7 @@ def reset_lora_v2(model: nn.Module) -> None:
def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]], apply_awq_mod: bool = True) -> bool:
del apply_awq_mod # retained for interface compatibility
reset_lora_v2(model)
aggregated_weights: Dict[str, List[Dict[str, object]]] = defaultdict(list)
aggregated_weights: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
saw_supported_format = False
unresolved_targets = 0
@@ -471,7 +470,7 @@ def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path,
class ComfyQwenImageWrapperLM(nn.Module):
def __init__(self, model: nn.Module, config=None, apply_awq_mod: bool = True):
super().__init__()
self.model = model
self.model: Any = model
self.config = {} if config is None else config
self.dtype = next(model.parameters()).dtype
self.loras: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]] = []
+2 -2
View File
@@ -67,7 +67,7 @@ class PromptLM:
stack = inspect.stack()
if len(stack) > 2 and stack[2].function == "get_input_info":
optional_inputs = _PromptOptionalInputs(optional_inputs) # type: ignore[assignment]
optional_inputs = _PromptOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType]
return {
"required": {
@@ -126,7 +126,7 @@ class PromptLM:
else:
prompt = expanded_text
from nodes import CLIPTextEncode # type: ignore
from nodes import CLIPTextEncode # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
conditioning = CLIPTextEncode().encode(clip, prompt)[0]
return (conditioning, prompt)
+7 -7
View File
@@ -5,7 +5,7 @@ import time
import uuid
from typing import Any, Dict, Optional
import numpy as np
import folder_paths # type: ignore
import folder_paths # pyright: ignore[reportMissingImports]
from ..services.service_registry import ServiceRegistry
from ..metadata_collector.metadata_processor import MetadataProcessor
from ..metadata_collector import get_metadata
@@ -13,7 +13,7 @@ from ..utils.constants import CARD_PREVIEW_WIDTH
from ..utils.exif_utils import ExifUtils
from ..utils.utils import calculate_recipe_fingerprint, sanitize_folder_name
from PIL import Image, PngImagePlugin
import piexif
import piexif # pyright: ignore[reportMissingTypeStubs]
import logging
# Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name
@@ -355,7 +355,7 @@ class SaveImageLM:
type_lower = model_type.lower() if model_type else "other"
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
def format_metadata(self, metadata_dict: dict, add_loras_to_prompt: bool = False) -> str:
def format_metadata(self, metadata_dict: dict[str, Any], add_loras_to_prompt: bool = False) -> str:
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
if not metadata_dict: return ""
@@ -396,7 +396,7 @@ class SaveImageLM:
ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0]
# Resolve LoRA hash and Civitai data from local cache
loras_data: list[dict] = []
loras_data: list[dict[str, Any]] = []
for lora_name, strength in lora_entries:
lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
"lora_scanner", lora_name
@@ -418,9 +418,9 @@ class SaveImageLM:
hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper()
# Build Civitai resources JSON array
civitai_resources: list[dict] = []
civitai_resources: list[dict[str, Any]] = []
if ckpt_civitai.get("id", 0) > 0:
ckpt_resource: dict = {}
ckpt_resource: dict[str, Any] = {}
ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint")
model_id = ckpt_civitai.get("modelId", 0)
version_id = ckpt_civitai.get("id", 0)
@@ -439,7 +439,7 @@ class SaveImageLM:
lora_civitai = lora["civitai"]
if not lora_civitai or lora_civitai.get("id", 0) <= 0:
continue
lora_resource: dict = {"weight": lora["strength"]}
lora_resource: dict[str, Any] = {"weight": lora["strength"]}
lora_type = (lora_civitai.get("model") or {}).get("type", "LORA")
model_id = lora_civitai.get("modelId", 0)
version_id = lora_civitai.get("id", 0)
+6 -6
View File
@@ -1,7 +1,7 @@
import logging
import os
from typing import List, Tuple
import comfy.sd # type: ignore
from typing import Any, List, Tuple
import comfy.sd # pyright: ignore[reportMissingImports]
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
logger = logging.getLogger(__name__)
@@ -34,9 +34,9 @@ class UNETLoaderLM:
CATEGORY = "Lora Manager/loaders"
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
# Get list of unet names from scanner (includes extra folder paths)
unet_names = s._get_unet_names()
unet_names = cls._get_unet_names()
return {
"required": {
"unet_name": (
@@ -105,7 +105,7 @@ class UNETLoaderLM:
logger.error(f"Error getting unet names: {e}")
return []
def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple:
def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple[Any, ...]:
"""Load a diffusion model by name, supporting extra folder paths
Args:
@@ -148,7 +148,7 @@ class UNETLoaderLM:
def _load_gguf_unet(
self, unet_path: str, unet_name: str, weight_dtype: str
) -> Tuple:
) -> Tuple[Any, ...]:
"""Load a GGUF format diffusion model
Args:
+7 -3
View File
@@ -1,3 +1,6 @@
from typing import Any
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
@@ -6,7 +9,7 @@ class AnyType(str):
# Credit to Regis Gaughan, III (rgthree)
class FlexibleOptionalInputType(dict):
class FlexibleOptionalInputType(dict[str, Any]):
"""A special class to make flexible nodes that pass data to our python handlers.
Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs
@@ -23,6 +26,7 @@ class FlexibleOptionalInputType(dict):
"""
def __init__(self, type):
super().__init__()
self.type = type
def __getitem__(self, key):
@@ -40,7 +44,7 @@ import re
import logging
import copy
import sys
import folder_paths # type: ignore
import folder_paths # pyright: ignore[reportMissingImports]
logger = logging.getLogger(__name__)
@@ -70,7 +74,7 @@ def extract_lora_name(lora_path):
return apply_lora_syntax_format(name_no_ext)
def parse_lora_syntax(text: str) -> list[dict]:
def parse_lora_syntax(text: str) -> list[dict[str, Any]]:
"""Parse <lora:name:strength> syntax from text input into a list of dicts.
Each entry contains: name, model_strength, clip_strength.
+17 -5
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
"""Base classes for recipe parsers."""
import json
@@ -38,7 +42,7 @@ class RecipeMetadataParser(ABC):
pass
@staticmethod
async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any], Optional[str]],
async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any] | None, str | None] | Dict[str, Any],
recipe_scanner=None, base_model_counts=None, hash_value=None) -> Optional[Dict[str, Any]]:
"""
Populate a lora entry with information from Civitai API response
@@ -175,10 +179,18 @@ class RecipeMetadataParser(ABC):
lora_entry['localPath'] = local_path
lora_entry['file_name'] = os.path.splitext(os.path.basename(local_path))[0]
# Get thumbnail from local preview if available
# Get thumbnail from local preview if available.
# Match the cache item by local path first (get_path_by_hash
# cascade: 10-char autov2 / 12-char autov3), then by hash.
lora_cache = await lora_scanner.get_cached_data()
lora_item = next((item for item in lora_cache.raw_data
if item['sha256'].lower() == lora_entry['hash'].lower()), None)
h = (lora_entry.get("hash") or "").lower()
lora_item = next((item for item in lora_cache.raw_data
if (item.get("file_path") or "") == local_path), None)
if lora_item is None:
lora_item = next((item for item in lora_cache.raw_data
if (item.get("sha256") or "").lower() == h
or (item.get("autov3") or "").lower() == h
or (item.get("sha256") or "")[:10].lower() == h), None)
if lora_item and 'preview_url' in lora_item:
lora_entry['thumbnailUrl'] = config.get_preview_static_url(lora_item['preview_url'])
except Exception as e:
@@ -194,7 +206,7 @@ class RecipeMetadataParser(ABC):
return lora_entry
@staticmethod
async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any]) -> Dict[str, Any]:
async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any] | Tuple[Dict[str, Any] | None, str | None] | None) -> Dict[str, Any]:
"""
Populate checkpoint information from Civitai API response
+4
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import logging
import json
import os
+3 -1
View File
@@ -1,6 +1,7 @@
"""Factory for creating recipe metadata parsers."""
import logging
from typing import Any
from .parsers import (
RecipeFormatParser,
ComfyMetadataParser,
@@ -31,7 +32,8 @@ class RecipeParserFactory:
# First, try CivitaiApiMetadataParser for dict input
if isinstance(metadata, dict):
try:
if CivitaiApiMetadataParser().is_metadata_matching(metadata):
user_comment: Any = metadata
if CivitaiApiMetadataParser().is_metadata_matching(user_comment):
return CivitaiApiMetadataParser()
except Exception as e:
logger.debug(f"CivitaiApiMetadataParser check failed: {e}")
+1 -1
View File
@@ -52,7 +52,7 @@ class AutomaticMetadataParser(RecipeMetadataParser):
negative_and_params = ""
# Initialize metadata
metadata = {
metadata: Dict[str, Any] = {
"prompt": prompt,
"loras": []
}
+117 -56
View File
@@ -4,7 +4,7 @@ import json
import logging
from typing import Dict, Any, Union
from ..base import RecipeMetadataParser
from ..constants import GEN_PARAM_KEYS
from ..constants import GEN_PARAM_KEYS, VALID_LORA_TYPES
from ...services.metadata_service import get_default_metadata_provider
from ...config import config
@@ -14,15 +14,16 @@ logger = logging.getLogger(__name__)
class CivitaiApiMetadataParser(RecipeMetadataParser):
"""Parser for Civitai image metadata format"""
def is_metadata_matching(self, metadata) -> bool:
def is_metadata_matching(self, user_comment) -> bool:
"""Check if the metadata matches the Civitai image metadata format
Args:
metadata: The metadata from the image (dict)
user_comment: The metadata from the image (dict)
Returns:
bool: True if this parser can handle the metadata
"""
metadata = user_comment
if not metadata or not isinstance(metadata, dict):
return False
@@ -73,7 +74,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
return False
async def parse_metadata( # type: ignore[override]
async def parse_metadata( # pyright: ignore[reportIncompatibleMethodOverride]
self, user_comment, recipe_scanner=None, civitai_client=None,
local_cache: dict[str, Any] | None = None,
) -> Dict[str, Any]:
@@ -89,8 +90,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
Returns:
Dict containing parsed recipe data
"""
metadata: Dict[str, Any] = user_comment # type: ignore[assignment]
metadata = user_comment
metadata: Dict[str, Any] = user_comment
try:
# Get metadata provider instead of using civitai_client directly
metadata_provider = await get_default_metadata_provider()
@@ -116,7 +116,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
metadata = inner_meta
# Initialize result structure
result = {
result: Dict[str, Any] = {
"base_model": None,
"loras": [],
"model": None,
@@ -125,10 +125,10 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
}
# Track already added LoRAs to prevent duplicates
added_loras = {} # key: model_version_id or hash, value: index in result["loras"]
added_loras: Dict[str, Any] = {} # key: model_version_id or hash, value: index in result["loras"]
# Extract hash information from hashes field for LoRA matching
lora_hashes = {}
lora_hashes: Dict[str, Any] = {}
if "hashes" in metadata and isinstance(metadata["hashes"], dict):
for key, hash_value in metadata["hashes"].items():
key_str = str(key)
@@ -184,7 +184,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
if model_info:
result["base_model"] = model_info.get("baseModel", "")
base_model_counts = {}
base_model_counts: Dict[str, int] = {}
# Process standard resources array
if "resources" in metadata and isinstance(metadata["resources"], list):
@@ -196,7 +196,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# identification because it has an explicit type field and hash,
# unlike modelVersionIds which is a flat list with no type info.
if resource_type == "model":
checkpoint_entry = {
checkpoint_entry: Dict[str, Any] = {
"id": 0,
"modelId": 0,
"name": resource.get("name", "Unknown Model"),
@@ -216,7 +216,8 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# Try to look up base model from the checkpoint hash
cp_hash = checkpoint_entry.get("hash")
if cp_hash and metadata_provider:
local_cached = local_cache.get(cp_hash) if local_cache else None
# local_cache keys are stored lowercase
local_cached = local_cache.get(cp_hash.lower()) if local_cache else None
if local_cached:
self._populate_entry_from_cache(
checkpoint_entry, local_cached
@@ -294,8 +295,15 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# Try to get info from Civitai if hash is available
if lora_hash and metadata_provider:
local_cached = local_cache.get(lora_hash) if local_cache else None
# local_cache keys are stored lowercase
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
if local_cached:
cached_type = self._cache_item_model_type(local_cached)
if cached_type and cached_type not in VALID_LORA_TYPES:
logger.debug(
f"Skipping non-LoRA cache item for hash {lora_hash}"
)
continue
self._populate_entry_from_cache(
lora_entry, local_cached
)
@@ -304,6 +312,12 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
added_loras[str(lora_entry["id"])] = len(
result["loras"]
)
# Mirror base.py:150-151 counts for API-path loras
bm = local_cached.get("base_model") or ""
if bm:
base_model_counts[bm] = base_model_counts.get(
bm, 0
) + 1
else:
try:
civitai_info = (
@@ -649,30 +663,47 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
}
if metadata_provider:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
# local_cache keys are stored lowercase
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
if local_cached:
cached_type = self._cache_item_model_type(local_cached)
if cached_type and cached_type not in VALID_LORA_TYPES:
logger.debug(
f"Skipping non-LoRA cache item for hash {lora_hash}"
)
continue
lora_entry = populated_entry
self._populate_entry_from_cache(lora_entry, local_cached)
# Mirror base.py:150-151 counts for API-path loras
bm = local_cached.get("base_model") or ""
if bm:
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}"
)
else:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
continue
lora_entry = populated_entry
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}"
)
added_loras[lora_hash] = len(result["loras"])
result["loras"].append(lora_entry)
@@ -711,32 +742,51 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# Try to get info from Civitai if hash is available
if lora_entry["hash"] and metadata_provider:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
# local_cache keys are stored lowercase
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
if local_cached:
cached_type = self._cache_item_model_type(local_cached)
if cached_type and cached_type not in VALID_LORA_TYPES:
logger.debug(
f"Skipping non-LoRA cache item for hash {lora_hash}"
)
lora_index += 1
continue # Skip invalid LoRA types
lora_entry = populated_entry
continue # Skip non-LoRA cache items
self._populate_entry_from_cache(lora_entry, local_cached)
# Mirror base.py:150-151 counts for API-path loras
bm = local_cached.get("base_model") or ""
if bm:
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
# If we have a version ID from Civitai, track it for deduplication
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}"
)
else:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
lora_index += 1
continue # Skip invalid LoRA types
lora_entry = populated_entry
# If we have a version ID from Civitai, track it for deduplication
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}"
)
# Track by hash if we have it
if lora_hash:
@@ -795,3 +845,14 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
base_model = cache_item.get("base_model", "")
if base_model:
entry["baseModel"] = base_model
@staticmethod
def _cache_item_model_type(cache_item: dict[str, Any]) -> str:
"""Lowercased civitai.model.type of a cache item, or '' when unknown."""
civ = cache_item.get("civitai")
if not isinstance(civ, dict):
return ""
model_info = civ.get("model")
if not isinstance(model_info, dict):
return ""
return (model_info.get("type") or "").lower()
+1 -1
View File
@@ -30,7 +30,7 @@ class MetaFormatParser(RecipeMetadataParser):
prompt = parts[0].strip()
# Initialize metadata
metadata = {"prompt": prompt, "loras": []}
metadata: Dict[str, Any] = {"prompt": prompt, "loras": []}
# Extract negative prompt and parameters if available
if len(parts) > 1:
+10 -2
View File
@@ -91,7 +91,15 @@ class RecipeFormatParser(RecipeMetadataParser):
exists_locally = lora_scanner.has_hash(lora['hash'])
if exists_locally:
lora_cache = await lora_scanner.get_cached_data()
lora_item = next((item for item in lora_cache.raw_data if item['sha256'].lower() == lora['hash'].lower()), None)
# Cascade match: full sha256, stored autov3, or autov2 (sha256[:10]).
h = (lora.get('hash') or '').lower()
lora_item = next(
(item for item in lora_cache.raw_data
if (item.get("sha256") or "").lower() == h
or (item.get("autov3") or "").lower() == h
or (item.get("sha256") or "")[:10].lower() == h),
None
)
if lora_item:
lora_entry['existsLocally'] = True
lora_entry['inLibrary'] = True
@@ -148,7 +156,7 @@ class RecipeFormatParser(RecipeMetadataParser):
checkpoint_data = recipe_metadata.get('checkpoint') or {}
if isinstance(checkpoint_data, dict) and checkpoint_data:
version_id = checkpoint_data.get('modelVersionId') or checkpoint_data.get('id')
checkpoint_entry = {
checkpoint_entry: Dict[str, Any] = {
'id': version_id or 0,
'modelId': checkpoint_data.get('modelId', 0),
'name': checkpoint_data.get('name', 'Unknown Checkpoint'),
+7 -7
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Callable, Dict, Mapping
from typing import TYPE_CHECKING, Awaitable, Callable, Dict, Mapping
import jinja2
from aiohttp import web
@@ -84,7 +84,7 @@ class BaseModelRoutes(ABC):
self.metadata_progress_callback = WebSocketBroadcastCallback()
self._handler_set: ModelHandlerSet | None = None
self._handler_mapping: Dict[str, Callable[[web.Request], web.StreamResponse]] | None = None
self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None
self._preview_service = PreviewAssetService(
metadata_manager=MetadataManager,
@@ -131,7 +131,7 @@ class BaseModelRoutes(ABC):
self._handler_set = None
self._handler_mapping = None
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
if self._handler_mapping is None:
handler_set = self._create_handler_set()
self._handler_set = handler_set
@@ -220,7 +220,7 @@ class BaseModelRoutes(ABC):
)
@property
def route_handlers(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def route_handlers(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
return self._ensure_handler_mapping()
def setup_routes(self, app: web.Application, prefix: str) -> None:
@@ -237,7 +237,7 @@ class BaseModelRoutes(ABC):
"""Setup model-specific routes."""
raise NotImplementedError
def _parse_specific_params(self, request: web.Request) -> Dict:
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse model-specific parameters - to be overridden by subclasses."""
return {}
@@ -253,7 +253,7 @@ class BaseModelRoutes(ABC):
"""Find the appropriate model file from the files list - can be overridden by subclasses."""
return next((file for file in files if file.get("type") in ("Model", "Diffusion Model") and file.get("primary") is True), None)
def get_handler(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
def get_handler(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
"""Expose handlers for subclasses or tests."""
return self._ensure_handler_mapping()[name]
@@ -285,7 +285,7 @@ class BaseModelRoutes(ABC):
)
return self.model_lifecycle_service
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
async def proxy(request: web.Request) -> web.StreamResponse:
try:
handler = self.get_handler(name)
+13 -9
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import logging
import os
from typing import Callable, Mapping
from typing import Awaitable, Callable, Mapping
import jinja2
from aiohttp import web
@@ -61,7 +61,9 @@ class BaseRecipeRoutes:
self._i18n_registered = False
self._startup_hooks_registered = False
self._handler_set: RecipeHandlerSet | None = None
self._handler_mapping: dict[str, Callable] | None = None
self._handler_mapping: Mapping[
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
] | None = None
async def attach_dependencies(self, app: web.Application | None = None) -> None:
"""Resolve shared services from the registry."""
@@ -84,7 +86,9 @@ class BaseRecipeRoutes:
app.on_startup.append(self.attach_dependencies)
self._startup_hooks_registered = True
def to_route_mapping(self) -> Mapping[str, Callable]:
def to_route_mapping(
self,
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
"""Return a mapping of handler name to coroutine for registrar binding."""
if self._handler_mapping is None:
@@ -124,17 +128,17 @@ class BaseRecipeRoutes:
or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
)
if not standalone_mode:
from ..metadata_collector import get_metadata # type: ignore[import-not-found]
from ..metadata_collector.metadata_processor import ( # type: ignore[import-not-found]
from ..metadata_collector import get_metadata # pyright: ignore[reportMissingImports]
from ..metadata_collector.metadata_processor import ( # pyright: ignore[reportMissingImports]
MetadataProcessor,
)
from ..metadata_collector.metadata_registry import ( # type: ignore[import-not-found]
from ..metadata_collector.metadata_registry import ( # pyright: ignore[reportMissingImports]
MetadataRegistry,
)
else: # pragma: no cover - optional dependency path
get_metadata = None # type: ignore[assignment]
MetadataProcessor = None # type: ignore[assignment]
MetadataRegistry = None # type: ignore[assignment]
get_metadata = None # pyright: ignore[reportAssignmentType]
MetadataProcessor = None # pyright: ignore[reportAssignmentType]
MetadataRegistry = None # pyright: ignore[reportAssignmentType]
analysis_service = RecipeAnalysisService(
exif_utils=ExifUtils,
+9 -9
View File
@@ -1,5 +1,5 @@
import logging
from typing import Dict, List, Set
from typing import Any, Dict, List, Set
from aiohttp import web
from .base_model_routes import BaseModelRoutes
@@ -28,13 +28,13 @@ class CheckpointRoutes(BaseModelRoutes):
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application):
def setup_routes(self, app: web.Application, prefix: str = "checkpoints"):
"""Setup Checkpoint routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'checkpoints' prefix (includes page route)
super().setup_routes(app, 'checkpoints')
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup Checkpoint-specific routes"""
@@ -53,9 +53,9 @@ class CheckpointRoutes(BaseModelRoutes):
"""Get expected model types string for error messages"""
return "Checkpoint"
def _parse_specific_params(self, request: web.Request) -> Dict:
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse Checkpoint-specific parameters"""
params: Dict = {}
params: Dict[str, Any] = {}
if 'checkpoint_hash' in request.query:
params['hash_filters'] = {'single_hash': request.query['checkpoint_hash'].lower()}
@@ -70,7 +70,7 @@ class CheckpointRoutes(BaseModelRoutes):
"""Get detailed information for a specific checkpoint by name"""
try:
name = request.match_info.get('name', '')
checkpoint_info = await self.service.get_model_info_by_name(name)
checkpoint_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
if checkpoint_info:
return web.json_response(checkpoint_info)
@@ -89,7 +89,7 @@ class CheckpointRoutes(BaseModelRoutes):
roots.extend(config.checkpoints_roots or [])
roots.extend(config.extra_checkpoints_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:
@@ -114,7 +114,7 @@ class CheckpointRoutes(BaseModelRoutes):
roots.extend(config.unet_roots or [])
roots.extend(config.extra_unet_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:
+4 -4
View File
@@ -26,13 +26,13 @@ class EmbeddingRoutes(BaseModelRoutes):
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application):
def setup_routes(self, app: web.Application, prefix: str = "embeddings"):
"""Setup Embedding routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'embeddings' prefix (includes page route)
super().setup_routes(app, 'embeddings')
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup Embedding-specific routes"""
@@ -51,7 +51,7 @@ class EmbeddingRoutes(BaseModelRoutes):
"""Get detailed information for a specific embedding by name"""
try:
name = request.match_info.get('name', '')
embedding_info = await self.service.get_model_info_by_name(name)
embedding_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
if embedding_info:
return web.json_response(embedding_info)
+8 -4
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
from typing import Callable, Mapping
from typing import Any, Awaitable, Callable, Mapping
from aiohttp import web
@@ -35,7 +35,7 @@ class ExampleImagesRoutes:
*,
ws_manager,
download_manager: DownloadManager | None = None,
processor=ExampleImagesProcessor,
processor: Any = ExampleImagesProcessor,
file_manager=ExampleImagesFileManager,
cleanup_service: ExampleImagesCleanupService | None = None,
) -> None:
@@ -46,7 +46,9 @@ class ExampleImagesRoutes:
self._file_manager = file_manager
self._cleanup_service = cleanup_service or ExampleImagesCleanupService()
self._handler_set: ExampleImagesHandlerSet | None = None
self._handler_mapping: Mapping[str, Callable[[web.Request], web.StreamResponse]] | None = None
self._handler_mapping: Mapping[
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
] | None = None
@classmethod
def setup_routes(cls, app: web.Application, *, ws_manager) -> None:
@@ -61,7 +63,9 @@ class ExampleImagesRoutes:
registrar = ExampleImagesRouteRegistrar(app)
registrar.register_routes(self.to_route_mapping())
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def to_route_mapping(
self,
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
"""Return the registrar-compatible mapping of handler names to callables."""
if self._handler_mapping is None:
@@ -3,7 +3,7 @@ from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Callable, Mapping
from typing import Awaitable, Callable, Mapping
from aiohttp import web
@@ -170,7 +170,7 @@ class ExampleImagesHandlerSet:
management: ExampleImagesManagementHandler
files: ExampleImagesFileHandler
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
"""Flatten handler methods into the registrar mapping."""
return {
+53 -45
View File
@@ -276,7 +276,7 @@ def _collect_comfyui_session_logs(
) -> dict[str, Any]:
if log_entries is None:
try:
import app.logger as comfy_logger
import app.logger as comfy_logger # pyright: ignore[reportMissingImports]
log_entries = list(comfy_logger.get_logs() or [])
except Exception as exc: # pragma: no cover - environment dependent
@@ -422,10 +422,10 @@ class PromptServerProtocol(Protocol):
"""Subset of PromptServer used by the handlers."""
instance: "PromptServerProtocol"
sockets: dict # maps clientId (sid) → WebSocketResponse
sockets: dict[str, Any] # maps clientId (sid) → WebSocketResponse
def send_sync(
self, event: str, payload: dict | None = None, sid: str | None = None
self, event: str, payload: dict[str, Any] | None = None, sid: str | None = None
) -> None: # pragma: no cover - protocol
...
@@ -443,7 +443,12 @@ class UsageStatsFactory(Protocol):
class MetadataProviderProtocol(Protocol):
async def get_model_versions(
self, model_id: int
) -> dict | None: # pragma: no cover - protocol
) -> dict[str, Any] | None: # pragma: no cover - protocol
...
async def get_user_models(
self, username: str, cursor: str | None = None
) -> Any: # pragma: no cover - protocol
...
@@ -466,16 +471,16 @@ class MetadataArchiveManagerProtocol(Protocol):
class BackupServiceProtocol(Protocol):
async def create_snapshot(
self, *, snapshot_type: str = "manual", persist: bool = False
) -> dict: # pragma: no cover - protocol
) -> dict[str, Any]: # pragma: no cover - protocol
...
async def restore_snapshot(self, archive_path: str) -> dict: # pragma: no cover - protocol
async def restore_snapshot(self, archive_path: str) -> dict[str, Any]: # pragma: no cover - protocol
...
def get_status(self) -> dict: # pragma: no cover - protocol
def get_status(self) -> dict[str, Any]: # pragma: no cover - protocol
...
def get_available_snapshots(self) -> list[dict]: # pragma: no cover - protocol
def get_available_snapshots(self) -> list[dict[str, Any]]: # pragma: no cover - protocol
...
@@ -491,7 +496,7 @@ class NodeRegistry:
def __init__(self) -> None:
self._lock = asyncio.Lock()
# sid → {unique_id → node_info}
self._tab_nodes: Dict[str, Dict[str, dict]] = {}
self._tab_nodes: Dict[str, Dict[str, dict[str, Any]]] = {}
self._ready = asyncio.Event()
self._waiting_clients: set[str] = set()
@@ -504,7 +509,7 @@ class NodeRegistry:
# Helpers to build one node dict (extracted so it's reused for each tab)
# ------------------------------------------------------------------
@staticmethod
def _build_node_dict(node: dict) -> dict:
def _build_node_dict(node: dict[str, Any]) -> dict[str, Any]:
node_id = node["node_id"]
graph_id = str(node["graph_id"])
unique_id = f"{graph_id}:{node_id}"
@@ -513,11 +518,11 @@ class NodeRegistry:
bgcolor = node.get("bgcolor") or DEFAULT_NODE_COLOR
raw_capabilities = node.get("capabilities")
capabilities: dict = {}
capabilities: dict[str, Any] = {}
if isinstance(raw_capabilities, dict):
capabilities = dict(raw_capabilities)
raw_widget_names: list | None = node.get("widget_names")
raw_widget_names: list[Any] | None = node.get("widget_names")
if not isinstance(raw_widget_names, list):
capability_widget_names = capabilities.get("widget_names")
raw_widget_names = (
@@ -565,9 +570,9 @@ class NodeRegistry:
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
async def register_nodes(self, sid: str, nodes: list[dict]) -> None:
async def register_nodes(self, sid: str, nodes: list[dict[str, Any]]) -> None:
"""Register/replace the node list for a single ComfyUI tab (identified by *sid*)."""
tab_nodes: dict[str, dict] = {}
tab_nodes: dict[str, dict[str, Any]] = {}
for node in nodes:
nd = self._build_node_dict(node)
tab_nodes[nd["unique_id"]] = nd
@@ -602,7 +607,7 @@ class NodeRegistry:
except asyncio.TimeoutError:
return False
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict:
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict[str, Any]:
"""Return the union of all known tab nodes, pruning any tab that is no
longer connected."""
async with self._lock:
@@ -619,8 +624,8 @@ class NodeRegistry:
len(stale_sids), stale_sids,
)
merged: dict[str, dict] = {}
tab_info: dict[str, dict] = {}
merged: dict[str, dict[str, Any]] = {}
tab_info: dict[str, dict[str, Any]] = {}
for sid, nodes in self._tab_nodes.items():
tab_info[sid] = {
"node_count": len(nodes),
@@ -653,7 +658,7 @@ class SupportersHandler:
def __init__(self, logger: logging.Logger | None = None) -> None:
self._logger = logger or logging.getLogger(__name__)
def _load_supporters(self) -> dict:
def _load_supporters(self) -> dict[str, Any]:
"""Load supporters data from JSON file."""
try:
current_file = os.path.abspath(__file__)
@@ -1229,10 +1234,8 @@ class DoctorHandler:
settings_snapshot = _sanitize_sensitive_data(
getattr(self._settings, "settings", {}) or {}
)
startup_messages_getter = getattr(self._settings, "get_startup_messages", None)
startup_messages = (
list(startup_messages_getter()) if callable(startup_messages_getter) else []
)
startup_messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
startup_messages = list(startup_messages_getter()) if startup_messages_getter else []
environment = {
"app_version": app_version,
@@ -1439,7 +1442,7 @@ class SettingsHandler:
*,
settings_service=None,
metadata_provider_updater: Callable[
[], Awaitable[None]
[], Awaitable[Any]
] = update_metadata_providers,
downloader_factory: Callable[
[], Awaitable[DownloaderProtocol]
@@ -1484,8 +1487,8 @@ class SettingsHandler:
settings_file = getattr(self._settings, "settings_file", None)
if settings_file:
response_data["settings_file"] = settings_file
messages_getter = getattr(self._settings, "get_startup_messages", None)
messages = list(messages_getter()) if callable(messages_getter) else []
messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
messages = list(messages_getter()) if messages_getter else []
return web.json_response(
{
"success": True,
@@ -2005,11 +2008,11 @@ async def _noop_backup_service() -> None:
@dataclass
class ServiceRegistryAdapter:
get_lora_scanner: Callable[[], Awaitable]
get_checkpoint_scanner: Callable[[], Awaitable]
get_embedding_scanner: Callable[[], Awaitable]
get_downloaded_version_history_service: Callable[[], Awaitable]
get_backup_service: Callable[[], Awaitable] = _noop_backup_service
get_lora_scanner: Callable[[], Awaitable[Any]]
get_checkpoint_scanner: Callable[[], Awaitable[Any]]
get_embedding_scanner: Callable[[], Awaitable[Any]]
get_downloaded_version_history_service: Callable[[], Awaitable[Any]]
get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service
class ModelLibraryHandler:
@@ -2050,8 +2053,8 @@ class ModelLibraryHandler:
return await self._service_registry.get_downloaded_version_history_service()
@staticmethod
def _with_downloaded_flag(versions: list[dict]) -> list[dict]:
enriched: list[dict] = []
def _with_downloaded_flag(versions: list[dict[str, Any]]) -> list[dict[str, Any]]:
enriched: list[dict[str, Any]] = []
for version in versions:
entry = dict(version)
entry.setdefault("hasBeenDownloaded", True)
@@ -2244,7 +2247,7 @@ class ModelLibraryHandler:
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
embedding_scanner = await self._service_registry.get_embedding_scanner()
results: list[dict] = []
results: list[dict[str, Any]] = []
for model_id in model_ids:
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
if lora_versions:
@@ -2353,7 +2356,7 @@ class ModelLibraryHandler:
)
try:
model_version_id = int(data.get("modelVersionId"))
model_version_id = int(data.get("modelVersionId")) # pyright: ignore[reportArgumentType]
except (TypeError, ValueError):
return web.json_response(
{"success": False, "error": "Parameter modelVersionId must be an integer"},
@@ -2465,10 +2468,11 @@ class ModelLibraryHandler:
"checkpoint": checkpoint_scanner,
"embedding": embedding_scanner,
}
scanner = scanner_map.get(found_type)
scanner = scanner_map.get(found_type or "")
if scanner:
persist = getattr(scanner, "_persist_current_cache", None)
if callable(persist):
scanner.bump_cache_version()
persist: Any = getattr(scanner, "_persist_current_cache", None)
if persist:
await persist()
history_service = await self._get_download_history_service()
@@ -2649,13 +2653,13 @@ class ModelLibraryHandler:
}
lora_type_aliases = {model_type.lower() for model_type in VALID_LORA_TYPES}
type_scanner_map: Dict[str, object | None] = {
type_scanner_map: Dict[str, Any] = {
**{alias: lora_scanner for alias in lora_type_aliases},
"checkpoint": checkpoint_scanner,
"textualinversion": embedding_scanner,
}
versions: list[dict] = []
versions: list[dict[str, Any]] = []
history_service = await self._get_download_history_service()
model_ids: list[int] = []
model_count = 0
@@ -2707,6 +2711,8 @@ class ModelLibraryHandler:
tags_value = model.get("tags")
tags = tags_value if isinstance(tags_value, list) else []
model_id = model.get("id")
if model_id is None:
continue
try:
model_id_int = int(model_id)
except (TypeError, ValueError):
@@ -2722,6 +2728,8 @@ class ModelLibraryHandler:
continue
version_id = version.get("id")
if version_id is None:
continue
try:
version_id_int = int(version_id)
except (TypeError, ValueError):
@@ -2783,7 +2791,7 @@ class MetadataArchiveHandler:
] = get_metadata_archive_manager,
settings_service=None,
metadata_provider_updater: Callable[
[], Awaitable[None]
[], Awaitable[Any]
] = update_metadata_providers,
) -> None:
self._metadata_archive_manager_factory = metadata_archive_manager_factory
@@ -2930,7 +2938,7 @@ class BackupHandler:
if request.content_type.startswith("multipart/"):
reader = await request.multipart()
field = await reader.next()
field: Any = await reader.next()
uploaded = False
while field is not None:
if getattr(field, "filename", None):
@@ -3549,7 +3557,7 @@ class NodeRegistryHandler:
except (TypeError, ValueError):
parsed_node_id = node_identifier
payload: dict = {
payload: dict[str, Any] = {
"id": parsed_node_id,
"value": value,
"mode": mode,
@@ -3673,7 +3681,7 @@ class NodeRegistryHandler:
except (TypeError, ValueError):
parsed_node_id = node_identifier
payload: dict = {
payload: dict[str, Any] = {
"id": parsed_node_id,
"value": value,
"mode": mode,
@@ -3740,8 +3748,8 @@ class MiscHandlerSet:
doctor: DoctorHandler,
example_workflows: ExampleWorkflowsHandler,
base_model: BaseModelHandlerSet,
hf_handler: HfHandler | None = None,
agent_handler: AgentHandler | None = None,
hf_handler: Any = None,
agent_handler: Any = None,
) -> None:
self.health = health
self.settings = settings
+86 -19
View File
@@ -71,7 +71,7 @@ class ModelPageView:
self._server_i18n = server_i18n
self._logger = logger
def _load_supporters(self) -> dict:
def _load_supporters(self) -> dict[str, Any]:
"""Load supporters data from JSON file."""
try:
current_file = os.path.abspath(__file__)
@@ -152,7 +152,7 @@ class ModelPageView:
self._template_env.filters["t"] = (
self._server_i18n.create_template_filter()
)
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
self._template_env._i18n_filter_added = True # pyright: ignore[reportAttributeAccessIssue]
from ...services.llm_service import PROVIDER_PRESETS
@@ -199,7 +199,7 @@ class ModelListingHandler:
self,
*,
service,
parse_specific_params: Callable[[web.Request], Dict],
parse_specific_params: Callable[[web.Request], Dict[str, Any]],
logger: logging.Logger,
) -> None:
self._service = service
@@ -287,7 +287,7 @@ class ModelListingHandler:
)
return web.json_response({"error": str(exc)}, status=500)
def _parse_common_params(self, request: web.Request) -> Dict:
def _parse_common_params(self, request: web.Request) -> Dict[str, Any]:
page = int(request.query.get("page", "1"))
page_size = min(int(request.query.get("page_size", "20")), 100)
sort_by = request.query.get("sort_by", "name")
@@ -658,7 +658,7 @@ class ModelManagementHandler:
try:
reader = await request.multipart()
field = await reader.next()
field: Any = await reader.next()
if field is None or field.name != "preview_file":
raise ValueError("Expected 'preview_file' field")
content_type = field.headers.get("Content-Type", "image/png")
@@ -700,7 +700,7 @@ class ModelManagementHandler:
{
"success": True,
"preview_url": config.get_preview_static_url(
result["preview_path"]
str(result["preview_path"])
),
"preview_nsfw_level": result["preview_nsfw_level"],
}
@@ -781,7 +781,7 @@ class ModelManagementHandler:
result = await self._preview_service.replace_preview(
model_path=model_path,
preview_data=preview_data,
preview_data=preview_bytes,
content_type=content_type,
original_filename=original_filename,
nsfw_level=nsfw_level,
@@ -793,7 +793,7 @@ class ModelManagementHandler:
{
"success": True,
"preview_url": config.get_preview_static_url(
result["preview_path"]
str(result["preview_path"])
),
"preview_nsfw_level": result["preview_nsfw_level"],
}
@@ -1488,8 +1488,73 @@ class ModelQueryHandler:
search = request.query.get("search", "").strip()
limit = min(int(request.query.get("limit", "15")), 100)
offset = max(0, int(request.query.get("offset", "0")))
folder = request.query.get("folder")
recursive = request.query.get("recursive", "true").lower() == "true"
base_models = list(request.query.getall("base_model", []))
model_types = list(request.query.getall("model_type", []))
tag_filters: Dict[str, str] = {}
for tag in request.query.getall("tag_include", []):
if tag:
tag_filters[tag] = "include"
for tag in request.query.getall("tag_exclude", []):
if tag:
tag_filters[tag] = "exclude"
auto_tag_filters: Dict[str, str] = {}
for tag in request.query.getall("auto_tag_include", []):
if tag:
auto_tag_filters[tag] = "include"
for tag in request.query.getall("auto_tag_exclude", []):
if tag:
auto_tag_filters[tag] = "exclude"
tag_logic = request.query.get("tag_logic", "any").lower()
if tag_logic not in ("any", "all"):
tag_logic = "any"
credit_required = request.query.get("credit_required")
if credit_required is not None:
credit_required = credit_required.lower() not in ("false", "0", "")
allow_selling_generated_content = request.query.get(
"allow_selling_generated_content"
)
if allow_selling_generated_content is not None:
allow_selling_generated_content = (
allow_selling_generated_content.lower() not in ("false", "0", "")
)
# The presence of the recursive param (always sent by the loras
# widget when filter mode is on) signals that the filter pipeline
# must run even when no concrete filter is set, so global settings
# like show_only_sfw stay consistent with the list endpoint.
apply_filters = (
"recursive" in request.query
or folder is not None
or bool(base_models)
or bool(model_types)
or bool(tag_filters)
or bool(auto_tag_filters)
or credit_required is not None
or allow_selling_generated_content is not None
)
matching_paths = await self._service.search_relative_paths(
search, limit, offset
search,
limit,
offset,
folder=folder,
recursive=recursive,
base_models=base_models,
model_types=model_types,
tags=tag_filters,
auto_tags=auto_tag_filters,
tag_logic=tag_logic,
credit_required=credit_required,
allow_selling_generated_content=allow_selling_generated_content,
apply_filters=apply_filters,
)
return web.json_response(
{"success": True, "relative_paths": matching_paths}
@@ -1995,7 +2060,7 @@ class ModelCivitaiHandler:
settings_service: SettingsManager,
ws_manager: WebSocketManager,
logger: logging.Logger,
metadata_provider_factory: Callable[[], Awaitable],
metadata_provider_factory: Callable[[], Awaitable[Any]],
validate_model_type: Callable[[str], bool],
expected_model_types: Callable[[], str],
find_model_file: Callable[
@@ -2060,7 +2125,7 @@ class ModelCivitaiHandler:
downloaded_version_ids = set(
await history_service.get_downloaded_version_ids(
self._service.model_type,
model_id,
int(model_id),
)
)
except Exception as exc: # pragma: no cover - defensive logging
@@ -2337,8 +2402,8 @@ class ModelUpdateHandler:
self._logger.error("Failed to fetch license info: %s", exc, exc_info=True)
return web.json_response({"success": False, "error": str(exc)}, status=500)
updated: List[Dict[str, str]] = []
errors: List[Dict[str, str]] = []
updated: List[Dict[str, Any]] = []
errors: List[Dict[str, Any]] = []
for model_id in model_ids:
license_payload = license_map.get(model_id)
if not license_payload:
@@ -2351,6 +2416,7 @@ class ModelUpdateHandler:
model_section = civitai_section.get("model")
if not isinstance(model_section, Mapping):
model_section = {}
model_section = dict(model_section)
model_section.update(resolved_payload)
civitai_section["model"] = model_section
metadata_payload["civitai"] = civitai_section
@@ -2366,7 +2432,7 @@ class ModelUpdateHandler:
)
errors.append({"filePath": metadata_path, "error": str(exc)})
response_payload = {"success": True, "updated": updated}
response_payload: Dict[str, Any] = {"success": True, "updated": updated}
missing_model_ids = [mid for mid in model_ids if mid not in license_map]
if missing_model_ids:
response_payload["missingModelIds"] = missing_model_ids
@@ -2715,6 +2781,7 @@ class ModelUpdateHandler:
civitai_payload = metadata_payload.get("civitai")
if not isinstance(civitai_payload, Mapping):
civitai_payload = {}
civitai_payload = dict(civitai_payload)
model_payload = civitai_payload.get("model")
if not isinstance(model_payload, Mapping):
@@ -2759,7 +2826,7 @@ class ModelUpdateHandler:
return aggregated
def _extract_target_model_ids(self, payload: Dict) -> Optional[List[int]]:
def _extract_target_model_ids(self, payload: Dict[str, Any]) -> Optional[List[int]]:
if not isinstance(payload, Mapping):
return None
@@ -2787,7 +2854,7 @@ class ModelUpdateHandler:
return {}
to_dict = getattr(metadata, "to_dict", None)
if callable(to_dict):
if to_dict:
try:
return to_dict()
except Exception:
@@ -2798,7 +2865,7 @@ class ModelUpdateHandler:
return {}
async def _read_json(self, request: web.Request) -> Dict:
async def _read_json(self, request: web.Request) -> Dict[str, Any]:
if not request.can_read_body:
return {}
try:
@@ -2830,7 +2897,7 @@ class ModelUpdateHandler:
record,
*,
version_context: Optional[Dict[int, Dict[str, Any]]] = None,
) -> Dict:
) -> Dict[str, Any]:
context = version_context or {}
# Check user setting for hiding early access versions
hide_early_access = False
@@ -2859,7 +2926,7 @@ class ModelUpdateHandler:
@staticmethod
def _serialize_version(
version, context: Optional[Dict[str, Any]]
) -> Dict:
) -> Dict[str, Any]:
context = context or {}
preview_override = context.get("preview_override")
preview_url = (
+148 -53
View File
@@ -10,7 +10,7 @@ import asyncio
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Tuple
from aiohttp import web
@@ -44,6 +44,22 @@ EnsureDependenciesCallable = Callable[[], Awaitable[None]]
RecipeScannerGetter = Callable[[], Any]
CivitaiClientGetter = Callable[[], Any]
# Cap concurrent preview-dimension reads across requests. With a cold LRU
# cache one page can touch up to page_size image files; 16 balances SSD and
# HDD throughput without starving the event loop.
_DIMS_READ_SEMAPHORE = asyncio.Semaphore(16)
async def _read_preview_dims(path: str) -> Optional[Tuple[int, int]]:
"""Read preview dimensions off the event loop under the concurrency cap.
PIL I/O runs in a worker thread so it never blocks the event loop, and the
semaphore bounds how many files are opened at once even when many list
requests land together.
"""
async with _DIMS_READ_SEMAPHORE:
return await asyncio.to_thread(ExifUtils.get_image_dimensions, path)
@dataclass(frozen=True)
class RecipeHandlerSet:
@@ -246,7 +262,8 @@ class RecipeListingHandler:
recursive=recursive,
)
for item in result.get("items", []):
items = result.get("items", [])
for item in items:
file_path = item.get("file_path")
if file_path:
item["file_url"] = self.format_recipe_file_url(file_path)
@@ -255,6 +272,26 @@ class RecipeListingHandler:
item.setdefault("loras", [])
item.setdefault("base_model", "")
# Batch preview dimension reads with asyncio.gather. The previous
# loop awaited asyncio.to_thread once per item, so a page_size=100
# request submitted 100 sequential thread calls (50-300ms cold-page
# latency). gather runs them concurrently while the semaphore caps
# disk opens; dimensions stay omitted (not null) when a preview has
# no readable size (video, missing file).
to_read = [
(i, item.get("file_path"))
for i, item in enumerate(items)
if item.get("file_path")
]
if to_read:
dims_list = await asyncio.gather(
*(_read_preview_dims(path) for _, path in to_read)
)
for (idx, _), dims in zip(to_read, dims_list):
if dims:
item = items[idx]
item["width"], item["height"] = dims
return web.json_response(result)
except Exception as exc:
self._logger.error("Error retrieving recipes: %s", exc, exc_info=True)
@@ -1045,10 +1082,10 @@ class RecipeManagementHandler:
*,
image_url: str,
name: str,
lora_entries: list,
checkpoint_entry: dict,
gen_params_request: dict,
tags: list,
lora_entries: list[Any],
checkpoint_entry: Dict[str, Any] | None,
gen_params_request: Dict[str, Any] | None,
tags: list[Any],
base_model: str,
source_path: str,
) -> web.Response:
@@ -1081,6 +1118,12 @@ class RecipeManagementHandler:
_original_image_url,
) = await self._download_remote_media(image_url)
# Build a version-cached map of local model hashes to cache items so
# CivitaiApiMetadataParser can skip CivitAI API calls for models that
# exist on disk. Built once and shared by every parse pass below.
local_cache = await recipe_scanner.build_local_hash_cache()
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
# Extract embedded EXIF metadata (offloaded to thread pool in this call)
embedded_gen_params = {}
parsed_embedded = None
@@ -1102,9 +1145,16 @@ class RecipeManagementHandler:
)
)
if parser:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
parsed_embedded = await parser.parse_metadata(
raw_embedded,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if parsed_embedded and "gen_params" in parsed_embedded:
embedded_gen_params = parsed_embedded["gen_params"]
else:
@@ -1135,9 +1185,16 @@ class RecipeManagementHandler:
civitai_inner_meta
)
if parser:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if civitai_parsed and "gen_params" in civitai_parsed:
# Merge: API gen_params override EXIF at field level,
# EXIF fills in fields the API doesn't have.
@@ -1641,7 +1698,7 @@ class RecipeManagementHandler:
if not provider:
return ""
version_info = await provider.get_model_version_info(version_id)
version_info = await provider.get_model_version_info(str(version_id))
if isinstance(version_info, tuple):
version_info = version_info[0]
@@ -1761,6 +1818,12 @@ class RecipeManagementHandler:
await self._download_remote_media(image_url)
)
# Build a version-cached map of local model hashes to cache items so
# CivitaiApiMetadataParser can skip CivitAI API calls for models that
# exist on disk. Built once and shared by every parse pass below.
local_cache = await recipe_scanner.build_local_hash_cache()
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
# Extract embedded EXIF metadata
embedded_gen_params = {}
parsed_embedded = None
@@ -1782,9 +1845,16 @@ class RecipeManagementHandler:
)
)
if parser:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
parsed_embedded = await parser.parse_metadata(
raw_embedded,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if parsed_embedded and "gen_params" in parsed_embedded:
embedded_gen_params = parsed_embedded["gen_params"]
finally:
@@ -1822,9 +1892,16 @@ class RecipeManagementHandler:
)
)
if parser:
parsed_embedded = await parser.parse_metadata(
raw_orig, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
parsed_embedded = await parser.parse_metadata(
raw_orig,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
parsed_embedded = await parser.parse_metadata(
raw_orig, recipe_scanner=recipe_scanner
)
if (
parsed_embedded
and "gen_params" in parsed_embedded
@@ -1858,9 +1935,16 @@ class RecipeManagementHandler:
civitai_inner_meta
)
if parser:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if civitai_parsed and "gen_params" in civitai_parsed:
# Merge: API gen_params override EXIF at field level,
# EXIF fills in fields the API doesn't have.
@@ -2072,33 +2156,44 @@ class RecipeManagementHandler:
parsed_input = {**image_data, **inner_meta}
parsed_input.pop("meta", None)
# Build a local cache of {hash cache_item} so the parser can
# skip CivitAI API calls for models that exist on disk.
local_cache: Dict[str, Dict[str, Any]] = {}
lora_scanner = getattr(recipe_scanner, "_lora_scanner", None)
if lora_scanner and model_hash:
try:
parent_cache_data = await lora_scanner.get_cached_data()
for item in getattr(parent_cache_data, "raw_data", []):
if item.get("sha256", "").lower() == model_hash.lower():
local_cache[model_hash.lower()] = item
# Compute AutoV3 so the parser can also match on
# that hash type (CivitAI metadata resources use
# AutoV3).
file_path = item.get("file_path")
if file_path and os.path.exists(file_path):
try:
from ...utils.file_utils import (
calculate_autov3,
)
autov3 = calculate_autov3(file_path)
if autov3:
local_cache[autov3.lower()] = item
except Exception:
pass
break
except Exception:
pass
# Build the shared local hash cache so the parser can skip CivitAI
# API calls for models that exist on disk.
local_cache: Dict[str, Dict[str, Any]] = (
await recipe_scanner.build_local_hash_cache()
)
# Bounded supplement for un-backfilled parents. The shared builder
# never computes autov3; when the parent model exists on disk but
# its cached entry has no stored AutoV3, compute it for that single
# file and register the AutoV3 key so the parser can also match on
# that hash type (CivitAI metadata resources use AutoV3). This runs
# whenever the parent is found with an empty autov3, independent of
# whether the sha256 key is already present in the shared cache.
if model_hash:
lora_scanner = getattr(recipe_scanner, "_lora_scanner", None)
if lora_scanner:
try:
parent_cache_data = await lora_scanner.get_cached_data()
for item in getattr(parent_cache_data, "raw_data", []):
if item.get("sha256", "").lower() == model_hash.lower():
autov3 = (item.get("autov3") or "").lower()
if not autov3:
file_path = item.get("file_path")
if file_path and os.path.exists(file_path):
try:
from ...utils.file_utils import (
calculate_autov3,
)
autov3 = (
calculate_autov3(file_path) or ""
).lower()
except Exception:
pass
if autov3:
local_cache[autov3] = item
break
except Exception:
pass
parser = self._analysis_service._recipe_parser_factory.create_parser(
parsed_input
@@ -2130,10 +2225,10 @@ class RecipeManagementHandler:
parent_model_id: int | None = None
parent_version_name: str | None = None
parent_model_name: str | None = None
# Prefer sha256 key; fall back to any cached entry.
# Resolve the parent strictly by its sha256 key. There is no
# arbitrary fallback: with a full-library cache, picking any entry
# would corrupt the isDeleted reconciliation below.
parent_item = local_cache.get(model_hash.lower()) if model_hash else None
if parent_item is None and local_cache:
parent_item = next(iter(local_cache.values()))
if parent_item:
civ = parent_item.get("civitai") or {}
if isinstance(civ, dict):
@@ -2349,7 +2444,7 @@ class RecipeAnalysisHandler:
content_type = request.headers.get("Content-Type", "")
if "multipart/form-data" in content_type:
reader = await request.multipart()
field = await reader.next()
field: Any = await reader.next()
if field is None or field.name != "image":
raise RecipeValidationError("No image field found")
image_chunks = bytearray()
+6 -71
View File
@@ -1,8 +1,8 @@
import asyncio
import logging
from aiohttp import web
from typing import Dict
from server import PromptServer # type: ignore
from typing import Any, Dict
from server import PromptServer # pyright: ignore[reportMissingImports]
from .base_model_routes import BaseModelRoutes
from .model_route_registrar import ModelRouteRegistrar
@@ -31,13 +31,13 @@ class LoraRoutes(BaseModelRoutes):
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application):
def setup_routes(self, app: web.Application, prefix: str = "loras"):
"""Setup LoRA routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'loras' prefix (includes page route)
super().setup_routes(app, "loras")
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup LoRA-specific routes"""
@@ -73,7 +73,7 @@ class LoraRoutes(BaseModelRoutes):
"POST", "/api/lm/{prefix}/get_trigger_words", prefix, self.get_trigger_words
)
def _parse_specific_params(self, request: web.Request) -> Dict:
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse LoRA-specific parameters"""
params = {}
@@ -119,25 +119,6 @@ class LoraRoutes(BaseModelRoutes):
logger.error(f"Error getting letter counts: {e}")
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_notes(self, request: web.Request) -> web.Response:
"""Get notes for a specific LoRA file"""
try:
lora_name = request.query.get("name")
if not lora_name:
return web.Response(text="Lora file name is required", status=400)
notes = await self.service.get_lora_notes(lora_name)
if notes is not None:
return web.json_response({"success": True, "notes": notes})
else:
return web.json_response(
{"success": False, "error": "LoRA not found in cache"}, status=404
)
except Exception as e:
logger.error(f"Error getting lora notes: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_trigger_words(self, request: web.Request) -> web.Response:
"""Get trigger words for a specific LoRA file"""
try:
@@ -168,52 +149,6 @@ class LoraRoutes(BaseModelRoutes):
logger.error(f"Error getting lora usage tips by path: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_preview_url(self, request: web.Request) -> web.Response:
"""Get the static preview URL for a LoRA file"""
try:
lora_name = request.query.get("name")
if not lora_name:
return web.Response(text="Lora file name is required", status=400)
preview_url = await self.service.get_lora_preview_url(lora_name)
if preview_url:
return web.json_response({"success": True, "preview_url": preview_url})
else:
return web.json_response(
{
"success": False,
"error": "No preview URL found for the specified lora",
},
status=404,
)
except Exception as e:
logger.error(f"Error getting lora preview URL: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_civitai_url(self, request: web.Request) -> web.Response:
"""Get the Civitai URL for a LoRA file"""
try:
lora_name = request.query.get("name")
if not lora_name:
return web.Response(text="Lora file name is required", status=400)
result = await self.service.get_lora_civitai_url(lora_name)
if result["civitai_url"]:
return web.json_response({"success": True, **result})
else:
return web.json_response(
{
"success": False,
"error": "No Civitai data found for the specified lora",
},
status=404,
)
except Exception as e:
logger.error(f"Error getting lora Civitai URL: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_random_loras(self, request: web.Request) -> web.Response:
"""Get random LoRAs based on filters and strength ranges"""
try:
@@ -337,7 +272,7 @@ class LoraRoutes(BaseModelRoutes):
graph_identifier = entry.get("graph_id")
try:
parsed_node_id = int(node_identifier)
parsed_node_id = int(node_identifier) # pyright: ignore[reportArgumentType]
except (TypeError, ValueError):
parsed_node_id = node_identifier
+2 -2
View File
@@ -5,7 +5,7 @@ miscellaneous endpoints share a consistent registration flow.
"""
from dataclasses import dataclass
from typing import Callable, Iterable, Mapping
from typing import Any, Callable, Iterable, Mapping
from aiohttp import web
@@ -147,7 +147,7 @@ class MiscRouteRegistrar:
handler_lookup[definition.handler_name],
)
def _bind(self, method: str, path: str, handler: Callable) -> None:
def _bind(self, method: str, path: str, handler: Callable[..., Any]) -> None:
add_method_name = self._METHOD_MAP[method.upper()]
add_method = getattr(self._app.router, add_method_name)
add_method(path, handler)
+1 -1
View File
@@ -7,7 +7,7 @@ import os
from typing import Awaitable, Callable, Mapping
from aiohttp import web
from server import PromptServer # type: ignore
from server import PromptServer # pyright: ignore[reportMissingImports]
from ..services.metadata_service import (
get_metadata_archive_manager,
+4 -4
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable, Iterable, Mapping
from typing import Any, Callable, Iterable, Mapping
from aiohttp import web
@@ -174,15 +174,15 @@ class ModelRouteRegistrar:
handler_lookup[definition.handler_name],
)
def add_route(self, method: str, path: str, handler: Callable) -> None:
def add_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
self._bind_route(method, path, handler)
def add_prefixed_route(
self, method: str, path_template: str, prefix: str, handler: Callable
self, method: str, path_template: str, prefix: str, handler: Callable[..., Any]
) -> None:
self._bind_route(method, path_template.replace("{prefix}", prefix), handler)
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
add_method_name = self._METHOD_MAP[method.upper()]
add_method = getattr(self._app.router, add_method_name)
add_method(path, handler)
+2 -2
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable, Mapping
from typing import Any, Callable, Mapping
from aiohttp import web
@@ -105,7 +105,7 @@ class RecipeRouteRegistrar:
handler = handler_lookup[definition.handler_name]
self._bind_route(definition.method, definition.path, handler)
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
add_method_name = self._METHOD_MAP[method.upper()]
add_method = getattr(self._app.router, add_method_name)
add_method(path, handler)
+11 -10
View File
@@ -40,10 +40,11 @@ class StatsRoutes:
"""Route handlers for Statistics page and API endpoints"""
def __init__(self):
self.lora_scanner = None
self.checkpoint_scanner = None
self.embedding_scanner = None
self.usage_stats = None
self.lora_scanner: Any = None
self.checkpoint_scanner: Any = None
self.embedding_scanner: Any = None
self.usage_stats: Any = None
self._i18n_filter_added = False
self.template_env = jinja2.Environment(
loader=jinja2.FileSystemLoader(config.templates_path),
autoescape=True
@@ -95,9 +96,9 @@ class StatsRoutes:
server_i18n.set_locale(user_language)
# 为模板环境添加i18n过滤器
if not hasattr(self.template_env, '_i18n_filter_added'):
if not self._i18n_filter_added:
self.template_env.filters['t'] = server_i18n.create_template_filter()
self.template_env._i18n_filter_added = True
self._i18n_filter_added = True
template = self.template_env.get_template('statistics.html')
rendered = template.render(
@@ -549,7 +550,7 @@ class StatsRoutes:
'error': str(e)
}, status=500)
def _count_unused_models(self, models: List[Dict], usage_data: Dict) -> int:
def _count_unused_models(self, models: List[Dict[str, Any]], usage_data: Dict[str, Any]) -> int:
"""Count models that have never been used"""
used_hashes = set(usage_data.keys())
unused_count = 0
@@ -560,7 +561,7 @@ class StatsRoutes:
return unused_count
def _get_top_used_models(self, usage_data: Dict, model_map: Dict, limit: int) -> List[Dict]:
def _get_top_used_models(self, usage_data: Dict[str, Any], model_map: Dict[str, Any], limit: int) -> List[Dict[str, Any]]:
"""Get top used models with their metadata"""
sorted_usage = sorted(usage_data.items(), key=lambda x: x[1].get('total', 0), reverse=True)
@@ -578,7 +579,7 @@ class StatsRoutes:
return top_models
def _get_usage_timeline(self, usage_data: Dict, days: int) -> List[Dict]:
def _get_usage_timeline(self, usage_data: Dict[str, Any], days: int) -> List[Dict[str, Any]]:
"""Get usage timeline for the past N days"""
timeline = []
today = datetime.now()
@@ -614,7 +615,7 @@ class StatsRoutes:
return list(reversed(timeline)) # Oldest to newest
def _format_size(self, size_bytes: int) -> str:
def _format_size(self, size_bytes: float) -> str:
"""Format file size in human readable format"""
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
if size_bytes < 1024.0:
+16 -13
View File
@@ -6,7 +6,7 @@ import shutil
import tempfile
import asyncio
from aiohttp import web, ClientError
from typing import Dict, List
from typing import Any, Dict, List, cast
from ..utils.settings_paths import ensure_settings_file
from ..services.downloader import get_downloader
@@ -467,9 +467,10 @@ class UpdateRoutes:
if not success:
logger.error(f"Failed to fetch release info: {data}")
return False, ""
zip_url = data.get("zipball_url")
version = data.get("tag_name", "unknown")
release_payload = cast(dict[str, Any], data)
zip_url = release_payload.get("zipball_url", "")
version = release_payload.get("tag_name", "unknown")
# Download ZIP to temporary file
with tempfile.NamedTemporaryFile(delete=False, suffix=".zip") as tmp_zip:
@@ -580,9 +581,10 @@ class UpdateRoutes:
logger.warning("Failed to fetch GitHub commit: %s", data)
return "main", [], 0, ""
commit_sha = data.get('sha', '')[:7]
commit_message = data.get('commit', {}).get('message', '')
commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10]
commit_payload = cast(dict[str, Any], data)
commit_sha = commit_payload.get('sha', '')[:7]
commit_message = commit_payload.get('commit', {}).get('message', '')
commit_date = commit_payload.get('commit', {}).get('committer', {}).get('date', '')[:10]
version = f"main-{commit_sha}"
changelog = [commit_message] if commit_message else []
@@ -598,10 +600,11 @@ class UpdateRoutes:
custom_headers={'Accept': 'application/vnd.github+json'}
)
if c_ok:
if c_data.get('status') in ('ahead', 'diverged'):
behind_by = c_data.get('ahead_by', 0)
compare_payload = cast(dict[str, Any], c_data)
if compare_payload.get('status') in ('ahead', 'diverged'):
behind_by = compare_payload.get('ahead_by', 0)
else:
behind_by = c_data.get('behind_by', 0)
behind_by = compare_payload.get('behind_by', 0)
return version, changelog, behind_by, commit_date
@@ -706,7 +709,7 @@ class UpdateRoutes:
logger.info(f"Successfully updated to {new_version}")
return True, new_version
except git.exc.GitError as e:
except git.exc.GitError as e: # pyright: ignore[reportAttributeAccessIssue]
logger.error(f"Git error during update: {e}")
return False, ""
except Exception as e:
@@ -767,7 +770,7 @@ class UpdateRoutes:
return git_info
@staticmethod
async def _get_remote_version() -> tuple[str, List[str], List[Dict]]:
async def _get_remote_version() -> tuple[str, List[str], List[Dict[str, Any]]]:
"""
Fetch remote version from GitHub
Returns:
@@ -789,7 +792,7 @@ class UpdateRoutes:
# Parse releases
releases = []
for i, release in enumerate(data):
for i, release in enumerate(cast(list[dict[str, Any]], data)):
version = release.get('tag_name', '')
if not version.startswith('v'):
version = f"v{version}"
+1 -1
View File
@@ -117,7 +117,7 @@ def _render_prompt(template: str, variables: Dict[str, Any]) -> str:
Uses simple regex substitution no Jinja2 dependency needed.
"""
def replace(match: re.Match) -> str:
def replace(match: re.Match[str]) -> str:
key = match.group(1).strip()
value = variables.get(key, "")
if isinstance(value, (dict, list)):
+1 -1
View File
@@ -295,7 +295,7 @@ class PostProcessor:
normalises every tag to lowercase for case-insensitive dedup.
"""
merged: List[str] = []
seen: set = set()
seen: set[str] = set()
for tag in list(existing) + list(new):
t = tag.strip().lower()
if t and t not in seen:
+1 -1
View File
@@ -49,7 +49,7 @@ _FRONTMATTER_RE = re.compile(
)
def _parse_skill_file(path: Path) -> tuple[dict, str]:
def _parse_skill_file(path: Path) -> tuple[dict[str, Any], str]:
"""Read a prompt definition file (``prompt.md`` or legacy ``SKILL.md``) and
return (frontmatter_dict, body_text).
@@ -9,7 +9,7 @@ from __future__ import annotations
import html as html_module
import re
from typing import List, Tuple
from typing import Any, List, Tuple
_REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
@@ -18,10 +18,10 @@ _REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
def extract_simple_markdown_images(
markdown_text: str,
repo: str,
existing_urls: set | None = None,
existing_urls: set[str] | None = None,
default_width: int = 512,
default_height: int = 512,
) -> list[dict]:
) -> list[dict[str, Any]]:
"""Extract standalone markdown images from the README body.
Matches ``![alt](url)`` on lines that are NOT part of a markdown table
@@ -36,8 +36,8 @@ def extract_simple_markdown_images(
return []
base_url = f"https://huggingface.co/{repo}/resolve/main"
images: list[dict] = []
seen_urls: set = set(existing_urls) if existing_urls else set()
images: list[dict[str, Any]] = []
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
# Collect lines that are NOT inside fenced code blocks
lines = markdown_text.split("\n")
@@ -86,10 +86,10 @@ def extract_simple_markdown_images(
def extract_html_img_tags(
markdown_text: str,
repo: str,
existing_urls: set | None = None,
existing_urls: set[str] | None = None,
default_width: int = 512,
default_height: int = 512,
) -> list[dict]:
) -> list[dict[str, Any]]:
"""Extract image URLs from HTML ``<img src=\"...\">`` tags in the README.
Many HF collection repos (e.g. ``deadman44/Z-Image_LoRA``) use raw HTML
@@ -103,8 +103,8 @@ def extract_html_img_tags(
return []
base_url = f"https://huggingface.co/{repo}/resolve/main"
images: list[dict] = []
seen_urls: set = set(existing_urls) if existing_urls else set()
images: list[dict[str, Any]] = []
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
for m in re.finditer(
r'<img\s[^>]*src=\"([^\"]+)\"',
@@ -175,7 +175,7 @@ def extract_gallery_images(
repo: str,
default_width: int = 512,
default_height: int = 512,
) -> List[dict]:
) -> List[dict[str, Any]]:
"""Extract widget/gallery images from the YAML frontmatter of a HF README.
Args:
@@ -196,7 +196,7 @@ def extract_gallery_images(
if not frontmatter:
return []
images: List[dict] = []
images: List[dict[str, Any]] = []
base_url = f"https://huggingface.co/{repo}/resolve/main"
w = default_width or 512
h = default_height or 512
@@ -258,7 +258,7 @@ def extract_gallery_images(
text = raw_text
if url:
image: dict = {
image: dict[str, Any] = {
"url": url,
"type": "image",
"nsfwLevel": 0,
@@ -276,10 +276,10 @@ def extract_gallery_images(
def extract_gallery_table_images(
markdown_text: str,
repo: str,
existing_urls: set | None = None,
existing_urls: set[str] | None = None,
default_width: int = 512,
default_height: int = 512,
) -> list[dict]:
) -> list[dict[str, Any]]:
"""Extract images from ``| Preview | Prompt |`` markdown gallery tables.
Many HF READMEs include a sample-gallery table in the body (outside
@@ -295,8 +295,8 @@ def extract_gallery_table_images(
return []
base_url = f"https://huggingface.co/{repo}/resolve/main"
images: list[dict] = []
seen_urls: set = set(existing_urls) if existing_urls else set()
images: list[dict[str, Any]] = []
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
lines = markdown_text.split("\n")
n = len(lines)
i = 0
@@ -514,7 +514,7 @@ def _strip_standalone_images(text: str) -> str:
URL was stripped entirely, making it impossible for the LLM to return
a ``preview_url`` for repos that use HTML ``<img>`` tags exclusively.
"""
def _img_to_md(match: re.Match) -> str:
def _img_to_md(match: re.Match[str]) -> str:
"""Convert an ``<img>`` tag to markdown image syntax ``![alt](src)``."""
tag = match.group(0)
src_m = re.search(r'src="([^"]+)"', tag) or re.search(r"src='([^']+)'", tag)
@@ -942,7 +942,7 @@ def _strip_badge_images(text: str) -> str:
"twitter", "colab", "gradio", "space",
)
def _should_remove(m: re.Match) -> str:
def _should_remove(m: re.Match[str]) -> str:
alt = (m.group(1) or "").lower()
for kw in badge_keywords:
if kw in alt:
+7 -3
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
from __future__ import annotations
import asyncio
@@ -23,7 +27,7 @@ logger = logging.getLogger(__name__)
def _try_certifi_ca_path() -> str | None:
"""Return the certifi CA bundle path if available, else None."""
try:
import certifi # type: ignore[import-untyped]
import certifi # pyright: ignore[reportMissingTypeStubs]
path = certifi.where()
if os.path.isfile(path):
@@ -84,7 +88,7 @@ class Aria2Downloader:
self._transfers: Dict[str, Aria2Transfer] = {}
self._poll_interval = 0.5
self._state_store = Aria2TransferStateStore()
self._stderr_reader_task: Optional[asyncio.Task] = None
self._stderr_reader_task: Optional[asyncio.Task[Any]] = None
@property
def is_running(self) -> bool:
@@ -190,7 +194,7 @@ class Aria2Downloader:
download_id,
)
options: Dict[str, str] = {
options: Dict[str, Any] = {
"dir": save_dir,
"out": out_name,
"continue": "true",
+3 -3
View File
@@ -8,7 +8,7 @@ from filename, base_model, and CivitAI version name — no manual tagging requir
from __future__ import annotations
import re
from typing import Dict, List, Set
from typing import Any, Dict, List, Set
# ── Tag category definitions ──────────────────────────────────────────
# Each category maps a display label to a regex pattern.
@@ -52,7 +52,7 @@ AUTO_TAG_GROUPS = {
DEFAULT_ENABLED_GROUPS = {"mode", "video"}
def _collect_sources(model_data: Dict) -> List[str]:
def _collect_sources(model_data: Dict[str, Any]) -> List[str]:
"""Collect all text sources from model data for tag matching."""
sources: List[str] = []
@@ -73,7 +73,7 @@ def _collect_sources(model_data: Dict) -> List[str]:
return sources
def extract_auto_tags(model_data: Dict) -> List[str]:
def extract_auto_tags(model_data: Dict[str, Any]) -> List[str]:
"""Extract auto-detected tags from model metadata.
Uses a two-layer approach:
+144
View File
@@ -0,0 +1,144 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
"""Backfill the AutoV3 checked state for models loaded from a persisted snapshot.
The SQLite persistent cache predates the AutoV3 feature, so entries hydrated
from it have a NULL ``autov3`` column (the "not checked yet" state). This
service computes the embedded AutoV3 hash for each such model once per
process and persists it through the scanner's single write path
(:meth:`ModelScanner.update_autov3_for_model`), marking every visited row so a
subsequent run finds nothing left to do.
Three-state contract honored here:
- ``NULL`` (sqlite) / absent (dict) = not checked yet backfill computes it
- ``''`` (sqlite/dict) / JSON null = checked, no value available never recompute
- 12-char lowercase hex = value never recompute
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import threading
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING: # pragma: no cover - type-check only; runtime imports are local
from .model_scanner import ModelScanner
logger = logging.getLogger(__name__)
def _resolve_autov3(file_path: str) -> str:
"""Resolve the AutoV3 hash for a model file.
Prefers the Civitai AutoV3 reported for the file whose SHA256 matches
(the authoritative value for recipe matching); falls back to the embedded
safetensors header hash. Returns ``''`` when neither is available.
"""
try:
metadata_path = f"{os.path.splitext(file_path)[0]}.metadata.json"
if os.path.exists(metadata_path):
with open(metadata_path, "r", encoding="utf-8") as handle:
payload = json.load(handle)
if isinstance(payload, dict):
from ..utils.models import autov3_from_civitai_files # local import avoids cycles
sha256 = (payload.get("sha256") or "").lower()
civitai_autov3 = autov3_from_civitai_files(payload.get("civitai"), sha256)
if civitai_autov3:
return civitai_autov3
except Exception:
pass
from ..utils.file_utils import calculate_autov3 # local import avoids cycles
return calculate_autov3(file_path) or ""
class Autov3BackfillService:
"""Compute and persist AutoV3 hashes for models missing a checked state."""
_instance: Optional["Autov3BackfillService"] = None
_instance_lock = threading.Lock()
def __init__(self) -> None:
# Re-entrancy guard per model type: scanners for different model types
# initialize concurrently (lora_manager.py), so a global guard would
# silently skip every type but the first to start. Each model type
# runs its own backfill; a duplicate trigger for the same type no-ops.
self._running_types: set[str] = set()
@classmethod
def get_instance(cls) -> "Autov3BackfillService":
"""Return the process-wide singleton instance."""
if cls._instance is None:
with cls._instance_lock:
if cls._instance is None:
cls._instance = cls()
return cls._instance
async def backfill(self, scanner: "ModelScanner") -> int:
"""Compute AutoV3 for every un-checked model of ``scanner.model_type``.
Each candidate file is read once via :func:`~py.utils.file_utils.calculate_autov3`
(cheap: safetensors header only) and the result is persisted through
``scanner.update_autov3_for_model``. Files that no longer exist on
disk are skipped they are intentionally NOT marked, because scanner
cleanup removes the stale row later.
Returns:
The number of models successfully updated. Never raises; on any
failure a warning is logged and ``0`` is returned. A duplicate
trigger for a model type that is already being backfilled returns
``0`` immediately; different model types run concurrently.
"""
model_type = scanner.model_type
if model_type in self._running_types:
return 0
self._running_types.add(model_type)
try:
# Local imports avoid import cycles at module load time.
from .persistent_model_cache import get_persistent_cache
from ..utils.file_utils import calculate_autov3
persistent = getattr(scanner, "_persistent_cache", None) or get_persistent_cache()
paths = persistent.get_models_missing_autov3(model_type)
loop = asyncio.get_running_loop()
count = 0
for path in paths:
# A file that no longer exists must not be marked; scanner
# cleanup removes the stale row later. The existence check and
# hash resolution run in the executor so the loop stays
# responsive to API requests while the backfill iterates a
# large library.
if not await loop.run_in_executor(None, os.path.exists, path):
continue
autov3 = await loop.run_in_executor(None, _resolve_autov3, path)
if await scanner.update_autov3_for_model(model_type, path, autov3):
count += 1
if paths:
logger.info(
"AutoV3 backfill: updated %d/%d models for %s",
count,
len(paths),
model_type,
)
else:
# Steady state after the first run: nothing left to backfill.
logger.debug("AutoV3 backfill: nothing to process for %s", model_type)
return count
except Exception as exc:
logger.warning(
"AutoV3 backfill failed for %s: %s",
getattr(scanner, "model_type", "?"),
exc,
)
return 0
finally:
self._running_types.discard(model_type)
+4
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
from __future__ import annotations
import asyncio
+146 -67
View File
@@ -2,7 +2,7 @@ from abc import ABC, abstractmethod
import asyncio
import re
import random
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
from typing import Any, Awaitable, Dict, List, Optional, Type, Union, TYPE_CHECKING, cast
import logging
import os
import time
@@ -70,24 +70,24 @@ class BaseModelService(ABC):
page: int,
page_size: int,
sort_by: str = "name",
folder: str = None,
folder_include: list = None,
folder_exclude: list = None,
search: str = None,
folder: str | None = None,
folder_include: list[str] | None = None,
folder_exclude: list[str] | None = None,
search: str | None = None,
fuzzy_search: bool = False,
base_models: list = None,
model_types: list = None,
base_models: list[str] | None = None,
model_types: list[str] | None = None,
tags: Optional[Dict[str, str]] = None,
auto_tags: Optional[Dict[str, str]] = None,
search_options: dict = None,
hash_filters: dict = None,
search_options: dict[str, Any] | None = None,
hash_filters: dict[str, Any] | None = None,
favorites_only: bool = False,
update_available_only: bool = False,
credit_required: Optional[bool] = None,
allow_selling_generated_content: Optional[bool] = None,
tag_logic: str = "any",
**kwargs,
) -> Dict:
) -> Dict[str, Any]:
"""Get paginated and filtered model data"""
overall_start = time.perf_counter()
@@ -178,8 +178,8 @@ class BaseModelService(ABC):
ufs = self.settings.get("version_grouping", "same_base")
group_by_base = ufs == "same_base"
model_groups: Dict[Any, List[Dict]] = {}
ungrouped_standalone: List[Dict] = []
model_groups: Dict[Any, List[Dict[str, Any]]] = {}
ungrouped_standalone: List[Dict[str, Any]] = []
for item in sorted_data:
mid = self._extract_group_key(item)
if mid is None:
@@ -249,7 +249,7 @@ class BaseModelService(ABC):
filter_duration = time.perf_counter() - t1
post_filter_count = len(filtered_data)
annotated_for_filter: Optional[List[Dict]] = None
annotated_for_filter: Optional[List[Dict[str, Any]]] = None
t2 = time.perf_counter()
if update_available_only:
annotated_for_filter = await self._annotate_update_flags(filtered_data)
@@ -296,11 +296,11 @@ class BaseModelService(ABC):
page: int,
page_size: int,
sort_by: str = "name",
search: str = None,
search: str | None = None,
fuzzy_search: bool = False,
search_options: dict = None,
search_options: dict[str, Any] | None = None,
**kwargs,
) -> Dict:
) -> Dict[str, Any]:
"""Get paginated excluded model data."""
excluded_paths = list(self.scanner.get_excluded_models())
excluded_entries: List[Dict[str, Any]] = []
@@ -326,7 +326,7 @@ class BaseModelService(ABC):
]
persist_current_cache = getattr(self.scanner, "_persist_current_cache", None)
if callable(persist_current_cache):
await persist_current_cache()
await cast(Awaitable[Any], persist_current_cache())
excluded_entries = self._sort_entries(excluded_entries, sort_by)
@@ -444,39 +444,50 @@ class BaseModelService(ABC):
return entry
async def _apply_hash_filters(
self, data: List[Dict], hash_filters: Dict
) -> List[Dict]:
"""Apply hash-based filtering"""
self, data: List[Dict[str, Any]], hash_filters: Dict[str, Any]
) -> List[Dict[str, Any]]:
"""Apply hash-based filtering (SHA256 and AutoV3)."""
def matches_hash_set(item: Dict[str, Any], hash_set: set[str]) -> bool:
"""Check whether an item matches any hash in the set.
Compares the item's ``sha256`` field and its non-empty ``autov3``
field, both case-insensitively.
"""
if item.get("sha256", "").lower() in hash_set:
return True
autov3 = item.get("autov3", "")
return bool(autov3) and autov3.lower() in hash_set
single_hash = hash_filters.get("single_hash")
multiple_hashes = hash_filters.get("multiple_hashes")
if single_hash:
# Filter by single hash
single_hash = single_hash.lower()
# Filter by single hash (SHA256 or AutoV3)
return [
item for item in data if item.get("sha256", "").lower() == single_hash
item for item in data if matches_hash_set(item, {single_hash.lower()})
]
elif multiple_hashes:
# Filter by multiple hashes
hash_set = set(hash.lower() for hash in multiple_hashes)
return [item for item in data if item.get("sha256", "").lower() in hash_set]
# Filter by multiple hashes (SHA256 or AutoV3)
hash_set = {hash.lower() for hash in multiple_hashes}
return [item for item in data if matches_hash_set(item, hash_set)]
return data
async def _apply_common_filters(
self,
data: List[Dict],
folder: str = None,
folder_include: list = None,
folder_exclude: list = None,
base_models: list = None,
model_types: list = None,
data: List[Dict[str, Any]],
folder: str | None = None,
folder_include: list[str] | None = None,
folder_exclude: list[str] | None = None,
base_models: list[str] | None = None,
model_types: list[str] | None = None,
tags: Optional[Dict[str, str]] = None,
auto_tags: Optional[Dict[str, str]] = None,
favorites_only: bool = False,
search_options: dict = None,
search_options: dict[str, Any] | None = None,
tag_logic: str = "any",
) -> List[Dict]:
) -> List[Dict[str, Any]]:
"""Apply common filters that work across all model types"""
normalized_options = self.search_strategy.normalize_options(search_options)
criteria = FilterCriteria(
@@ -495,24 +506,24 @@ class BaseModelService(ABC):
async def _apply_search_filters(
self,
data: List[Dict],
data: List[Dict[str, Any]],
search: str,
fuzzy_search: bool,
search_options: dict,
) -> List[Dict]:
search_options: dict[str, Any] | None,
) -> List[Dict[str, Any]]:
"""Apply search filtering"""
normalized_options = self.search_strategy.normalize_options(search_options)
return self.search_strategy.apply(
data, search, normalized_options, fuzzy_search
)
async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]:
async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]:
"""Apply model-specific filters - to be overridden by subclasses if needed"""
return data
async def _apply_credit_required_filter(
self, data: List[Dict], credit_required: bool
) -> List[Dict]:
self, data: List[Dict[str, Any]], credit_required: bool
) -> List[Dict[str, Any]]:
"""Apply credit required filtering based on license_flags.
Args:
@@ -542,8 +553,8 @@ class BaseModelService(ABC):
return filtered_data
async def _apply_allow_selling_filter(
self, data: List[Dict], allow_selling: bool
) -> List[Dict]:
self, data: List[Dict[str, Any]], allow_selling: bool
) -> List[Dict[str, Any]]:
"""Apply allow selling generated content filtering based on license_flags.
Args:
@@ -575,8 +586,8 @@ class BaseModelService(ABC):
async def _annotate_update_flags(
self,
items: List[Dict],
) -> List[Dict]:
items: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""Attach an update_available flag to each response item.
Items without a civitai model id default to False.
@@ -591,7 +602,7 @@ class BaseModelService(ABC):
item["update_available"] = False
return annotated
id_to_items: Dict[int, List[Dict]] = {}
id_to_items: Dict[int, List[Dict[str, Any]]] = {}
ordered_ids: List[int] = []
for item in annotated:
model_id = self._extract_model_id(item)
@@ -628,7 +639,7 @@ class BaseModelService(ABC):
record_method = getattr(self.update_service, "get_records_bulk", None)
if callable(record_method):
try:
records = await record_method(self.model_type, ordered_ids)
records = await cast(Awaitable[Any], record_method(self.model_type, ordered_ids))
resolved = {
model_id: record.has_update(hide_early_access=hide_early_access)
for model_id, record in records.items()
@@ -648,11 +659,11 @@ class BaseModelService(ABC):
bulk_method = getattr(self.update_service, "has_updates_bulk", None)
if callable(bulk_method):
try:
resolved = await bulk_method(
resolved = await cast(Awaitable[Any], bulk_method(
self.model_type,
ordered_ids,
hide_early_access=hide_early_access,
)
))
except Exception as exc:
logger.error(
"Failed to resolve update status in bulk for %s models (%s): %s",
@@ -714,7 +725,7 @@ class BaseModelService(ABC):
return annotated
@staticmethod
def _extract_hf_group_key(item: Dict) -> Optional[str]:
def _extract_hf_group_key(item: Dict[str, Any]) -> Optional[str]:
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
hf_url = item.get("hf_url") if isinstance(item, dict) else None
if not hf_url or not isinstance(hf_url, str):
@@ -727,7 +738,7 @@ class BaseModelService(ABC):
return f"hf:{m.group(1)}"
@staticmethod
def _extract_group_key(item: Dict) -> Union[int, str, None]:
def _extract_group_key(item: Dict[str, Any]) -> Union[int, str, None]:
"""Return the group identity key: CivitAI modelId (int) or HF repo (str).
Preference order:
@@ -741,7 +752,7 @@ class BaseModelService(ABC):
return BaseModelService._extract_hf_group_key(item)
@staticmethod
def _extract_model_id(item: Dict) -> Optional[int]:
def _extract_model_id(item: Dict[str, Any]) -> Optional[int]:
civitai = item.get("civitai") if isinstance(item, dict) else None
if not isinstance(civitai, dict):
return None
@@ -754,7 +765,7 @@ class BaseModelService(ABC):
return None
@staticmethod
def _extract_version_id(item: Dict) -> Optional[int]:
def _extract_version_id(item: Dict[str, Any]) -> Optional[int]:
civitai = item.get("civitai") if isinstance(item, dict) else None
if not isinstance(civitai, dict):
return None
@@ -767,7 +778,7 @@ class BaseModelService(ABC):
return None
@staticmethod
def _extract_base_model(item: Dict) -> Optional[str]:
def _extract_base_model(item: Dict[str, Any]) -> Optional[str]:
value = item.get("base_model")
if value is None:
return None
@@ -819,7 +830,7 @@ class BaseModelService(ABC):
return highest_by_base
def _paginate(self, data: List[Dict], page: int, page_size: int) -> Dict:
def _paginate(self, data: List[Dict[str, Any]], page: int, page_size: int) -> Dict[str, Any]:
"""Apply pagination to filtered data"""
total_items = len(data)
start_idx = (page - 1) * page_size
@@ -834,7 +845,7 @@ class BaseModelService(ABC):
}
@abstractmethod
async def format_response(self, model_data: Dict) -> Optional[Dict]:
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Format model data for API response - must be implemented by subclasses.
Subclasses should return None for corrupted entries so the handler
@@ -843,17 +854,17 @@ class BaseModelService(ABC):
pass
# Common service methods that delegate to scanner
async def get_top_tags(self, limit: int = 20) -> List[Dict]:
async def get_top_tags(self, limit: int = 20) -> List[Dict[str, Any]]:
"""Get top tags sorted by frequency"""
return await self.scanner.get_top_tags(limit)
async def search_tags(
self, query: str, limit: int = 50
) -> List[Dict]:
) -> List[Dict[str, Any]]:
"""Search tags by substring, sorted by frequency"""
return await self.scanner.search_tags(query, limit)
async def get_base_models(self, limit: int = 20) -> List[Dict]:
async def get_base_models(self, limit: int = 20) -> List[Dict[str, Any]]:
"""Get base models sorted by frequency"""
return await self.scanner.get_base_models(limit)
@@ -920,7 +931,7 @@ class BaseModelService(ABC):
"""Get model root directories"""
return self.scanner.get_model_roots()
def filter_civitai_data(self, data: Dict, minimal: bool = False) -> Dict:
def filter_civitai_data(self, data: Dict[str, Any], minimal: bool = False) -> Dict[str, Any]:
"""Filter relevant fields from CivitAI data"""
if not data:
return {}
@@ -946,7 +957,7 @@ class BaseModelService(ABC):
)
return {k: data[k] for k in fields if k in data}
async def get_folder_tree(self, model_root: str) -> Dict:
async def get_folder_tree(self, model_root: str) -> Dict[str, Any]:
"""Get hierarchical folder tree for a specific model root"""
cache = await self.scanner.get_cached_data()
@@ -975,7 +986,7 @@ class BaseModelService(ABC):
return tree
async def get_unified_folder_tree(self) -> Dict:
async def get_unified_folder_tree(self) -> Dict[str, Any]:
"""Get unified folder tree across all model roots"""
cache = await self.scanner.get_cached_data()
@@ -1004,7 +1015,7 @@ class BaseModelService(ABC):
return unified_tree
async def get_model_notes(self, model_name: str) -> Optional[dict]:
async def get_model_notes(self, model_name: str) -> Optional[dict[str, Any]]:
"""Get notes and file_path for a specific model file.
Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path
@@ -1136,7 +1147,7 @@ class BaseModelService(ABC):
return {"civitai_url": None, "model_id": None, "version_id": None}
async def get_model_metadata(self, file_path: str) -> Optional[Dict]:
async def get_model_metadata(self, file_path: str) -> Optional[Dict[str, Any]]:
"""Load full metadata for a single model.
Listing/search endpoints return lightweight cache entries; this method performs
@@ -1232,7 +1243,7 @@ class BaseModelService(ABC):
return True
@staticmethod
def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple:
def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple[int, int, int, str]:
"""Sort paths by how well they satisfy the include tokens.
Sorts based on path without extension for consistent ordering.
@@ -1259,19 +1270,87 @@ class BaseModelService(ABC):
)
async def search_relative_paths(
self, search_term: str, limit: int = 15, offset: int = 0
self,
search_term: str,
limit: int = 15,
offset: int = 0,
*,
folder: Optional[str] = None,
folder_include: Optional[list[str]] = None,
folder_exclude: Optional[list[str]] = None,
base_models: Optional[list[str]] = None,
model_types: Optional[list[str]] = None,
tags: Optional[dict[str, str]] = None,
auto_tags: Optional[dict[str, str]] = None,
tag_logic: str = "any",
credit_required: Optional[bool] = None,
allow_selling_generated_content: Optional[bool] = None,
recursive: bool = True,
apply_filters: bool = False,
) -> List[str]:
"""Search model relative file paths for autocomplete functionality"""
"""Search model relative file paths for autocomplete functionality.
Optional filter kwargs mirror the filters used by the list endpoint
(/api/lm/{prefix}/list). When no filter kwargs are provided the
behavior is identical to plain token-based path matching.
"""
cache = await self.scanner.get_cached_data()
include_terms, exclude_terms = self._parse_search_tokens(search_term)
data = cache.raw_data
has_filters = any(
[
apply_filters,
folder is not None,
folder_include,
folder_exclude,
base_models,
model_types,
tags,
auto_tags,
credit_required is not None,
allow_selling_generated_content is not None,
]
)
if has_filters:
# Auto-tags are not stored in the scanner cache — they are computed
# on the fly. Pre-compute them only when an auto-tag filter is
# active to avoid mutating cache entries unnecessarily.
if auto_tags:
from .auto_tag_service import extract_auto_tags
for item in data:
if not item.get("auto_tags"):
item["auto_tags"] = extract_auto_tags(item)
criteria = FilterCriteria(
folder=folder,
folder_include=folder_include,
folder_exclude=folder_exclude,
base_models=base_models,
model_types=model_types,
tags=tags,
auto_tags=auto_tags,
search_options={"recursive": recursive},
tag_logic=tag_logic,
)
data = self.filter_set.apply(data, criteria)
if credit_required is not None:
data = await self._apply_credit_required_filter(
data, credit_required
)
if allow_selling_generated_content is not None:
data = await self._apply_allow_selling_filter(
data, allow_selling_generated_content
)
matching_paths = []
# Get model roots for path calculation
model_roots = self.scanner.get_model_roots()
# Collect all matching paths first (needed for proper sorting and offset)
for model in cache.raw_data:
for model in data:
file_path = model.get("file_path", "")
if not file_path:
continue
+30 -2
View File
@@ -59,6 +59,7 @@ class CacheEntryValidator:
'notes': ('', False),
'usage_tips': ('', False),
'hash_status': ('completed', False),
'autov3': (None, False),
}
@classmethod
@@ -119,8 +120,13 @@ class CacheEntryValidator:
if is_required:
errors.append(f"Required field '{field_name}' is missing or None")
if auto_repair:
working_entry[field_name] = cls._get_default_copy(default_value)
repaired = True
# A missing optional field whose default is None is already
# semantically equal to its default (e.g. autov3: absent
# means "not checked") — writing None back is a no-op, not
# a repair.
if default_value is not None:
working_entry[field_name] = cls._get_default_copy(default_value)
repaired = True
continue
# Validate field type and value
@@ -175,6 +181,15 @@ class CacheEntryValidator:
# that invalidates the entry, but we also don't mark it repaired.
pass
# Normalize autov3 to lowercase if needed (optional field, never stripped).
autov3 = working_entry.get('autov3')
if isinstance(autov3, str) and autov3:
normalized_autov3 = autov3.lower()
if normalized_autov3 != autov3:
if auto_repair:
working_entry['autov3'] = normalized_autov3
repaired = True
# Determine if entry is valid
# Entry is valid if no critical required field errors remain after repair
# Critical fields are file_path and sha256
@@ -242,6 +257,19 @@ class CacheEntryValidator:
"""
expected_type = type(default_value)
# Special case: autov3 is optional with a three-state contract.
# None = not checked, "" = checked but unavailable, otherwise a
# 12-character hex string (case-insensitive here; normalized to
# lowercase separately).
if field_name == 'autov3':
if value is None or value == "":
return None
if not isinstance(value, str):
return f"Field 'autov3' should be string or None, got {type(value).__name__}"
if len(value) != 12 or any(c not in '0123456789abcdefABCDEF' for c in value):
return "Field 'autov3' should be a 12-character hex string"
return None
# Special handling for numeric types
if expected_type == int:
if not isinstance(value, (int, float)):
+33 -6
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import asyncio
import json
import logging
@@ -6,7 +10,7 @@ from datetime import datetime
from typing import Any, Dict, List, Optional
from ..utils.models import CheckpointMetadata
from ..utils.file_utils import find_preview_file, normalize_path
from ..utils.file_utils import find_preview_file, normalize_path, calculate_autov3
from ..utils.metadata_manager import MetadataManager
from ..config import config
from .model_scanner import ModelScanner
@@ -62,6 +66,11 @@ class CheckpointScanner(ModelScanner):
# Find preview image
preview_url = find_preview_file(base_name, dir_path)
# AutoV3 reads only the safetensors header, so it is cheap even for
# large checkpoints; record the checked state at creation time ("" =
# checked but unavailable).
autov3 = calculate_autov3(real_path)
# Create metadata WITHOUT calculating hash
metadata = CheckpointMetadata(
file_name=base_name,
@@ -77,6 +86,7 @@ class CheckpointScanner(ModelScanner):
sub_type="checkpoint",
from_civitai=False, # Mark as local model since no hash yet
hash_status="pending", # Mark hash as pending
autov3=autov3 or "",
)
# Save the created metadata
@@ -120,7 +130,11 @@ class CheckpointScanner(ModelScanner):
# that queries get_hash_by_filename first) will miss on every
# lookup and keep calling back into this method, creating a
# tight loop that never populates the index.
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
async with self._hash_calculation_lock:
@@ -132,7 +146,11 @@ class CheckpointScanner(ModelScanner):
and metadata.hash_status == "completed"
and metadata.sha256
):
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
task = self._hash_calculation_tasks.get(real_path)
@@ -185,7 +203,11 @@ class CheckpointScanner(ModelScanner):
if metadata.hash_status == "completed" and metadata.sha256:
# Populate the in-memory hash index even for pre-computed
# hashes, mirroring the fix in calculate_hash_for_model.
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
# Update status to calculating
@@ -202,7 +224,11 @@ class CheckpointScanner(ModelScanner):
await MetadataManager.save_metadata(file_path, metadata)
# Update hash index
self._hash_index.add_entry(sha256.lower(), file_path)
self._hash_index.add_entry(
sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
# Update the in-memory cache entry so that subsequent
# _persist_current_cache / _save_persistent_cache calls
@@ -216,6 +242,7 @@ class CheckpointScanner(ModelScanner):
if entry.get("file_path") == file_path:
entry["sha256"] = sha256.lower()
entry["hash_status"] = "completed"
self.bump_cache_version()
break
logger.info(f"Hash calculated for checkpoint: {file_path}")
@@ -405,7 +432,7 @@ class CheckpointScanner(ModelScanner):
roots.extend(config.extra_checkpoints_roots or [])
roots.extend(config.extra_unet_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root not in seen:
+28 -28
View File
@@ -1,6 +1,6 @@
import os
import logging
from typing import Dict, Optional
from typing import Any, Dict, Optional
from .base_model_service import BaseModelService
from .auto_tag_service import extract_auto_tags
@@ -21,58 +21,58 @@ class CheckpointService(BaseModelService):
"""
super().__init__("checkpoint", scanner, CheckpointMetadata, update_service=update_service)
async def format_response(self, checkpoint_data: Dict) -> Optional[Dict]:
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Format Checkpoint data for API response.
Returns None when the entry is missing critical fields (corrupted cache
row), so the handler layer can filter it out. See issue #730.
"""
# Guard against corrupted cache entries missing critical fields
file_path = checkpoint_data.get("file_path")
file_path = model_data.get("file_path")
if not file_path or not isinstance(file_path, str):
logger.warning(
"Skipping corrupted checkpoint entry (missing file_path): %s",
checkpoint_data.get("file_name", "<unknown>"),
model_data.get("file_name", "<unknown>"),
)
return None
# Get sub_type from cache entry (new canonical field)
sub_type = checkpoint_data.get("sub_type", "checkpoint")
sub_type = model_data.get("sub_type", "checkpoint")
file_name = checkpoint_data.get("file_name") or ""
model_name = checkpoint_data.get("model_name") or file_name
folder = checkpoint_data.get("folder") or ""
file_name = model_data.get("file_name") or ""
model_name = model_data.get("model_name") or file_name
folder = model_data.get("folder") or ""
return {
"model_name": model_name,
"file_name": file_name,
"preview_url": config.get_preview_static_url(checkpoint_data.get("preview_url", "")),
"preview_nsfw_level": checkpoint_data.get("preview_nsfw_level", 0),
"base_model": checkpoint_data.get("base_model", ""),
"preview_url": config.get_preview_static_url(model_data.get("preview_url", "")),
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
"base_model": model_data.get("base_model", ""),
"folder": folder,
"sha256": checkpoint_data.get("sha256", ""),
"sha256": model_data.get("sha256", ""),
"file_path": file_path.replace(os.sep, "/"),
"file_size": checkpoint_data.get("size", 0),
"modified": checkpoint_data.get("modified", ""),
"tags": checkpoint_data.get("tags", []),
"from_civitai": checkpoint_data.get("from_civitai", True),
"usage_count": checkpoint_data.get("usage_count", 0),
"notes": checkpoint_data.get("notes", ""),
"file_size": model_data.get("size", 0),
"modified": model_data.get("modified", ""),
"tags": model_data.get("tags", []),
"from_civitai": model_data.get("from_civitai", True),
"usage_count": model_data.get("usage_count", 0),
"notes": model_data.get("notes", ""),
"sub_type": sub_type,
"favorite": checkpoint_data.get("favorite", False),
"exclude": bool(checkpoint_data.get("exclude", False)),
"update_available": bool(checkpoint_data.get("update_available", False)),
"skip_metadata_refresh": bool(checkpoint_data.get("skip_metadata_refresh", False)),
"civitai": self.filter_civitai_data(checkpoint_data.get("civitai", {}), minimal=True),
"auto_tags": checkpoint_data.get("auto_tags") or extract_auto_tags(checkpoint_data),
"version_count": checkpoint_data.get("version_count"),
"hf_url": checkpoint_data.get("hf_url", ""),
"favorite": model_data.get("favorite", False),
"exclude": bool(model_data.get("exclude", False)),
"update_available": bool(model_data.get("update_available", False)),
"skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)),
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
"version_count": model_data.get("version_count"),
"hf_url": model_data.get("hf_url", ""),
}
def find_duplicate_hashes(self) -> Dict:
def find_duplicate_hashes(self) -> Dict[str, Any]:
"""Find Checkpoints with duplicate SHA256 hashes"""
return self.scanner._hash_index.get_duplicate_hashes()
def find_duplicate_filenames(self) -> Dict:
def find_duplicate_filenames(self) -> Dict[str, Any]:
"""Find Checkpoints with conflicting filenames"""
return self.scanner._hash_index.get_duplicate_filenames()
+41 -36
View File
@@ -1,8 +1,12 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import json
import logging
import asyncio
from copy import deepcopy
from typing import Optional, Dict, Tuple, List
from typing import Any, Optional, Dict, Tuple, List, cast
from .model_metadata_provider import CivArchiveModelMetadataProvider, ModelMetadataProviderManager
from .downloader import get_downloader
from .errors import RateLimitError
@@ -37,8 +41,8 @@ class CivArchiveClient:
async def _request_json(
self,
path: str,
params: Optional[Dict[str, str]] = None
) -> Tuple[Optional[Dict], Optional[str]]:
params: Optional[Dict[str, Any]] = None
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Call CivArchive API and return JSON payload"""
success, payload = await self._make_request(path, params=params)
if not success:
@@ -52,12 +56,12 @@ class CivArchiveClient:
self,
path: str,
*,
params: Optional[Dict[str, str]] = None,
) -> Tuple[bool, Dict | str]:
params: Optional[Dict[str, Any]] = None,
) -> Tuple[bool, Dict[str, Any] | str]:
"""Wrapper around downloader.make_request that surfaces rate limits."""
downloader = await get_downloader()
kwargs: Dict[str, Dict[str, str]] = {}
kwargs: Dict[str, Dict[str, Any]] = {}
if params:
safe_params = {str(key): str(value) for key, value in params.items() if value is not None}
if safe_params:
@@ -73,10 +77,11 @@ class CivArchiveClient:
if payload.provider is None:
payload.provider = "civarchive_api"
raise payload
return success, payload
# RateLimitError is always raised above, so the returned payload is a dict or str.
return success, cast(Dict[str, Any] | str, payload)
@staticmethod
def _normalize_payload(payload: Dict) -> Dict:
def _normalize_payload(payload: Dict[str, Any]) -> Dict[str, Any]:
"""Unwrap CivArchive responses that wrap content under a data key"""
if not isinstance(payload, dict):
return {}
@@ -86,12 +91,12 @@ class CivArchiveClient:
return payload
@staticmethod
def _split_context(payload: Dict) -> Tuple[Dict, Dict, List[Dict]]:
def _split_context(payload: Dict[str, Any]) -> Tuple[Dict[str, Any], Dict[str, Any], List[Dict[str, Any]]]:
"""Separate version payload from surrounding model context"""
data = CivArchiveClient._normalize_payload(payload)
context: Dict = {}
fallback_files: List[Dict] = []
version: Dict = {}
context: Dict[str, Any] = {}
fallback_files: List[Dict[str, Any]] = []
version: Dict[str, Any] = {}
for key, value in data.items():
if key in {"version", "model"}:
@@ -115,7 +120,7 @@ class CivArchiveClient:
return context, version, fallback_files
@staticmethod
def _ensure_list(value) -> List:
def _ensure_list(value: Any) -> List[Any]:
if isinstance(value, list):
return value
if value is None:
@@ -123,7 +128,7 @@ class CivArchiveClient:
return [value]
@staticmethod
def _build_model_info(context: Dict) -> Dict:
def _build_model_info(context: Dict[str, Any]) -> Dict[str, Any]:
tags = context.get("tags")
if not isinstance(tags, list):
tags = list(tags) if isinstance(tags, (set, tuple)) else ([] if tags is None else [tags])
@@ -136,7 +141,7 @@ class CivArchiveClient:
}
@staticmethod
def _build_creator_info(context: Dict) -> Dict:
def _build_creator_info(context: Dict[str, Any]) -> Dict[str, Any]:
username = context.get("creator_username") or context.get("username") or ""
image = context.get("creator_image") or context.get("creator_avatar") or ""
creator: Dict[str, Optional[str]] = {
@@ -150,7 +155,7 @@ class CivArchiveClient:
return creator
@staticmethod
def _transform_file_entry(file_data: Dict) -> Dict:
def _transform_file_entry(file_data: Dict[str, Any]) -> Dict[str, Any]:
mirrors = file_data.get("mirrors") or []
if not isinstance(mirrors, list):
mirrors = [mirrors]
@@ -165,7 +170,7 @@ class CivArchiveClient:
if not name and available_mirror:
name = available_mirror.get("filename")
transformed: Dict = {
transformed: Dict[str, Any] = {
"id": file_data.get("id"),
"sizeKB": file_data.get("sizeKB"),
"name": name,
@@ -216,23 +221,23 @@ class CivArchiveClient:
def _transform_files(
self,
files: Optional[List[Dict]],
fallback_files: Optional[List[Dict]] = None
) -> List[Dict]:
candidates: List[Dict] = []
files: Optional[List[Dict[str, Any]]],
fallback_files: Optional[List[Dict[str, Any]]] = None
) -> List[Dict[str, Any]]:
candidates: List[Dict[str, Any]] = []
if isinstance(files, list) and files:
candidates = files
elif isinstance(fallback_files, list):
candidates = fallback_files
transformed_files: List[Dict] = []
transformed_files: List[Dict[str, Any]] = []
for file_data in candidates:
if isinstance(file_data, dict):
transformed_files.append(self._transform_file_entry(file_data))
# Sort: .safetensors first, .ckpt second, others last
# so the backend fallback (no file_params) prefers safetensors
def _sort_key(f: Dict) -> int:
def _sort_key(f: Dict[str, Any]) -> int:
fname = f.get("name") or ""
if isinstance(fname, str):
lower = fname.lower()
@@ -247,10 +252,10 @@ class CivArchiveClient:
def _transform_version(
self,
context: Dict,
version: Dict,
fallback_files: Optional[List[Dict]] = None
) -> Optional[Dict]:
context: Dict[str, Any],
version: Dict[str, Any],
fallback_files: Optional[List[Dict[str, Any]]] = None
) -> Optional[Dict[str, Any]]:
if not version:
return None
@@ -291,7 +296,7 @@ class CivArchiveClient:
return version_copy
async def _resolve_version_from_files(self, payload: Dict) -> Optional[Dict]:
async def _resolve_version_from_files(self, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Fallback to fetch version data when only file metadata is available"""
data = self._normalize_payload(payload)
files = data.get("files") or payload.get("files") or []
@@ -323,7 +328,7 @@ class CivArchiveClient:
return resolved
return None
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Find model by SHA256 hash value using CivArchive API"""
try:
payload, error = await self._request_json(f"/sha256/{model_hash.lower()}")
@@ -332,12 +337,12 @@ class CivArchiveClient:
return None, "Model not found"
return None, error
context, version_data, fallback_files = self._split_context(payload)
context, version_data, fallback_files = self._split_context(cast(Dict[str, Any], payload))
transformed = self._transform_version(context, version_data, fallback_files)
if transformed:
return transformed, None
resolved = await self._resolve_version_from_files(payload)
resolved = await self._resolve_version_from_files(cast(Dict[str, Any], payload))
if resolved:
return resolved, None
@@ -350,7 +355,7 @@ class CivArchiveClient:
logger.error(f"Error fetching CivArchive model by hash {model_hash[:10]}: {e}")
return None, str(e)
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
"""Get all versions of a model using CivArchive API"""
try:
payload, error = await self._request_json(f"/models/{model_id}")
@@ -364,7 +369,7 @@ class CivArchiveClient:
context, version_data, fallback_files = self._split_context(payload)
versions_meta = data.get("versions") or []
transformed_versions: List[Dict] = []
transformed_versions: List[Dict[str, Any]] = []
for meta in versions_meta:
if not isinstance(meta, dict):
continue
@@ -381,7 +386,7 @@ class CivArchiveClient:
if primary_version:
transformed_versions.insert(0, primary_version)
ordered_versions: List[Dict] = []
ordered_versions: List[Dict[str, Any]] = []
seen_ids = set()
for version in transformed_versions:
version_id = version.get("id")
@@ -402,7 +407,7 @@ class CivArchiveClient:
logger.error(f"Error fetching CivArchive model versions for {model_id}: {e}")
return None
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: int | str | None = None, version_id: int | str | None = None) -> Optional[Dict[str, Any]]:
"""Get specific model version using CivArchive API
Args:
@@ -459,7 +464,7 @@ class CivArchiveClient:
logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {e}")
return None
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[str, Any]], Optional[str]]:
""" Fetch model version metadata using a known bogus model lookup
CivArchive lacks a direct version lookup API, this uses a workaround (which we handle in the main model request now)
+1 -1
View File
@@ -283,7 +283,7 @@ class CivitaiBaseModelService:
return None
if isinstance(result, str):
data = json.loads(result)
data: Any = json.loads(result)
else:
data = result
+40 -35
View File
@@ -1,10 +1,14 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import asyncio
import copy
import logging
import os
import time
from collections import OrderedDict
from typing import Any, Optional, Dict, Tuple, List, Sequence
from typing import Any, Optional, Dict, Tuple, List, Sequence, cast
from .connectivity_guard import (
OFFLINE_FRIENDLY_MESSAGE,
is_expected_offline_error,
@@ -58,7 +62,7 @@ class CivitaiClient:
# Uses OrderedDict with LRU eviction at MAX_CACHE_ENTRIES to prevent
# unbounded growth in long-running server processes.
self._version_info_cache: OrderedDict[
str, Tuple[Optional[Dict], Optional[str]]
str, Tuple[Optional[Dict[str, Any]], Optional[str]]
] = OrderedDict()
self._MAX_CACHE_ENTRIES = 500
@@ -72,7 +76,7 @@ class CivitaiClient:
*,
use_auth: bool = False,
**kwargs,
) -> Tuple[bool, Dict | str]:
) -> Tuple[bool, Dict[str, Any] | str]:
"""Wrapper around downloader.make_request that surfaces rate limits,
with retry for transient server errors (5xx, Cloudflare 524, network flakiness)."""
@@ -86,7 +90,8 @@ class CivitaiClient:
**kwargs,
)
if success:
return True, result
# RateLimitError is raised below; a successful result is dict or str.
return True, cast(Dict[str, Any] | str, result)
if isinstance(result, RateLimitError):
if result.provider is None:
@@ -126,7 +131,7 @@ class CivitaiClient:
return False, "Unexpected error in _make_request"
@staticmethod
def _remove_comfy_metadata(model_version: Optional[Dict]) -> None:
def _remove_comfy_metadata(model_version: Optional[Dict[str, Any]]) -> None:
"""Remove Comfy-specific metadata from model version images."""
if not isinstance(model_version, dict):
return
@@ -173,7 +178,7 @@ class CivitaiClient:
async def get_model_by_hash(
self, model_hash: str
) -> Tuple[Optional[Dict], Optional[str]]:
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
try:
success, version = await self._make_request(
"GET",
@@ -220,7 +225,7 @@ class CivitaiClient:
# Ensure directory exists
os.makedirs(os.path.dirname(save_path), exist_ok=True)
with open(save_path, "wb") as f:
f.write(content)
f.write(content if isinstance(content, bytes) else content.encode("utf-8"))
return True
return False
except Exception as e:
@@ -275,7 +280,7 @@ class CivitaiClient:
return True
return False
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
"""Get all versions of a model with local availability info"""
try:
success, result = await self._make_request(
@@ -283,7 +288,7 @@ class CivitaiClient:
f"{self.base_url}/models/{model_id}",
use_auth=True,
)
if success:
if success and isinstance(result, dict):
# Also return model type along with versions
return {
"modelVersions": result.get("modelVersions", []),
@@ -317,7 +322,7 @@ class CivitaiClient:
async def get_model_versions_bulk(
self, model_ids: Sequence[int]
) -> Optional[Dict[int, Dict]]:
) -> Optional[Dict[int, Dict[str, Any]]]:
"""Fetch model metadata for multiple ids using the batch API."""
deduped: Dict[int, None] = {}
@@ -347,13 +352,13 @@ class CivitaiClient:
if not isinstance(items, list):
return {}
payload: Dict[int, Dict] = {}
payload: Dict[int, Dict[str, Any]] = {}
for item in items:
if not isinstance(item, dict):
continue
model_id = item.get("id")
try:
normalized_id = int(model_id)
normalized_id = int(cast(Any, model_id))
except (TypeError, ValueError):
continue
payload[normalized_id] = {
@@ -373,8 +378,8 @@ class CivitaiClient:
return None
async def get_model_version(
self, model_id: int = None, version_id: int = None
) -> Optional[Dict]:
self, model_id: int | None = None, version_id: int | None = None
) -> Optional[Dict[str, Any]]:
"""Get specific model version with additional metadata."""
try:
if model_id is None and version_id is not None:
@@ -392,7 +397,7 @@ class CivitaiClient:
logger.error(f"Error fetching model version: {e}")
return None
async def _get_version_by_id_only(self, version_id: int) -> Optional[Dict]:
async def _get_version_by_id_only(self, version_id: int) -> Optional[Dict[str, Any]]:
version = await self._fetch_version_by_id(version_id)
if version is None:
return None
@@ -411,7 +416,7 @@ class CivitaiClient:
async def _get_version_with_model_id(
self, model_id: int, version_id: Optional[int]
) -> Optional[Dict]:
) -> Optional[Dict[str, Any]]:
model_data = await self._fetch_model_data(model_id)
if not model_data:
return None
@@ -464,20 +469,20 @@ class CivitaiClient:
self._remove_comfy_metadata(version)
return version
async def _fetch_model_data(self, model_id: int) -> Optional[Dict]:
async def _fetch_model_data(self, model_id: int) -> Optional[Dict[str, Any]]:
success, data = await self._make_request(
"GET",
f"{self.base_url}/models/{model_id}",
use_auth=True,
)
if success:
if success and isinstance(data, dict):
return data
if is_expected_offline_error(data):
return None
logger.warning(f"Failed to fetch model data for model {model_id}")
return None
async def _fetch_version_by_id(self, version_id: Optional[int]) -> Optional[Dict]:
async def _fetch_version_by_id(self, version_id: Optional[int]) -> Optional[Dict[str, Any]]:
if version_id is None:
return None
@@ -486,7 +491,7 @@ class CivitaiClient:
f"{self.base_url}/model-versions/{version_id}",
use_auth=True,
)
if success:
if success and isinstance(version, dict):
return version
if is_expected_offline_error(version):
return None
@@ -494,7 +499,7 @@ class CivitaiClient:
logger.warning(f"Failed to fetch version by id {version_id}")
return None
async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict]:
async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict[str, Any]]:
if not model_hash:
return None
@@ -503,7 +508,7 @@ class CivitaiClient:
f"{self.base_url}/model-versions/by-hash/{model_hash}",
use_auth=True,
)
if success:
if success and isinstance(version, dict):
return version
if is_expected_offline_error(version):
return None
@@ -512,8 +517,8 @@ class CivitaiClient:
return None
def _select_target_version(
self, model_data: Dict, model_id: int, version_id: Optional[int]
) -> Optional[Dict]:
self, model_data: Dict[str, Any], model_id: int, version_id: Optional[int]
) -> Optional[Dict[str, Any]]:
model_versions = model_data.get("modelVersions", [])
if not model_versions:
logger.warning(f"No model versions found for model {model_id}")
@@ -532,7 +537,7 @@ class CivitaiClient:
return model_versions[0]
def _extract_primary_model_hash(self, version_entry: Dict) -> Optional[str]:
def _extract_primary_model_hash(self, version_entry: Dict[str, Any]) -> Optional[str]:
for file_info in version_entry.get("files", []):
if file_info.get("type") == "Model" and file_info.get("primary"):
hashes = file_info.get("hashes", {})
@@ -542,8 +547,8 @@ class CivitaiClient:
return None
def _build_version_from_model_data(
self, version_entry: Dict, model_id: int, model_data: Dict
) -> Dict:
self, version_entry: Dict[str, Any], model_id: int, model_data: Dict[str, Any]
) -> Dict[str, Any]:
version = copy.deepcopy(version_entry)
version.pop("index", None)
version["modelId"] = model_id
@@ -555,7 +560,7 @@ class CivitaiClient:
}
return version
def _enrich_version_with_model_data(self, version: Dict, model_data: Dict) -> None:
def _enrich_version_with_model_data(self, version: Dict[str, Any], model_data: Dict[str, Any]) -> None:
model_info = version.get("model")
if not isinstance(model_info, dict):
model_info = {}
@@ -571,7 +576,7 @@ class CivitaiClient:
async def get_model_version_info(
self, version_id: str
) -> Tuple[Optional[Dict], Optional[str]]:
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Fetch model version metadata from Civitai
Args:
@@ -596,7 +601,7 @@ class CivitaiClient:
logger.debug("Resolving Civitai model version info: %s", url)
success, result = await self._make_request("GET", url, use_auth=True)
if success:
if success and isinstance(result, dict):
logger.debug("Successfully fetched model version info for: %s", version_id)
self._remove_comfy_metadata(result)
self._version_info_cache[version_id] = (result, None)
@@ -626,7 +631,7 @@ class CivitaiClient:
async def get_image_info(
self, image_id: str, source_url: str | None = None
) -> Optional[Dict]:
) -> Optional[Dict[str, Any]]:
"""Fetch image information from Civitai API
Args:
@@ -659,7 +664,7 @@ class CivitaiClient:
)
return None
if result and "items" in result and isinstance(result["items"], list):
if isinstance(result, dict) and "items" in result and isinstance(result["items"], list):
items = result["items"]
for item in items:
@@ -699,7 +704,7 @@ class CivitaiClient:
async def get_model_versions_by_hashes(
self, hashes: List[str]
) -> Optional[List[Dict]]:
) -> Optional[List[Dict[str, Any]]]:
"""Fetch full version details for up to 100 SHA256 hashes via the batch endpoint.
Uses POST /api/v1/model-versions/by-hash which returns full version
@@ -716,7 +721,7 @@ class CivitaiClient:
return []
BATCH_SIZE = 100
all_versions: List[Dict] = []
all_versions: List[Dict[str, Any]] = []
for start in range(0, len(hashes), BATCH_SIZE):
batch = hashes[start : start + BATCH_SIZE]
@@ -736,7 +741,7 @@ class CivitaiClient:
continue
if isinstance(result, list):
all_versions.extend(result)
all_versions.extend(cast(Any, result))
else:
logger.debug(
"Unexpected by-hash response type: %s", type(result)
+1 -1
View File
@@ -18,7 +18,7 @@ class DownloadCoordinator:
self,
*,
ws_manager,
download_manager_factory: Callable[[], Awaitable],
download_manager_factory: Callable[[], Awaitable[Any]],
) -> None:
self._ws_manager = ws_manager
self._download_manager_factory = download_manager_factory
+113 -76
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import copy
import logging
import os
@@ -8,7 +12,7 @@ import zipfile
from concurrent.futures import ThreadPoolExecutor
from collections import OrderedDict
import uuid
from typing import Dict, List, Optional, Set, Tuple
from typing import Any, Dict, List, Optional, Set, Tuple, cast
from urllib.parse import urlparse
from ..utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata
from ..utils.constants import (
@@ -18,7 +22,7 @@ from ..utils.constants import (
VALID_LORA_TYPES,
)
from ..utils.civitai_utils import normalize_civitai_download_url, rewrite_preview_url
from ..utils.file_utils import calculate_sha256
from ..utils.file_utils import calculate_sha256, calculate_autov3
from ..utils.preview_selection import resolve_mature_threshold, select_preview_media
from ..utils.utils import sanitize_folder_name
from ..utils.exif_utils import ExifUtils
@@ -121,7 +125,7 @@ class DownloadManager:
"delay": 0,
}
)
except DownloadInProgressError:
except DownloadInProgressError: # pyright: ignore[reportPossiblyUnboundVariable]
logger.info(
"Skipping automatic example images download for %s; another example images download is already running",
model_hash,
@@ -170,7 +174,7 @@ class DownloadManager:
logger.error("aria2 download failed for %s: %s", download_url, exc)
return False, str(exc)
download_kwargs = {
download_kwargs: Dict[str, Any] = {
"progress_callback": progress_callback,
"use_auth": use_auth,
}
@@ -204,16 +208,16 @@ class DownloadManager:
async def download_from_civitai(
self,
model_id: int = None,
model_version_id: int = None,
save_dir: str = None,
model_id: int | None = None,
model_version_id: int | None = None,
save_dir: str | None = None,
relative_path: str = "",
progress_callback=None,
use_default_paths: bool = False,
download_id: str = None,
source: str = None,
file_params: Dict = None,
) -> Dict:
download_id: str | None = None,
source: str | None = None,
file_params: Dict[str, Any] | None = None,
) -> Dict[str, Any]:
"""Download model from Civitai with task tracking and concurrency control
Args:
@@ -309,14 +313,14 @@ class DownloadManager:
async def _download_with_semaphore(
self,
task_id: str,
model_id: int,
model_version_id: int,
save_dir: str,
model_id: int | None,
model_version_id: int | None,
save_dir: str | None,
relative_path: str,
progress_callback=None,
use_default_paths: bool = False,
source: str = None,
file_params: Dict = None,
source: str | None = None,
file_params: Dict[str, Any] | None = None,
):
"""Execute download with semaphore to limit concurrency"""
# Update status to waiting
@@ -380,7 +384,8 @@ class DownloadManager:
# Use original download implementation
try:
# Check for cancellation before starting
if asyncio.current_task().cancelled():
current_task = asyncio.current_task()
if current_task is not None and current_task.cancelled():
raise asyncio.CancelledError()
result = await self._execute_original_download(
@@ -484,11 +489,11 @@ class DownloadManager:
# Schedule cleanup of download record after delay
asyncio.create_task(self._cleanup_download_record(task_id))
def _start_background_download_task(self, download_id: str, coroutine) -> asyncio.Task:
def _start_background_download_task(self, download_id: str, coroutine) -> asyncio.Task[Any]:
task = asyncio.create_task(coroutine)
self._download_tasks[download_id] = task
def _cleanup_done_task(done_task: asyncio.Task) -> None:
def _cleanup_done_task(done_task: asyncio.Task[Any]) -> None:
current_task = self._download_tasks.get(download_id)
if current_task is done_task:
self._download_tasks.pop(download_id, None)
@@ -530,7 +535,7 @@ class DownloadManager:
async def _cleanup_cancelled_download_files(
self,
download_id: str,
download_info: Optional[Dict],
download_info: Optional[Dict[str, Any]],
) -> None:
target_files = set()
persisted = await self._aria2_state_store.get(download_id)
@@ -603,13 +608,13 @@ class DownloadManager:
self,
download_id: str,
*,
extra: Optional[Dict] = None,
extra: Optional[Dict[str, Any]] = None,
) -> None:
info = self._active_downloads.get(download_id)
if not info:
return
payload = {
payload: Dict[str, Any] = {
"download_id": download_id,
"model_id": info.get("model_id"),
"model_version_id": info.get("model_version_id"),
@@ -631,7 +636,7 @@ class DownloadManager:
await self._aria2_state_store.upsert(download_id, payload)
def _build_restored_download_info(self, record: Dict, save_path: str) -> Dict:
def _build_restored_download_info(self, record: Dict[str, Any], save_path: str) -> Dict[str, Any]:
return {
"model_id": record.get("model_id"),
"model_version_id": record.get("model_version_id"),
@@ -653,8 +658,8 @@ class DownloadManager:
def _is_same_aria2_download_request(
self,
current_info: Optional[Dict],
persisted_record: Dict,
current_info: Optional[Dict[str, Any]],
persisted_record: Dict[str, Any],
) -> bool:
if not isinstance(current_info, dict):
return False
@@ -666,13 +671,15 @@ class DownloadManager:
return current_version_id == persisted_version_id
def _build_download_urls_from_file_info(self, file_info: Dict, source: str = None) -> List[str]:
def _build_download_urls_from_file_info(self, file_info: Dict[str, Any], source: str | None = None) -> List[str]:
mirrors = file_info.get("mirrors") or []
download_urls: List[str] = []
if mirrors:
for mirror in mirrors:
if mirror.get("deletedAt") is None and mirror.get("url"):
download_urls.append(normalize_civitai_download_url(mirror["url"]))
normalized_url = normalize_civitai_download_url(mirror["url"])
if normalized_url:
download_urls.append(normalized_url)
if source == "civarchive" and len(download_urls) > 1:
civitai_urls = [
@@ -688,7 +695,9 @@ class DownloadManager:
if not download_urls:
download_url = file_info.get("downloadUrl")
if download_url:
download_urls.append(normalize_civitai_download_url(download_url))
normalized_url = normalize_civitai_download_url(download_url)
if normalized_url:
download_urls.append(normalized_url)
return download_urls
@@ -696,8 +705,8 @@ class DownloadManager:
self,
*,
model_type: str,
version_info: Dict,
file_info: Dict,
version_info: Dict[str, Any],
file_info: Dict[str, Any],
save_path: str,
):
if model_type == "checkpoint":
@@ -706,7 +715,7 @@ class DownloadManager:
return EmbeddingMetadata.from_civitai_info(version_info, file_info, save_path)
return LoraMetadata.from_civitai_info(version_info, file_info, save_path)
def _resolve_save_path_from_persisted_record(self, record: Dict) -> Optional[str]:
def _resolve_save_path_from_persisted_record(self, record: Dict[str, Any]) -> Optional[str]:
save_path = record.get("save_path") or record.get("file_path")
if isinstance(save_path, str) and save_path:
return os.path.abspath(save_path)
@@ -728,7 +737,7 @@ class DownloadManager:
return os.path.abspath(os.path.join(save_dir, file_name))
async def _resume_restored_aria2_download(self, download_id: str, record: Dict) -> Dict:
async def _resume_restored_aria2_download(self, download_id: str, record: Dict[str, Any]) -> Dict[str, Any]:
try:
if download_id in self._active_downloads:
self._active_downloads[download_id]["status"] = "downloading"
@@ -842,7 +851,7 @@ class DownloadManager:
self,
previous_download_id: str,
new_download_id: str,
persisted_record: Dict,
persisted_record: Dict[str, Any],
save_path: str,
) -> None:
aria2_downloader = await get_aria2_downloader()
@@ -938,7 +947,7 @@ class DownloadManager:
except Exception:
status_payload = None
if status_payload is not None:
if status_payload is not None and isinstance(gid, str):
remote_status = status_payload.get("status", "")
if remote_status in {"active", "waiting", "paused"}:
await aria2_downloader.restore_transfer(download_id, gid, save_path)
@@ -1115,17 +1124,17 @@ class DownloadManager:
async def _execute_original_download(
self,
model_id,
model_version_id,
save_dir,
relative_path,
model_id: int | None,
model_version_id: int | None,
save_dir: str | None,
relative_path: str,
progress_callback,
use_default_paths,
download_id=None,
transfer_backend="python",
source=None,
file_params=None,
):
use_default_paths: bool,
download_id: str | None = None,
transfer_backend: str = "python",
source: str | None = None,
file_params: Dict[str, Any] | None = None,
) -> Dict[str, Any]:
"""Wrapper for original download_from_civitai implementation"""
try:
# Check if model version already exists in library
@@ -1172,7 +1181,7 @@ class DownloadManager:
# Get version info based on the provided identifier
version_info = await metadata_provider.get_model_version(
model_id, model_version_id
cast(int, model_id), cast(int, model_version_id)
)
if not version_info:
@@ -1183,7 +1192,7 @@ class DownloadManager:
)
metadata_provider = await get_default_metadata_provider()
version_info = await metadata_provider.get_model_version(
model_id, model_version_id
cast(int, model_id), cast(int, model_version_id)
)
if not version_info:
@@ -1388,6 +1397,8 @@ class DownloadManager:
relative_path = self._calculate_relative_path(version_info, model_type)
# Update save directory with relative path if provided
if not save_dir:
return {"success": False, "error": "No save directory specified"}
if relative_path:
base_save_dir = save_dir
save_dir = os.path.join(save_dir, relative_path)
@@ -1561,6 +1572,11 @@ class DownloadManager:
version_info, file_info, save_path
)
logger.info(f"Creating EmbeddingMetadata for {file_name}")
else:
return {
"success": False,
"error": f'Unsupported model type "{model_type}"',
}
# 6. Start download process
if transfer_backend == "aria2" and download_id:
@@ -1580,7 +1596,7 @@ class DownloadManager:
},
)
execute_kwargs = {
execute_kwargs: Dict[str, Any] = {
"download_urls": download_urls,
"save_dir": save_dir,
"metadata": metadata,
@@ -1627,7 +1643,8 @@ class DownloadManager:
)
# If early_access_msg exists and download failed, replace error message
if "early_access_msg" in locals() and not result.get("success", False):
early_access_msg = locals().get("early_access_msg")
if early_access_msg and not result.get("success", False):
result["error"] = early_access_msg
return result
@@ -1652,7 +1669,7 @@ class DownloadManager:
self,
model_type: str,
model_id_value,
version_info: Dict,
version_info: Dict[str, Any],
fallback_version_id=None,
file_path: str | None = None,
) -> None:
@@ -1683,8 +1700,8 @@ class DownloadManager:
try:
await history_service.mark_downloaded(
model_type,
int(version_id),
model_id=int(resolved_model_id) if resolved_model_id is not None else None,
int(cast(Any, version_id)),
model_id=int(cast(Any, resolved_model_id)) if resolved_model_id is not None else None,
source="download",
file_path=file_path,
)
@@ -1701,7 +1718,7 @@ class DownloadManager:
self,
model_type: str,
model_id_value,
version_info: Dict,
version_info: Dict[str, Any],
fallback_version_id=None,
) -> None:
"""Ensure update tracking reflects a newly downloaded version."""
@@ -1725,7 +1742,7 @@ class DownloadManager:
if isinstance(model_info, dict):
resolved_model_id = model_info.get("id")
try:
resolved_model_id = int(resolved_model_id)
resolved_model_id = int(cast(Any, resolved_model_id))
except (TypeError, ValueError):
logger.debug(
"Skipping update sync; invalid model id: %s", resolved_model_id
@@ -1736,7 +1753,7 @@ class DownloadManager:
if version_id is None:
version_id = fallback_version_id
try:
version_id = int(version_id)
version_id = int(cast(Any, version_id))
except (TypeError, ValueError):
logger.debug(
"Skipping update sync; invalid version id for model %s: %s",
@@ -1773,7 +1790,7 @@ class DownloadManager:
for entry in local_versions or []:
vid = entry.get("versionId")
try:
version_ids.add(int(vid))
version_ids.add(int(cast(Any, vid)))
except (TypeError, ValueError):
continue
@@ -1795,7 +1812,7 @@ class DownloadManager:
)
def _calculate_relative_path(
self, version_info: Dict, model_type: str = "lora"
self, version_info: Dict[str, Any], model_type: str = "lora"
) -> str:
"""Calculate relative path using template from settings
@@ -1871,21 +1888,22 @@ class DownloadManager:
download_urls: List[str],
save_dir: str,
metadata,
version_info: Dict,
version_info: Dict[str, Any],
relative_path: str,
progress_callback=None,
model_type: str = "lora",
download_id: str = None,
download_id: str | None = None,
transfer_backend: Optional[str] = None,
) -> Dict:
) -> Dict[str, Any]:
"""Execute the actual download process including preview images and model files"""
metadata_entries: List = []
metadata_entries: List[Any] = []
metadata_files_for_cleanup: List[str] = []
extracted_paths: List[str] = []
metadata_path = ""
preview_targets: List[str] = []
preview_path: str | None = None
preview_nsfw_level = 0
save_path: str | None = None
transfer_backend = (transfer_backend or self._get_model_download_backend()).lower()
try:
resolved, save_path = await self._resolve_download_target_path(
@@ -1933,9 +1951,9 @@ class DownloadManager:
mature_threshold=mature_threshold,
)
preview_url = selected_image.get("url") if selected_image else None
preview_url = cast(Optional[str], selected_image.get("url")) if selected_image else None
media_type = (
(selected_image.get("type") or "").lower() if selected_image else ""
cast(str, selected_image.get("type") or "").lower() if selected_image else ""
)
def _extension_from_url(url: str, fallback: str) -> str:
@@ -1959,9 +1977,10 @@ class DownloadManager:
preview_url, media_type="video"
)
attempt_urls: List[str] = []
if rewritten:
if rewritten and rewritten_url:
attempt_urls.append(rewritten_url)
attempt_urls.append(preview_url)
if preview_url:
attempt_urls.append(preview_url)
seen_attempts = set()
for attempt in attempt_urls:
@@ -1978,7 +1997,7 @@ class DownloadManager:
rewritten_url, rewritten = rewrite_preview_url(
preview_url, media_type="image"
)
if rewritten:
if rewritten and rewritten_url:
preview_ext = _extension_from_url(preview_url, ".png")
preview_path = os.path.splitext(save_path)[0] + preview_ext
success, _ = await downloader.download_file(
@@ -2004,7 +2023,9 @@ class DownloadManager:
)
if success:
with open(temp_path, "wb") as temp_file_handle:
temp_file_handle.write(content)
temp_file_handle.write(
content if isinstance(content, bytes) else content.encode("utf-8")
)
preview_path = (
os.path.splitext(save_path)[0] + ".webp"
)
@@ -2056,6 +2077,8 @@ class DownloadManager:
last_error = None
for download_url in download_urls:
download_url = normalize_civitai_download_url(download_url)
if download_url is None:
continue
use_auth = download_url.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
if transfer_backend == "aria2" and download_id:
await self._persist_aria2_state(
@@ -2160,6 +2183,10 @@ class DownloadManager:
"error": f"Zip archive does not contain any supported model files ({supported_text})",
}
actual_file_paths = extracted_paths
# The archive entry's AutoV3 (if any) describes the zip itself,
# not the extracted models; clear it so per-file header
# resolution applies to every extracted model.
metadata.autov3 = None
try:
os.remove(save_path)
except OSError as exc:
@@ -2235,7 +2262,7 @@ class DownloadManager:
entry, normalized_file_path, adjust_root
)
if adjusted_entry is not None:
entry = adjusted_entry
entry = cast(Any, adjusted_entry)
metadata_entries[index] = entry
metadata_file_path = (
@@ -2355,11 +2382,11 @@ class DownloadManager:
async def _build_metadata_entries(
self, base_metadata, file_paths: List[str]
) -> List:
) -> List[Any]:
if not file_paths:
return []
entries: List = []
entries: List[Any] = []
for index, file_path in enumerate(file_paths):
entry = base_metadata if index == 0 else copy.deepcopy(base_metadata)
# Update file paths without modifying size and modified timestamps
@@ -2374,6 +2401,16 @@ class DownloadManager:
sha256 = await calculate_sha256(file_path)
if sha256:
entry.sha256 = sha256.lower()
# AutoV3: the Civitai-reported value for the downloaded file (set
# by from_civitai_info) takes precedence. Only the un-checked
# state (None) triggers a header read; '' (checked-unavailable)
# is never re-read, honoring the three-state contract so rows
# marked at download time stay untouched by later passes.
if entry.autov3 is None:
autov3 = await asyncio.get_running_loop().run_in_executor(
None, calculate_autov3, file_path
)
entry.autov3 = (autov3 or "").lower()
entries.append(entry)
return entries
@@ -2392,7 +2429,7 @@ class DownloadManager:
return destination
def _distribute_preview_to_entries(
self, preview_path: str, entries: List
self, preview_path: str, entries: List[Any]
) -> List[str]:
if not preview_path or not entries:
return []
@@ -2451,7 +2488,7 @@ class DownloadManager:
progress_callback, normalized_snapshot, rounded_progress
)
async def cancel_download(self, download_id: str) -> Dict:
async def cancel_download(self, download_id: str) -> Dict[str, Any]:
"""Cancel an active download by download_id
Args:
@@ -2533,7 +2570,7 @@ class DownloadManager:
self._download_tasks.pop(download_id, None)
await self._aria2_state_store.remove(download_id)
async def skip_download(self, download_id: str) -> Dict:
async def skip_download(self, download_id: str) -> Dict[str, Any]:
"""Skip a download while preserving all partial files on disk.
Removes all in-memory tracking (asyncio task, semaphore, active/pause
@@ -2616,7 +2653,7 @@ class DownloadManager:
# Preserve aria2 state store entry so the partial download
# info survives restarts and can be resumed later
async def pause_download(self, download_id: str) -> Dict:
async def pause_download(self, download_id: str) -> Dict[str, Any]:
"""Pause an active download without losing progress."""
await self._restore_persisted_downloads()
@@ -2663,7 +2700,7 @@ class DownloadManager:
return {"success": True, "message": "Download paused successfully"}
async def resume_download(self, download_id: str) -> Dict:
async def resume_download(self, download_id: str) -> Dict[str, Any]:
"""Resume a previously paused download."""
await self._restore_persisted_downloads()
@@ -2680,7 +2717,7 @@ class DownloadManager:
self._pause_events[download_id] = pause_control
self._active_downloads[download_id] = self._build_restored_download_info(
persisted,
os.path.abspath(save_path),
os.path.abspath(cast(str, save_path)),
)
if pause_control.is_set():
@@ -2807,7 +2844,7 @@ class DownloadManager:
elif asyncio.iscoroutine(result):
await result
async def get_active_downloads(self) -> Dict:
async def get_active_downloads(self) -> Dict[str, Any]:
"""Get information about all active downloads
Returns:
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
from __future__ import annotations
import asyncio
+13 -8
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
"""
Unified download manager for all HTTP/HTTPS downloads in the application.
@@ -20,7 +24,7 @@ from dataclasses import dataclass
from datetime import datetime, timedelta
from email.utils import parsedate_to_datetime
from urllib.parse import urlparse
from typing import Optional, Dict, Tuple, Callable, Union, Awaitable
from typing import Optional, Dict, Tuple, Callable, Union, Awaitable, Any, cast
from ..services.settings_manager import get_settings_manager
from .connectivity_guard import (
OFFLINE_COOLDOWN_ERROR,
@@ -204,6 +208,7 @@ class Downloader:
# Double check after acquiring lock
if self._session is None or self._should_refresh_session():
await self._create_session()
assert self._session is not None
return self._session
@property
@@ -231,7 +236,7 @@ class Downloader:
)
try:
timeout_value = float(raw_value)
timeout_value = float(cast(Any, raw_value))
except (TypeError, ValueError):
timeout_value = default_timeout
@@ -243,7 +248,7 @@ class Downloader:
raw_value = os.environ.get("COMFYUI_DOWNLOAD_MAX_RETRIES")
try:
retries = int(raw_value)
retries = int(cast(Any, raw_value))
except (TypeError, ValueError):
retries = default_retries
@@ -320,7 +325,7 @@ class Downloader:
# CA coverage across different Python environments (especially
# embedded/compatibility Python builds).
try:
import certifi # type: ignore[import-untyped]
import certifi # pyright: ignore[reportMissingTypeStubs]
ca_path = certifi.where()
ssl_context = ssl.create_default_context(cafile=ca_path)
@@ -330,7 +335,7 @@ class Downloader:
logger.debug("SSL: certifi unavailable; using system default CA bundle")
# Optimize TCP connection parameters
connector_kwargs = dict(
connector_kwargs: Dict[str, Any] = dict(
ssl=ssl_context,
limit=8, # Concurrent connections
ttl_dns_cache=300, # DNS cache timeout
@@ -890,7 +895,7 @@ class Downloader:
use_auth: bool = False,
custom_headers: Optional[Dict[str, str]] = None,
return_headers: bool = False,
) -> Tuple[bool, Union[bytes, str], Optional[Dict]]:
) -> Tuple[bool, Union[bytes, str], Optional[Dict[str, Any]]]:
"""
Download a file to memory (for small files like preview images)
@@ -976,7 +981,7 @@ class Downloader:
url: str,
use_auth: bool = False,
custom_headers: Optional[Dict[str, str]] = None,
) -> Tuple[bool, Union[Dict, str]]:
) -> Tuple[bool, Union[Dict[str, Any], str]]:
"""
Get response headers without downloading the full content
@@ -1036,7 +1041,7 @@ class Downloader:
use_auth: bool = False,
custom_headers: Optional[Dict[str, str]] = None,
**kwargs,
) -> Tuple[bool, Union[Dict, str]]:
) -> Tuple[bool, Union[Dict[str, Any], str, RateLimitError]]:
"""
Make a generic HTTP request and return JSON response
+1 -1
View File
@@ -27,7 +27,7 @@ class EmbeddingScanner(ModelScanner):
roots.extend(config.embeddings_roots or [])
roots.extend(config.extra_embeddings_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:
+28 -28
View File
@@ -1,6 +1,6 @@
import os
import logging
from typing import Dict, Optional
from typing import Any, Dict, Optional
from .base_model_service import BaseModelService
from .auto_tag_service import extract_auto_tags
@@ -21,58 +21,58 @@ class EmbeddingService(BaseModelService):
"""
super().__init__("embedding", scanner, EmbeddingMetadata, update_service=update_service)
async def format_response(self, embedding_data: Dict) -> Optional[Dict]:
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Format Embedding data for API response.
Returns None when the entry is missing critical fields (corrupted cache
row), so the handler layer can filter it out. See issue #730.
"""
# Guard against corrupted cache entries missing critical fields
file_path = embedding_data.get("file_path")
file_path = model_data.get("file_path")
if not file_path or not isinstance(file_path, str):
logger.warning(
"Skipping corrupted embedding entry (missing file_path): %s",
embedding_data.get("file_name", "<unknown>"),
model_data.get("file_name", "<unknown>"),
)
return None
# Get sub_type from cache entry (new canonical field)
sub_type = embedding_data.get("sub_type", "embedding")
sub_type = model_data.get("sub_type", "embedding")
file_name = embedding_data.get("file_name") or ""
model_name = embedding_data.get("model_name") or file_name
folder = embedding_data.get("folder") or ""
file_name = model_data.get("file_name") or ""
model_name = model_data.get("model_name") or file_name
folder = model_data.get("folder") or ""
return {
"model_name": model_name,
"file_name": file_name,
"preview_url": config.get_preview_static_url(embedding_data.get("preview_url", "")),
"preview_nsfw_level": embedding_data.get("preview_nsfw_level", 0),
"base_model": embedding_data.get("base_model", ""),
"preview_url": config.get_preview_static_url(model_data.get("preview_url", "")),
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
"base_model": model_data.get("base_model", ""),
"folder": folder,
"sha256": embedding_data.get("sha256", ""),
"sha256": model_data.get("sha256", ""),
"file_path": file_path.replace(os.sep, "/"),
"file_size": embedding_data.get("size", 0),
"modified": embedding_data.get("modified", ""),
"tags": embedding_data.get("tags", []),
"from_civitai": embedding_data.get("from_civitai", True),
# "usage_count": embedding_data.get("usage_count", 0), # TODO: Enable when embedding usage tracking is implemented
"notes": embedding_data.get("notes", ""),
"file_size": model_data.get("size", 0),
"modified": model_data.get("modified", ""),
"tags": model_data.get("tags", []),
"from_civitai": model_data.get("from_civitai", True),
# "usage_count": model_data.get("usage_count", 0), # TODO: Enable when embedding usage tracking is implemented
"notes": model_data.get("notes", ""),
"sub_type": sub_type,
"favorite": embedding_data.get("favorite", False),
"exclude": bool(embedding_data.get("exclude", False)),
"update_available": bool(embedding_data.get("update_available", False)),
"skip_metadata_refresh": bool(embedding_data.get("skip_metadata_refresh", False)),
"civitai": self.filter_civitai_data(embedding_data.get("civitai", {}), minimal=True),
"auto_tags": embedding_data.get("auto_tags") or extract_auto_tags(embedding_data),
"version_count": embedding_data.get("version_count"),
"hf_url": embedding_data.get("hf_url", ""),
"favorite": model_data.get("favorite", False),
"exclude": bool(model_data.get("exclude", False)),
"update_available": bool(model_data.get("update_available", False)),
"skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)),
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
"version_count": model_data.get("version_count"),
"hf_url": model_data.get("hf_url", ""),
}
def find_duplicate_hashes(self) -> Dict:
def find_duplicate_hashes(self) -> Dict[str, Any]:
"""Find Embeddings with duplicate SHA256 hashes"""
return self.scanner._hash_index.get_duplicate_hashes()
def find_duplicate_filenames(self) -> Dict:
def find_duplicate_filenames(self) -> Dict[str, Any]:
"""Find Embeddings with conflicting filenames"""
return self.scanner._hash_index.get_duplicate_filenames()
@@ -35,7 +35,7 @@ class CleanupResult:
def to_dict(self) -> Dict[str, object]:
"""Convert the dataclass to a serialisable dictionary."""
data = {
data: Dict[str, object] = {
"success": self.success,
"checked_folders": self.checked_folders,
"moved_empty_folders": self.moved_empty_folders,
+14 -4
View File
@@ -1,10 +1,12 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import logging
from typing import List
from ..utils.models import LoraMetadata
from ..config import config
from .model_scanner import ModelScanner
from .model_hash_index import ModelHashIndex # Changed from LoraHashIndex to ModelHashIndex
import sys
logger = logging.getLogger(__name__)
@@ -15,8 +17,10 @@ class LoraScanner(ModelScanner):
def __init__(self):
# Define supported file extensions
file_extensions = {'.safetensors'}
# Initialize parent class with ModelHashIndex
from .model_hash_index import ModelHashIndex
super().__init__(
model_type="lora",
model_class=LoraMetadata,
@@ -26,11 +30,13 @@ class LoraScanner(ModelScanner):
def get_model_roots(self) -> List[str]:
"""Get lora root directories (including extra paths)"""
from ..config import config
roots: List[str] = []
roots.extend(config.loras_roots or [])
roots.extend(config.extra_loras_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:
@@ -68,8 +74,12 @@ class LoraScanner(ModelScanner):
test_hash = next(iter(self._hash_index._hash_to_path.keys()))
test_path = self._hash_index.get_path(test_hash)
logger.debug(f"\nTest lookup by hash: {test_hash[:8]}... -> {test_path}")
if test_path is None:
return
# Also test reverse lookup
test_hash_result = self._hash_index.get_hash(test_path)
if test_hash_result is None:
return
logger.debug(f"Test reverse lookup: {test_path} -> {test_hash_result[:8]}...\n\n")
+41 -41
View File
@@ -1,7 +1,7 @@
import logging
import json
import os
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional
from .base_model_service import BaseModelService
from .model_query import resolve_sub_type
@@ -24,7 +24,7 @@ class LoraService(BaseModelService):
"""
super().__init__("lora", scanner, LoraMetadata, update_service=update_service)
async def format_response(self, lora_data: Dict) -> Optional[Dict]:
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Format LoRA data for API response.
Returns None when the entry is missing critical fields (corrupted cache
@@ -32,56 +32,56 @@ class LoraService(BaseModelService):
whole listing request. See issue #730.
"""
# Guard against corrupted cache entries missing critical fields
file_path = lora_data.get("file_path")
file_path = model_data.get("file_path")
if not file_path or not isinstance(file_path, str):
logger.warning(
"Skipping corrupted LoRA entry (missing file_path): %s",
lora_data.get("file_name", "<unknown>"),
model_data.get("file_name", "<unknown>"),
)
return None
# Resolve sub_type using priority: sub_type > model_type > civitai.model.type > default
# Normalize to lowercase for consistent API responses
sub_type = resolve_sub_type(lora_data).lower()
sub_type = resolve_sub_type(model_data).lower()
file_name = lora_data.get("file_name") or ""
model_name = lora_data.get("model_name") or file_name
folder = lora_data.get("folder") or ""
file_name = model_data.get("file_name") or ""
model_name = model_data.get("model_name") or file_name
folder = model_data.get("folder") or ""
return {
"model_name": model_name,
"file_name": file_name,
"preview_url": config.get_preview_static_url(
lora_data.get("preview_url", "")
model_data.get("preview_url", "")
),
"preview_nsfw_level": lora_data.get("preview_nsfw_level", 0),
"base_model": lora_data.get("base_model", ""),
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
"base_model": model_data.get("base_model", ""),
"folder": folder,
"sha256": lora_data.get("sha256", ""),
"sha256": model_data.get("sha256", ""),
"file_path": file_path.replace(os.sep, "/"),
"file_size": lora_data.get("size", 0),
"modified": lora_data.get("modified", ""),
"tags": lora_data.get("tags", []),
"from_civitai": lora_data.get("from_civitai", True),
"usage_count": lora_data.get("usage_count", 0),
"usage_tips": lora_data.get("usage_tips", ""),
"notes": lora_data.get("notes", ""),
"favorite": lora_data.get("favorite", False),
"exclude": bool(lora_data.get("exclude", False)),
"update_available": bool(lora_data.get("update_available", False)),
"file_size": model_data.get("size", 0),
"modified": model_data.get("modified", ""),
"tags": model_data.get("tags", []),
"from_civitai": model_data.get("from_civitai", True),
"usage_count": model_data.get("usage_count", 0),
"usage_tips": model_data.get("usage_tips", ""),
"notes": model_data.get("notes", ""),
"favorite": model_data.get("favorite", False),
"exclude": bool(model_data.get("exclude", False)),
"update_available": bool(model_data.get("update_available", False)),
"skip_metadata_refresh": bool(
lora_data.get("skip_metadata_refresh", False)
model_data.get("skip_metadata_refresh", False)
),
"sub_type": sub_type,
"civitai": self.filter_civitai_data(
lora_data.get("civitai", {}), minimal=True
model_data.get("civitai", {}), minimal=True
),
"auto_tags": lora_data.get("auto_tags") or extract_auto_tags(lora_data),
"version_count": lora_data.get("version_count"),
"hf_url": lora_data.get("hf_url", ""),
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
"version_count": model_data.get("version_count"),
"hf_url": model_data.get("hf_url", ""),
}
async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]:
async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]:
"""Apply LoRA-specific filters"""
# Handle first_letter filter for LoRAs
first_letter = kwargs.get("first_letter")
@@ -152,7 +152,7 @@ class LoraService(BaseModelService):
return data
def _filter_by_first_letter(self, data: List[Dict], letter: str) -> List[Dict]:
def _filter_by_first_letter(self, data: List[Dict[str, Any]], letter: str) -> List[Dict[str, Any]]:
"""Filter data by first letter of model name
Special handling:
@@ -307,7 +307,7 @@ class LoraService(BaseModelService):
return None
@staticmethod
def get_recommended_strength_from_lora_data(lora_data: Dict) -> Optional[float]:
def get_recommended_strength_from_lora_data(lora_data: Dict[str, Any]) -> Optional[float]:
"""Parse usage_tips JSON and extract recommended model strength."""
try:
usage_tips = lora_data.get("usage_tips", "")
@@ -320,7 +320,7 @@ class LoraService(BaseModelService):
@staticmethod
def get_recommended_clip_strength_from_lora_data(
lora_data: Dict,
lora_data: Dict[str, Any],
) -> Optional[float]:
"""Parse usage_tips JSON and extract recommended clip strength."""
try:
@@ -332,7 +332,7 @@ class LoraService(BaseModelService):
except (json.JSONDecodeError, TypeError, AttributeError):
return None
async def get_lora_metadata_by_filename(self, filename: str) -> Optional[Dict]:
async def get_lora_metadata_by_filename(self, filename: str) -> Optional[Dict[str, Any]]:
"""Return cached raw metadata for a LoRA matching the given filename."""
cache = await self.scanner.get_cached_data(force_refresh=False)
@@ -357,11 +357,11 @@ class LoraService(BaseModelService):
return None
def find_duplicate_hashes(self) -> Dict:
def find_duplicate_hashes(self) -> Dict[str, Any]:
"""Find LoRAs with duplicate SHA256 hashes"""
return self.scanner._hash_index.get_duplicate_hashes()
def find_duplicate_filenames(self) -> Dict:
def find_duplicate_filenames(self) -> Dict[str, Any]:
"""Find LoRAs with conflicting filenames"""
return self.scanner._hash_index.get_duplicate_filenames()
@@ -373,8 +373,8 @@ class LoraService(BaseModelService):
use_same_clip_strength: bool = True,
clip_strength_min: float = 0.0,
clip_strength_max: float = 1.0,
locked_loras: Optional[List[Dict]] = None,
pool_config: Optional[Dict] = None,
locked_loras: Optional[List[Dict[str, Any]]] = None,
pool_config: Optional[Dict[str, Any]] = None,
count_mode: str = "fixed",
count_min: int = 3,
count_max: int = 7,
@@ -382,7 +382,7 @@ class LoraService(BaseModelService):
recommended_strength_scale_min: float = 0.5,
recommended_strength_scale_max: float = 1.0,
seed: Optional[int] = None,
) -> List[Dict]:
) -> List[Dict[str, Any]]:
"""
Get random LoRAs with specified strength ranges.
@@ -513,8 +513,8 @@ class LoraService(BaseModelService):
return result_loras
async def _apply_pool_filters(
self, available_loras: List[Dict], pool_config: Dict
) -> List[Dict]:
self, available_loras: List[Dict[str, Any]], pool_config: Dict[str, Any]
) -> List[Dict[str, Any]]:
"""
Apply pool_config filters to available LoRAs.
@@ -671,8 +671,8 @@ class LoraService(BaseModelService):
return available_loras
async def get_cycler_list(
self, pool_config: Optional[Dict] = None, sort_by: str = "filename"
) -> List[Dict]:
self, pool_config: Optional[Dict[str, Any]] = None, sort_by: str = "filename"
) -> List[Dict[str, Any]]:
"""
Get filtered and sorted LoRA list for cycling.
+5 -1
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import os
import logging
from .model_metadata_provider import (
@@ -170,7 +174,7 @@ def _wrap_provider_with_rate_limit(provider_name: str | None, provider: ModelMet
return RateLimitRetryingProvider(provider, label=provider_name)
async def get_metadata_provider(provider_name: str = None):
async def get_metadata_provider(provider_name: str | None = None):
"""Get a specific metadata provider or default provider with rate-limit handling."""
provider_manager = await ModelMetadataProviderManager.get_instance()
+20 -7
View File
@@ -6,25 +6,26 @@ import json
import logging
import os
from datetime import datetime
from typing import Any, Awaitable, Callable, Dict, Iterable, Optional
from typing import Any, Awaitable, Callable, Dict, Iterable, Optional, Protocol
from ..services.settings_manager import SettingsManager
from ..utils.civitai_utils import resolve_license_payload
from ..utils.model_utils import determine_base_model
from ..utils.models import autov3_from_civitai_files
from .connectivity_guard import OFFLINE_FRIENDLY_MESSAGE, is_expected_offline_error
from .errors import RateLimitError
logger = logging.getLogger(__name__)
class MetadataProviderProtocol:
class MetadataProviderProtocol(Protocol):
"""Subset of metadata provider interface consumed by the sync service."""
async def get_model_by_hash(self, sha256: str) -> tuple[Optional[Dict[str, Any]], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> tuple[Optional[Dict[str, Any]], Optional[str]]:
...
async def get_model_version(
self, model_id: int, model_version_id: Optional[int]
self, model_id: Any = None, version_id: Any = None
) -> Optional[Dict[str, Any]]:
...
@@ -38,8 +39,8 @@ class MetadataSyncService:
metadata_manager,
preview_service,
settings: SettingsManager,
default_metadata_provider_factory: Callable[[], Awaitable[MetadataProviderProtocol]],
metadata_provider_selector: Callable[[str], Awaitable[MetadataProviderProtocol]],
default_metadata_provider_factory: Callable[..., Awaitable[MetadataProviderProtocol]],
metadata_provider_selector: Callable[..., Awaitable[MetadataProviderProtocol]],
) -> None:
self._metadata_manager = metadata_manager
self._preview_service = preview_service
@@ -152,6 +153,18 @@ class MetadataSyncService:
civitai_metadata.get("baseModel")
)
# Civitai-first AutoV3 propagation: the freshly fetched version
# metadata may report an AutoV3 for the file whose SHA256 matches the
# local model. Persist it now so recipe matching sees it immediately —
# no full rescan or restart required (the header is never re-read to
# upgrade the checked-unavailable '' state).
sha256_value = (local_metadata.get("sha256") or "").lower()
civitai_autov3 = autov3_from_civitai_files(
local_metadata.get("civitai"), sha256_value
)
if civitai_autov3:
local_metadata["autov3"] = civitai_autov3
await self._preview_service.ensure_preview_for_metadata(
metadata_path, local_metadata, civitai_metadata.get("images", [])
)
@@ -479,7 +492,7 @@ class MetadataSyncService:
if not file_paths:
raise ValueError("No file paths provided for verification")
results = {
results: Dict[str, Any] = {
"verified_as_duplicates": True,
"mismatched_files": [],
"new_hash_map": {},
+21 -16
View File
@@ -31,17 +31,22 @@ DISPLAY_NAME_MODES = {"model_name", "file_name"}
class ModelCache:
"""Cache structure for model data with extensible sorting."""
raw_data: List[Dict]
raw_data: List[Dict[str, Any]]
folders: List[str]
version_index: Dict[int, Dict] = field(default_factory=dict)
version_index: Dict[int, Dict[str, Any]] = field(default_factory=dict)
model_id_index: Dict[int, List[Dict[str, Any]]] = field(default_factory=dict)
name_display_mode: str = "model_name"
_lock: Any = field(init=False, repr=False, default=None)
# Cache for last sort: (sort_key, order, seed) -> sorted list
_last_sort: Tuple[Optional[str], str, Optional[str]] = field(
init=False, repr=False, default=(None, "asc", None)
)
_last_sorted_data: List[Dict[str, Any]] = field(
init=False, repr=False, default_factory=list
)
def __post_init__(self):
self._lock = asyncio.Lock()
# Cache for last sort: (sort_key, order, seed) -> sorted list
self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
self._last_sorted_data: List[Dict] = []
self._normalize_raw_data()
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
# Default sort on init
@@ -64,7 +69,7 @@ class ModelCache:
return ""
return str(value)
def _normalize_item(self, item: Dict) -> None:
def _normalize_item(self, item: Dict[str, Any]) -> None:
"""Ensure core metadata fields are present and string typed."""
if not isinstance(item, dict):
@@ -80,7 +85,7 @@ class ModelCache:
for item in self.raw_data:
self._normalize_item(item)
def _get_display_name(self, item: Dict) -> str:
def _get_display_name(self, item: Dict[str, Any]) -> str:
"""Return the value used for name-based sorting based on display settings."""
if self.name_display_mode == "file_name":
@@ -114,7 +119,7 @@ class ModelCache:
for item in self.raw_data:
self.add_to_version_index(item)
def add_to_version_index(self, item: Dict) -> None:
def add_to_version_index(self, item: Dict[str, Any]) -> None:
"""Register a cache item in the version/model indexes if possible."""
civitai_data = item.get('civitai') if isinstance(item, dict) else None
@@ -143,7 +148,7 @@ class ModelCache:
else:
versions.append(descriptor)
def remove_from_version_index(self, item: Dict) -> None:
def remove_from_version_index(self, item: Dict[str, Any]) -> None:
"""Remove a cache item from the version/model indexes if present."""
civitai_data = item.get('civitai') if isinstance(item, dict) else None
@@ -177,7 +182,7 @@ class ModelCache:
def _build_version_descriptor(
self,
item: Dict,
item: Dict[str, Any],
civitai_data: Dict[str, Any],
version_id: int,
) -> Optional[Dict[str, Any]]:
@@ -204,8 +209,8 @@ class ModelCache:
async def resort(self):
"""Resort cached data according to last sort mode if set"""
async with self._lock:
if self._last_sort[0] is not None:
sort_key, order, seed = self._last_sort
sort_key, order, seed = self._last_sort
if sort_key is not None:
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sorted_data = sorted_data
# Update folder list
@@ -219,7 +224,7 @@ class ModelCache:
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
self.rebuild_version_index()
def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]:
def _sort_data(self, data: List[Dict[str, Any]], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict[str, Any]]:
"""Sort data by sort_key and order"""
start_time = time.perf_counter()
reverse = (order == 'desc')
@@ -293,7 +298,7 @@ class ModelCache:
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
return result
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]:
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict[str, Any]]:
"""Get sorted data by sort_key and order, using cache if possible"""
async with self._lock:
cache_key = (sort_key, order, seed)
@@ -321,8 +326,8 @@ class ModelCache:
self.name_display_mode = normalized
if self._last_sort[0] == 'name':
sort_key, order, seed = self._last_sort
sort_key, order, seed = self._last_sort
if sort_key == 'name':
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:
+3 -1
View File
@@ -41,7 +41,7 @@ class AutoOrganizeResult:
def to_dict(self) -> Dict[str, Any]:
"""Convert result to dictionary"""
result = {
result: Dict[str, Any] = {
'success': self.status != 'error',
'status': self.status,
'message': f'Auto-organize {self.operation_type} completed: {self.success_count} moved, {self.skipped_count} skipped, {self.failure_count} failed out of {self.total} total',
@@ -418,6 +418,8 @@ class ModelFileService:
"""Calculate the target directory for a model"""
if is_flat_structure:
file_path = model.get('file_path')
if not isinstance(file_path, str):
return None
current_dir = os.path.dirname(file_path)
# Check if already in root directory
+53 -4
View File
@@ -8,11 +8,12 @@ class ModelHashIndex:
self._hash_to_path: Dict[str, str] = {}
self._filename_to_hash: Dict[str, str] = {}
self._autov2_to_path: Dict[str, str] = {}
self._autov3_to_path: Dict[str, str] = {}
# New data structures for tracking duplicates
self._duplicate_hashes: Dict[str, List[str]] = {} # sha256 -> list of paths
self._duplicate_filenames: Dict[str, List[str]] = {} # filename -> list of paths
def add_entry(self, sha256: str, file_path: str) -> None:
def add_entry(self, sha256: str, file_path: str, autov3: Optional[str] = None) -> None:
"""Add or update hash index entry"""
if not sha256 or not file_path:
return
@@ -33,9 +34,14 @@ class ModelHashIndex:
self._duplicate_hashes.setdefault(sha256, []).append(file_path)
# Track duplicates by filename - FIXED LOGIC
is_re_registration = False
existing_hash: Optional[str] = None
if filename in self._filename_to_hash:
existing_hash = self._filename_to_hash[filename]
existing_path = self._hash_to_path.get(existing_hash)
# Same path registered again (e.g. a file replaced in place with
# new content) — used below to drop its stale autov3 mapping.
is_re_registration = existing_path == file_path
# If this is a different file with the same filename
if existing_path and existing_path != file_path:
@@ -67,12 +73,36 @@ class ModelHashIndex:
# AutoV2 = first 10 chars of SHA256
if len(sha256) >= 10:
self._autov2_to_path[sha256[:10]] = file_path
# AutoV3 is an independent hash (not derived from SHA256), stored as-is.
# Drop stale mappings for a path when it is re-registered with a NEW
# sha256 (file replaced in place) or with an explicit new autov3 value
# (correction). Re-registering the SAME file with the same sha256 and
# no autov3 (e.g. lazy-hash completion) must never clear its existing
# mapping. First-time registrations stay O(1).
if autov3:
autov3 = autov3.lower()
if is_re_registration and (existing_hash != sha256 or autov3):
stale_autov3_keys = [
key for key, mapped_path in self._autov3_to_path.items()
if mapped_path == file_path and key != autov3
]
for key in stale_autov3_keys:
del self._autov3_to_path[key]
if autov3:
self._autov3_to_path[autov3] = file_path
def add_autov3(self, autov3: str, file_path: str) -> None:
"""Add or update an AutoV3-only index entry (used when only AutoV3 is known)"""
if not autov3:
return
autov3 = autov3.lower()
self._autov3_to_path[autov3] = file_path
def _get_filename_from_path(self, file_path: str) -> str:
"""Extract filename without extension from path"""
return os.path.splitext(os.path.basename(file_path))[0]
def remove_by_path(self, file_path: str, hash_val: str = None) -> None:
def remove_by_path(self, file_path: str, hash_val: Optional[str] = None) -> None:
"""Remove entry by file path"""
filename = self._get_filename_from_path(file_path)
@@ -167,6 +197,11 @@ class ModelHashIndex:
for k in autov2_keys_to_remove:
del self._autov2_to_path[k]
# Remove from AutoV3 index
autov3_keys_to_remove = [k for k, v in self._autov3_to_path.items() if v == file_path]
for k in autov3_keys_to_remove:
del self._autov3_to_path[k]
def remove_by_hash(self, sha256: str) -> None:
"""Remove entry by hash"""
sha256 = sha256.lower()
@@ -189,6 +224,11 @@ class ModelHashIndex:
autov2_key = sha256[:10]
if autov2_key in self._autov2_to_path:
del self._autov2_to_path[autov2_key]
# Remove AutoV3 entries pointing to any removed path
autov3_keys_to_remove = [k for k, v in self._autov3_to_path.items() if v in paths_to_remove]
for k in autov3_keys_to_remove:
del self._autov3_to_path[k]
# Update filename-to-hash and duplicate filenames for all paths
for path_to_remove in paths_to_remove:
@@ -209,22 +249,26 @@ class ModelHashIndex:
del self._duplicate_filenames[fname]
def has_hash(self, hash_value: str) -> bool:
"""Check if hash exists in index (SHA256 or AutoV2)"""
"""Check if hash exists in index (SHA256, AutoV2, or AutoV3)"""
normalized = hash_value.lower()
if normalized in self._hash_to_path:
return True
if len(normalized) == 10:
return normalized in self._autov2_to_path
if len(normalized) == 12:
return normalized in self._autov3_to_path
return False
def get_path(self, hash_value: str) -> Optional[str]:
"""Get file path for a hash (SHA256 or AutoV2)"""
"""Get file path for a hash (SHA256, AutoV2, or AutoV3)"""
normalized = hash_value.lower()
path = self._hash_to_path.get(normalized)
if path is not None:
return path
if len(normalized) == 10:
return self._autov2_to_path.get(normalized)
if len(normalized) == 12:
return self._autov3_to_path.get(normalized)
return None
def get_hash(self, file_path: str) -> Optional[str]:
@@ -243,6 +287,7 @@ class ModelHashIndex:
self._hash_to_path.clear()
self._filename_to_hash.clear()
self._autov2_to_path.clear()
self._autov3_to_path.clear()
self._duplicate_hashes.clear()
self._duplicate_filenames.clear()
@@ -253,6 +298,10 @@ class ModelHashIndex:
def get_all_filenames(self) -> Set[str]:
"""Get all filenames in the index"""
return set(self._filename_to_hash.keys())
def get_all_autov3(self) -> Dict[str, str]:
"""Get a snapshot of all AutoV3 hashes mapped to their file paths"""
return dict(self._autov3_to_path)
def get_duplicate_hashes(self) -> Dict[str, List[str]]:
"""Get dictionary of duplicate hashes and their paths"""
+13 -6
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import logging
import os
from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING
from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING, cast
from ..services.service_registry import ServiceRegistry
from ..utils.constants import PREVIEW_EXTENSIONS
@@ -87,8 +87,8 @@ class ModelLifecycleService:
scanner,
metadata_manager,
metadata_loader: Callable[[str], Awaitable[Dict[str, object]]],
recipe_scanner_factory: Callable[[], Awaitable] | None = None,
update_service: "ModelUpdateService" | None = None,
recipe_scanner_factory: Callable[[], Awaitable[Any]] | None = None,
update_service: Optional["ModelUpdateService"] = None,
) -> None:
self._scanner = scanner
self._metadata_manager = metadata_manager
@@ -138,6 +138,9 @@ class ModelLifecycleService:
item for item in cache.raw_data if item.get("file_path") != file_path
]
await cache.resort()
bump_cache_version = getattr(self._scanner, "bump_cache_version", None)
if callable(bump_cache_version):
bump_cache_version()
if hasattr(self._scanner, "_hash_index") and self._scanner._hash_index:
self._scanner._hash_index.remove_by_path(file_path)
@@ -146,7 +149,7 @@ class ModelLifecycleService:
persist_current_cache = getattr(self._scanner, "_persist_current_cache", None)
if callable(persist_current_cache):
await persist_current_cache()
await cast(Awaitable[Any], persist_current_cache())
return {"success": True, "deleted_files": deleted_files}
@@ -244,6 +247,9 @@ class ModelLifecycleService:
item for item in cache.raw_data if item["file_path"] != file_path
]
await cache.resort()
bump_cache_version = getattr(self._scanner, "bump_cache_version", None)
if callable(bump_cache_version):
bump_cache_version()
excluded = getattr(self._scanner, "_excluded_models", None)
if isinstance(excluded, list):
@@ -252,7 +258,7 @@ class ModelLifecycleService:
persist_current_cache = getattr(self._scanner, "_persist_current_cache", None)
if callable(persist_current_cache):
await persist_current_cache()
await cast(Awaitable[Any], persist_current_cache())
message = f"Model {os.path.basename(file_path)} excluded"
return {"success": True, "message": message}
@@ -357,7 +363,8 @@ class ModelLifecycleService:
if os.path.exists(metadata_path):
metadata = await self._metadata_loader(metadata_path)
hash_value = metadata.get("sha256") if isinstance(metadata, dict) else None
raw_hash = metadata.get("sha256") if isinstance(metadata, dict) else None
hash_value = raw_hash if isinstance(raw_hash, str) else None
renamed_files: List[str] = []
new_metadata_path: Optional[str] = None
+52 -52
View File
@@ -10,7 +10,7 @@ from .errors import RateLimitError, ResourceNotFoundError
try:
from bs4 import BeautifulSoup
except ImportError as exc:
BeautifulSoup = None # type: ignore[assignment]
BeautifulSoup = None # pyright: ignore[reportAssignmentType]
_BS4_IMPORT_ERROR = exc
else:
_BS4_IMPORT_ERROR = None
@@ -18,7 +18,7 @@ else:
try:
import aiosqlite
except ImportError as exc:
aiosqlite = None # type: ignore[assignment]
aiosqlite = None # pyright: ignore[reportAssignmentType]
_AIOSQLITE_IMPORT_ERROR = exc
else:
_AIOSQLITE_IMPORT_ERROR = None
@@ -105,24 +105,24 @@ class ModelMetadataProvider(ABC):
"""Base abstract class for all model metadata providers"""
@abstractmethod
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Find model by hash value"""
pass
@abstractmethod
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
"""Get all versions of a model with their details"""
pass
async def get_model_versions_bulk(
self, model_ids: Sequence[int]
) -> Optional[Dict[int, Dict]]:
) -> Optional[Dict[int, Dict[str, Any]]]:
"""Fetch model versions for multiple model ids when supported."""
raise NotImplementedError
async def get_model_versions_by_hashes(
self, hashes: List[str]
) -> Optional[List[Dict]]:
) -> Optional[List[Dict[str, Any]]]:
"""Fetch full version details for multiple SHA256 hashes.
Used specifically to retrieve ``usageControl`` which is only
@@ -133,17 +133,17 @@ class ModelMetadataProvider(ABC):
raise NotImplementedError
@abstractmethod
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
"""Get specific model version with additional metadata"""
pass
@abstractmethod
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[str, Any]], Optional[str]]:
"""Fetch model version metadata"""
pass
@abstractmethod
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
"""Fetch one page of models owned by the specified user.
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
@@ -161,29 +161,29 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
def __init__(self, civitai_client):
self.client = civitai_client
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
return await self.client.get_model_by_hash(model_hash)
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
return await self.client.get_model_versions(model_id)
async def get_model_versions_bulk(
self, model_ids: Sequence[int]
) -> Optional[Dict[int, Dict]]:
) -> Optional[Dict[int, Dict[str, Any]]]:
return await self.client.get_model_versions_bulk(model_ids)
async def get_model_versions_by_hashes(
self, hashes: List[str]
) -> Optional[List[Dict]]:
) -> Optional[List[Dict[str, Any]]]:
return await self.client.get_model_versions_by_hashes(hashes)
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
return await self.client.get_model_version(model_id, version_id)
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[str, Any]], Optional[str]]:
return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
return await self.client.get_user_models(username, cursor)
async def get_creator_model_count(self, username: str) -> Optional[int]:
@@ -195,19 +195,19 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
def __init__(self, civarchive_client):
self.client = civarchive_client
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
return await self.client.get_model_by_hash(model_hash)
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
return await self.client.get_model_versions(model_id)
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
return await self.client.get_model_version(model_id, version_id)
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[str, Any]], Optional[str]]:
return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
"""Not supported by CivArchive provider"""
return None
@@ -218,7 +218,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
self.db_path = db_path
self._aiosqlite = _require_aiosqlite()
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Find model by hash value from SQLite database"""
async with self._aiosqlite.connect(self.db_path) as db:
# Look up in model_files table to get model_id and version_id
@@ -243,7 +243,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
result = await self._get_version_with_model_data(db, model_id, version_id)
return result, None if result else "Error retrieving model data"
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
"""Get all versions of a model from SQLite database"""
async with self._aiosqlite.connect(self.db_path) as db:
db.row_factory = self._aiosqlite.Row
@@ -299,7 +299,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
'name': model_name
}
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
"""Get specific model version with additional metadata from SQLite database"""
if not model_id and not version_id:
return None
@@ -339,7 +339,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
# Now we have both model_id and version_id, get the full data
return await self._get_version_with_model_data(db, model_id, version_id)
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[str, Any]], Optional[str]]:
"""Fetch model version metadata from SQLite database"""
async with self._aiosqlite.connect(self.db_path) as db:
db.row_factory = self._aiosqlite.Row
@@ -358,11 +358,11 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
version_data = await self._get_version_with_model_data(db, model_id, version_id)
return version_data, None
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
"""Listing models by username is not supported for archive database"""
return None
async def _get_version_with_model_data(self, db, model_id, version_id) -> Optional[Dict]:
async def _get_version_with_model_data(self, db, model_id, version_id) -> Optional[Dict[str, Any]]:
"""Helper to build version data with model information"""
# Get version details
version_query = "SELECT name, base_model, data FROM model_versions WHERE id = ? AND model_id = ?"
@@ -485,7 +485,7 @@ class FallbackMetadataProvider(ModelMetadataProvider):
jitter_ratio=self._rate_limit_jitter_ratio,
)
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
for provider, label in self._iter_providers():
try:
result, error = await self._call_with_rate_limit(
@@ -507,7 +507,7 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue
return None, "Model not found"
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
not_found_confirmed = False
for provider, label in self._iter_providers():
try:
@@ -538,7 +538,7 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue
return None
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
for provider, label in self._iter_providers():
try:
result = await self._call_with_rate_limit(
@@ -561,7 +561,7 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue
return None
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[str, Any]], Optional[str]]:
for provider, label in self._iter_providers():
try:
result, error = await self._call_with_rate_limit(
@@ -585,7 +585,7 @@ class FallbackMetadataProvider(ModelMetadataProvider):
async def get_model_versions_by_hashes(
self, hashes: List[str]
) -> Optional[List[Dict]]:
) -> Optional[List[Dict[str, Any]]]:
for provider, label in self._iter_providers():
try:
result = await self._call_with_rate_limit(
@@ -613,7 +613,7 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue
return None
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
for provider, label in self._iter_providers():
try:
result = await self._call_with_rate_limit(
@@ -681,14 +681,14 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
def __getattr__(self, item):
return getattr(self._provider, item)
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_model_by_hash,
model_hash,
)
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_model_versions,
@@ -698,7 +698,7 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
async def get_model_versions_bulk(
self,
model_ids: Sequence[int],
) -> Optional[Dict[int, Dict]]:
) -> Optional[Dict[int, Dict[str, Any]]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_model_versions_bulk,
@@ -707,14 +707,14 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
async def get_model_versions_by_hashes(
self, hashes: List[str]
) -> Optional[List[Dict]]:
) -> Optional[List[Dict[str, Any]]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_model_versions_by_hashes,
hashes,
)
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_model_version,
@@ -722,14 +722,14 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
version_id,
)
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[str, Any]], Optional[str]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_model_version_info,
version_id,
)
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
return await self._rate_limit_helper.run(
self._label,
self._provider.get_user_models,
@@ -762,12 +762,12 @@ class ModelMetadataProviderManager:
if is_default or self.default_provider is None:
self.default_provider = name
async def get_model_by_hash(self, model_hash: str, provider_name: str = None) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_by_hash(self, model_hash: str, provider_name: Optional[str] = None) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Find model by hash using specified or default provider"""
provider = self._get_provider(provider_name)
return await provider.get_model_by_hash(model_hash)
async def get_model_versions(self, model_id: str, provider_name: str = None) -> Optional[Dict]:
async def get_model_versions(self, model_id: str, provider_name: Optional[str] = None) -> Optional[Dict[str, Any]]:
"""Get model versions using specified or default provider"""
provider = self._get_provider(provider_name)
return await provider.get_model_versions(model_id)
@@ -775,8 +775,8 @@ class ModelMetadataProviderManager:
async def get_model_versions_bulk(
self,
model_ids: Sequence[int],
provider_name: str = None,
) -> Optional[Dict[int, Dict]]:
provider_name: Optional[str] = None,
) -> Optional[Dict[int, Dict[str, Any]]]:
"""Fetch model versions for multiple model ids when supported by provider."""
provider = self._get_provider(provider_name)
try:
@@ -784,12 +784,12 @@ class ModelMetadataProviderManager:
except NotImplementedError:
return None
async def get_model_version(self, model_id: int = None, version_id: int = None, provider_name: str = None) -> Optional[Dict]:
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None, provider_name: Optional[str] = None) -> Optional[Dict[str, Any]]:
"""Get specific model version using specified or default provider"""
provider = self._get_provider(provider_name)
return await provider.get_model_version(model_id, version_id)
async def get_model_version_info(self, version_id: str, provider_name: str = None) -> Tuple[Optional[Dict], Optional[str]]:
async def get_model_version_info(self, version_id: str, provider_name: Optional[str] = None) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
"""Fetch model version info using specified or default provider"""
provider = self._get_provider(provider_name)
return await provider.get_model_version_info(version_id)
@@ -797,8 +797,8 @@ class ModelMetadataProviderManager:
async def get_model_versions_by_hashes(
self,
hashes: List[str],
provider_name: str = None,
) -> Optional[List[Dict]]:
provider_name: Optional[str] = None,
) -> Optional[List[Dict[str, Any]]]:
provider = self._get_provider(provider_name)
try:
return await provider.get_model_versions_by_hashes(hashes)
@@ -808,19 +808,19 @@ class ModelMetadataProviderManager:
async def get_user_models(
self,
username: str,
provider_name: str = None,
provider_name: Optional[str] = None,
cursor: Optional[str] = None,
) -> Optional[Dict]:
) -> Optional[Dict[str, Any]]:
"""Fetch one page of models owned by the specified user"""
provider = self._get_provider(provider_name)
return await provider.get_user_models(username, cursor)
async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]:
async def get_creator_model_count(self, username: str, provider_name: Optional[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: Optional[str] = None) -> ModelMetadataProvider:
"""Get provider by name or default provider"""
if provider_name:
if provider_name not in self.providers:
+2 -1
View File
@@ -12,6 +12,7 @@ from typing import (
Tuple,
Protocol,
Callable,
cast,
)
from ..utils.constants import NSFW_LEVELS
@@ -309,7 +310,7 @@ class ModelFilterSet:
else:
include_tags.add(normalized)
else:
include_tags = {tag.strip().lower() for tag in tag_filters if tag}
include_tags = {tag.strip().lower() for tag in cast(Iterable[Any], tag_filters) if tag}
if include_tags:
tag_logic = criteria.tag_logic.lower() if criteria.tag_logic else "any"
+249 -35
View File
@@ -5,11 +5,11 @@ import asyncio
import time
import shutil
from dataclasses import dataclass
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Set, Type, Union
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Sequence, Set, Type, Union, cast
from ..utils.models import BaseModelMetadata
from ..utils.models import BaseModelMetadata, autov3_from_civitai_files
from ..config import config
from ..utils.file_utils import find_preview_file, get_preview_extension, calculate_sha256
from ..utils.file_utils import find_preview_file, get_preview_extension, calculate_sha256, calculate_autov3
from ..utils.metadata_manager import MetadataManager
from ..utils.civitai_utils import resolve_license_info
from .model_cache import ModelCache
@@ -29,7 +29,7 @@ logger = logging.getLogger(__name__)
class CacheBuildResult:
"""Represents the outcome of scanning model files for cache building."""
raw_data: List[Dict]
raw_data: List[Dict[str, Any]]
hash_index: ModelHashIndex
tags_count: Dict[str, int]
excluded_models: List[str]
@@ -59,7 +59,7 @@ class ModelScanner:
lock = cls._get_lock()
async with lock:
if cls not in cls._instances:
cls._instances[cls] = cls()
cls._instances[cls] = cls() # pyright: ignore[reportCallIssue]
return cls._instances[cls]
def __init__(self, model_type: str, model_class: Type[BaseModelMetadata], file_extensions: Set[str], hash_index: Optional[ModelHashIndex] = None):
@@ -78,7 +78,8 @@ class ModelScanner:
self.model_type = model_type
self.model_class = model_class
self.file_extensions = file_extensions
self._cache = None
self._cache: Any = None
self._cache_version: int = 0
self._hash_index = hash_index or ModelHashIndex()
self._tags_count = {} # Dictionary to store tag counts
self._is_initializing = False # Flag to track initialization state
@@ -86,6 +87,7 @@ class ModelScanner:
self._persistent_cache = get_persistent_cache()
self._name_display_mode = self._resolve_name_display_mode()
self._cancel_requested = False # Flag for cancellation
self._autov3_backfill_scheduled = False # One-time AutoV3 backfill trigger per process
try:
loop = asyncio.get_running_loop()
except RuntimeError:
@@ -97,6 +99,25 @@ class ModelScanner:
# Register this service
asyncio.create_task(self._register_service())
@property
def cache_version(self) -> int:
"""Monotonic version counter for the in-memory cache.
Every write path that mutates scanner cache state calls
:meth:`bump_cache_version`, so consumers (e.g. RecipeScanner) can
detect when a cached derivation of the raw data is stale. Reads never
bump.
"""
return self._cache_version
def bump_cache_version(self) -> None:
"""Invalidate derived caches by incrementing the cache version.
Public because external services (model lifecycle, route handlers)
rewrite scanner raw_data directly and must be able to invalidate it.
"""
self._cache_version += 1
def on_library_changed(self) -> None:
"""Reset caches when the active library changes."""
self._persistent_cache = get_persistent_cache()
@@ -106,6 +127,7 @@ class ModelScanner:
self._excluded_models = []
self._is_initializing = False
self._name_display_mode = self._resolve_name_display_mode()
self.bump_cache_version()
try:
loop = asyncio.get_running_loop()
@@ -182,7 +204,7 @@ class ModelScanner:
is_mapping = isinstance(source, Mapping)
def get_value(key: str, default: Any = None) -> Any:
if is_mapping:
if isinstance(source, Mapping):
return source.get(key, default)
sentinel = object()
@@ -225,6 +247,19 @@ class ModelScanner:
if not isinstance(notes, str):
notes = str(notes)
# AutoV3 three-state contract: absent key / None = "not checked yet",
# "" = "checked but unavailable" (never re-read the header), else the
# 12-char lowercase hex value. A metadata object already follows the
# contract and is passed through unchanged; a payload dict only carries
# an explicit checked state when the key is present.
if is_mapping:
if 'autov3' in source:
entry_autov3 = source['autov3'] or ''
else:
entry_autov3 = None
else:
entry_autov3 = get_value('autov3', None)
entry: Dict[str, Any] = {
'file_path': normalized_path,
# file_name is always stored WITHOUT extension (e.g. "OWSMianne_ANIMA_V1",
@@ -238,6 +273,7 @@ class ModelScanner:
'size': int(get_value('size', 0) or 0),
'modified': float(get_value('modified', 0.0) or 0.0),
'sha256': (get_value('sha256', '') or '').lower(),
'autov3': entry_autov3,
'base_model': get_value('base_model', '') or '',
'preview_url': preview_url,
'preview_nsfw_level': int(get_value('preview_nsfw_level', 0) or 0),
@@ -473,6 +509,13 @@ class ModelScanner:
if sha_value and path:
hash_index.add_entry(sha_value.lower(), path)
# Rebuild the AutoV3 index from the persisted autov3_index rows. These
# cover every known autov3 -> path mapping regardless of whether a
# sha256 row also exists for the same file.
for autov3_value, path in persisted.autov3_hash_rows:
if autov3_value and path:
hash_index.add_autov3(autov3_value.lower(), path)
tags_count: Dict[str, int] = {}
adjusted_raw_data: List[Dict[str, Any]] = []
for item in persisted.raw_data:
@@ -541,8 +584,30 @@ class ModelScanner:
'scanner_type': self.model_type,
'pageType': page_type
})
# Schedule the one-time AutoV3 backfill task (at most once per process)
# so entries loaded from a persisted snapshot that predates autov3 get
# their checked state computed in the background. The task never blocks
# or crashes the load path.
if not self._autov3_backfill_scheduled:
self._autov3_backfill_scheduled = True
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop is not None:
loop.create_task(self._run_autov3_backfill())
return True
async def _run_autov3_backfill(self) -> None:
"""Backfill autov3 for entries loaded from the persisted cache that lack it."""
try:
from ..services.autov3_backfill_service import Autov3BackfillService # lazy import (module created by another unit)
await Autov3BackfillService.get_instance().backfill(self)
except Exception as exc:
logger.warning("AutoV3 backfill failed: %s", exc)
async def _save_persistent_cache(self, scan_result: CacheBuildResult) -> None:
if not scan_result or not getattr(self, '_persistent_cache', None):
return
@@ -555,6 +620,7 @@ class ModelScanner:
return
hash_snapshot = self._build_hash_index_snapshot(scan_result.hash_index)
autov3_snapshot = self._build_autov3_index_snapshot(scan_result.hash_index)
loop = asyncio.get_event_loop()
try:
await loop.run_in_executor(
@@ -563,7 +629,8 @@ class ModelScanner:
self.model_type,
list(scan_result.raw_data),
hash_snapshot,
list(scan_result.excluded_models)
list(scan_result.excluded_models),
autov3_snapshot,
)
except Exception as exc:
logger.warning("%s Scanner: Failed to persist cache: %s", self.model_type.capitalize(), exc)
@@ -589,6 +656,20 @@ class ModelScanner:
bucket.append(path)
return snapshot
def _build_autov3_index_snapshot(self, hash_index: Optional[ModelHashIndex]) -> Dict[str, List[str]]:
"""Build the autov3 -> [paths] snapshot for the persisted cache."""
snapshot: Dict[str, List[str]] = {}
if not hash_index:
return snapshot
for autov3_value, path in hash_index.get_all_autov3().items():
if not autov3_value or not path:
continue
bucket = snapshot.setdefault(autov3_value.lower(), [])
if path not in bucket:
bucket.append(path)
return snapshot
async def _persist_current_cache(self) -> None:
if self._cache is None or not getattr(self, '_persistent_cache', None):
return
@@ -712,7 +793,7 @@ class ModelScanner:
else:
await self._reconcile_cache()
return self._cache
return cast(ModelCache, self._cache)
async def _initialize_cache(self) -> None:
"""Initialize or refresh the cache"""
@@ -872,6 +953,8 @@ class ModelScanner:
)
continue
model_data = validation_result.entry
if model_data is None:
continue
self._ensure_license_flags(model_data)
# Add to cache
@@ -880,7 +963,11 @@ class ModelScanner:
# Update hash index if available
if 'sha256' in model_data and 'file_path' in model_data:
self._hash_index.add_entry(model_data['sha256'].lower(), model_data['file_path'])
self._hash_index.add_entry(
model_data['sha256'].lower(),
model_data['file_path'],
model_data.get('autov3') or None
)
# Update tags count
if 'tags' in model_data and model_data['tags']:
@@ -928,8 +1015,8 @@ class ModelScanner:
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files]
dedup_removed = 0
seen_paths: set = set()
deduped: list = []
seen_paths: set[str] = set()
deduped: list[Dict[str, Any]] = []
for item in reversed(self._cache.raw_data):
path = item.get('file_path', '')
if path not in seen_paths:
@@ -964,6 +1051,7 @@ class ModelScanner:
logger.error(f"{self.model_type.capitalize()} Scanner: Error reconciling cache: {e}", exc_info=True)
finally:
self._is_initializing = False # Unset flag
self.bump_cache_version()
def is_initializing(self) -> bool:
"""Check if the scanner is currently initializing"""
@@ -1044,7 +1132,7 @@ class ModelScanner:
*,
hash_index: Optional[ModelHashIndex] = None,
excluded_models: Optional[List[str]] = None
) -> Dict:
) -> Optional[Dict[str, Any]]:
"""Process a single model file and return its metadata"""
hash_index = hash_index or self._hash_index
excluded_models = excluded_models if excluded_models is not None else self._excluded_models
@@ -1068,7 +1156,7 @@ class ModelScanner:
file_name = os.path.splitext(os.path.basename(file_path))[0]
file_info['name'] = file_name
metadata = self.model_class.from_civitai_info(version_info, file_info, file_path)
metadata = cast(Any, self.model_class).from_civitai_info(version_info, file_info, file_path)
metadata.preview_url = find_preview_file(file_name, os.path.dirname(file_path))
await MetadataManager.save_metadata(file_path, metadata)
logger.info(f"Created metadata from .civitai.info for {file_path} (Reason: .civitai.info was found but .metadata.json was missing)")
@@ -1105,6 +1193,8 @@ class ModelScanner:
if metadata is None:
metadata = await self._create_default_metadata(file_path)
assert metadata is not None
# Hook: allow subclasses to adjust metadata
metadata = self.adjust_metadata(metadata, file_path, root_path)
@@ -1130,6 +1220,36 @@ class ModelScanner:
except Exception as e:
logger.error(f"Failed to compute SHA256 for {file_path}: {e}")
# AutoV3 resolution: prefer the Civitai AutoV3 reported for the file
# whose SHA256 matches (authoritative for recipe matching), falling
# back to the embedded safetensors header hash only for models never
# checked before (autov3 is None). A checked-unavailable state ('')
# is only upgraded by Civitai data — the header is never re-read.
current_autov3 = model_data.get('autov3')
if current_autov3 in (None, ''):
try:
civitai_data = None
if isinstance(metadata, BaseModelMetadata):
civitai_data = metadata.civitai
elif isinstance(metadata, dict):
civitai_data = metadata.get("civitai")
autov3 = autov3_from_civitai_files(
civitai_data, model_data.get("sha256") or ""
) or ""
if not autov3 and current_autov3 is None:
autov3 = (calculate_autov3(os.path.realpath(file_path)) or '').lower()
if autov3 != current_autov3:
model_data['autov3'] = autov3
if isinstance(metadata, BaseModelMetadata):
metadata.autov3 = autov3
await MetadataManager.save_metadata(file_path, metadata)
elif isinstance(metadata, dict):
# Dict payload: JSON null encodes the checked-unavailable state.
metadata['autov3'] = autov3 or None
await MetadataManager.save_metadata(file_path, metadata)
except Exception as e:
logger.error(f"Failed to resolve AutoV3 for {file_path}: {e}")
# Skip excluded models
if model_data.get('exclude', False):
excluded_models.append(model_data['file_path'])
@@ -1169,6 +1289,8 @@ class ModelScanner:
self._log_duplicate_filename_summary()
self.bump_cache_version()
def _log_duplicate_filename_summary(self) -> None:
"""Log a batched summary of duplicate filename conflicts once per scan."""
# Duplicate filename detection is only relevant for LoRAs, which use
@@ -1202,7 +1324,7 @@ class ModelScanner:
async def _sync_download_history(
self,
raw_data: List[Mapping[str, Any]],
raw_data: Sequence[Mapping[str, Any]],
*,
source: str,
) -> None:
@@ -1251,7 +1373,7 @@ class ModelScanner:
) -> CacheBuildResult:
"""Collect metadata for all model files."""
raw_data: List[Dict] = []
raw_data: List[Dict[str, Any]] = []
hash_index = ModelHashIndex()
tags_count: Dict[str, int] = {}
excluded_models: List[str] = []
@@ -1315,6 +1437,8 @@ class ModelScanner:
)
continue
result = validation_result.entry
if result is None:
continue
self._ensure_license_flags(result)
raw_data.append(result)
@@ -1322,7 +1446,7 @@ class ModelScanner:
sha_value = result.get('sha256')
model_path = result.get('file_path')
if sha_value and model_path:
hash_index.add_entry(sha_value.lower(), model_path)
hash_index.add_entry(sha_value.lower(), model_path, result.get('autov3') or None)
for tag in result.get('tags') or []:
tags_count[tag] = tags_count.get(tag, 0) + 1
@@ -1354,7 +1478,7 @@ class ModelScanner:
excluded_models=excluded_models
)
async def add_model_to_cache(self, metadata_dict: Dict, folder: str = '') -> bool:
async def add_model_to_cache(self, metadata_dict: Dict[str, Any], folder: str = '') -> bool:
"""Add a model to the cache
Args:
@@ -1367,7 +1491,8 @@ class ModelScanner:
try:
if self._cache is None:
await self.get_cached_data()
assert self._cache is not None
# Update folder in metadata
metadata_dict['folder'] = folder
@@ -1391,14 +1516,19 @@ class ModelScanner:
await self._cache.resort()
# Update the hash index
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
self._hash_index.add_entry(
metadata_dict['sha256'],
metadata_dict['file_path'],
metadata_dict.get('autov3') or None,
)
await self._persist_current_cache()
self.bump_cache_version()
return True
except Exception as e:
logger.error(f"Error adding model to cache: {e}")
return False
async def move_model(self, source_path: str, target_path: str) -> Optional[str]:
async def move_model(self, source_path: str, target_path: str) -> Optional[Dict[str, Any]]:
"""Move a model and its associated files to a new location
Args:
@@ -1432,7 +1562,7 @@ class ModelScanner:
# Check for filename conflicts and auto-rename if necessary
from ..utils.models import BaseModelMetadata
final_filename = BaseModelMetadata.generate_unique_filename(
target_path, base_name, file_ext, get_source_hash
target_path, base_name, file_ext, lambda: get_source_hash() or ""
)
target_file = os.path.join(target_path, final_filename).replace(os.sep, '/')
@@ -1480,7 +1610,7 @@ class ModelScanner:
logger.error(f"Error moving associated file {source_file}: {e}")
# Handle metadata file specially to update paths
if source_metadata and os.path.exists(source_metadata):
if source_metadata and moved_metadata_path and os.path.exists(source_metadata):
try:
shutil.move(source_metadata, moved_metadata_path)
metadata = await self._update_metadata_paths(moved_metadata_path, target_file)
@@ -1498,7 +1628,7 @@ class ModelScanner:
logger.error(f"Error moving model: {e}", exc_info=True)
return None
async def _update_metadata_paths(self, metadata_path: str, model_path: str) -> Dict:
async def _update_metadata_paths(self, metadata_path: str, model_path: str) -> Optional[Dict[str, Any]]:
"""Update file paths in metadata file"""
try:
with open(metadata_path, 'r', encoding='utf-8') as f:
@@ -1524,7 +1654,7 @@ class ModelScanner:
logger.error(f"Error updating metadata paths: {e}", exc_info=True)
return None
async def update_single_model_cache(self, original_path: str, new_path: str, metadata: Dict, recalculate_type: bool = False) -> Union[bool, Dict]:
async def update_single_model_cache(self, original_path: str, new_path: str, metadata: Optional[Dict[str, Any]], recalculate_type: bool = False) -> Union[bool, Dict[str, Any]]:
"""Update cache after a model has been moved or modified"""
cache = await self.get_cached_data()
@@ -1547,6 +1677,7 @@ class ModelScanner:
]
cache_modified = bool(existing_item) or bool(metadata)
cache_entry: Optional[Dict[str, Any]] = None
if metadata:
normalized_new_path = new_path.replace(os.sep, '/')
@@ -1578,7 +1709,11 @@ class ModelScanner:
sha_value = cache_entry.get('sha256')
if sha_value:
self._hash_index.add_entry(sha_value.lower(), normalized_new_path)
self._hash_index.add_entry(
sha_value.lower(),
normalized_new_path,
cache_entry.get('autov3') or None,
)
all_folders = set(item['folder'] for item in cache.raw_data)
cache.folders = sorted(list(all_folders), key=lambda x: x.lower())
@@ -1592,8 +1727,11 @@ class ModelScanner:
if cache_modified:
await self._persist_current_cache()
self.bump_cache_version()
return cache_entry if metadata else True
if metadata and cache_entry is not None:
return cache_entry
return True
async def sync_cache_from_metadata(
self, file_path: str, metadata_dict: Dict[str, Any]
@@ -1716,10 +1854,11 @@ class ModelScanner:
# ---- In-place update of the cache entry ----
existing_entry.clear()
existing_entry.update(desired_entry)
self.bump_cache_version()
# ---- Incremental tag count update ----
new_tags: set = set(desired_entry.get("tags") or [])
old_tag_set: set = set(old_tags)
new_tags: set[str] = set(desired_entry.get("tags") or [])
old_tag_set: set[str] = set(old_tags)
for tag in old_tag_set - new_tags:
current = self._tags_count.get(tag, 0)
if current <= 1:
@@ -1736,7 +1875,11 @@ class ModelScanner:
if old_sha:
self._hash_index.remove_by_path(file_path)
if new_sha:
self._hash_index.add_entry(new_sha, file_path)
self._hash_index.add_entry(
new_sha,
file_path,
desired_entry.get('autov3') or None,
)
# ---- Incremental version index update ----
new_civitai = desired_entry.get("civitai")
@@ -1787,6 +1930,75 @@ class ModelScanner:
return True
async def update_autov3_for_model(self, model_type: str, file_path: str, autov3: str) -> bool:
"""Persist an AutoV3 hash for a single model (single write path used by the backfill service).
Locates the in-memory cache entry by ``file_path`` and updates only its
``autov3`` field: the in-memory hash index, the SQLite snapshot via
:meth:`PersistentModelCache.update_single_model`, and the
``.metadata.json`` sidecar. sha256, tags, and every other field are
left untouched, so the persistent delta only ever differs in autov3.
Returns:
``True`` when the entry was found and updated, ``False`` otherwise.
Never raises failures are logged and swallowed.
"""
try:
if self._cache is None:
return False
entry = next(
(item for item in self._cache.raw_data if item.get('file_path') == file_path),
None,
)
if entry is None:
return False
# Normalize once so the memory entry, sidecar, and SQLite row agree.
autov3 = (autov3 or "").lower()
# Capture the pre-mutation state so update_single_model only sees
# an autov3 delta between old and new.
old_item = dict(entry)
entry['autov3'] = autov3 or ''
# Prefer add_entry when a sha256 is known so the sha256 and autov3
# maps stay in sync; fall back to an autov3-only registration.
sha_value = entry.get('sha256')
checked_autov3 = entry.get('autov3') or None
if sha_value:
self._hash_index.add_entry(sha_value.lower(), file_path, checked_autov3)
elif checked_autov3:
self._hash_index.add_autov3(checked_autov3, file_path)
persistent = getattr(self, '_persistent_cache', None)
if persistent is not None:
await asyncio.get_event_loop().run_in_executor(
None,
persistent.update_single_model,
model_type,
entry,
old_item,
)
# Sidecar write-back: JSON null encodes the checked-unavailable
# state. Skip silently when the sidecar does not exist.
metadata_path = f"{os.path.splitext(file_path)[0]}.metadata.json"
if os.path.exists(metadata_path):
with open(metadata_path, 'r', encoding='utf-8') as handle:
payload = json.load(handle)
if not isinstance(payload, dict):
payload = {}
payload['autov3'] = entry['autov3'] or None
await MetadataManager.save_metadata(metadata_path, payload)
self.bump_cache_version()
return True
except Exception as exc:
logger.warning("Failed to update AutoV3 for %s: %s", file_path, exc)
return False
@staticmethod
def _cache_entries_differ(a: Dict[str, Any], b: Dict[str, Any]) -> bool:
"""Return ``True`` when two cache-entry dicts differ in any field.
@@ -1846,7 +2058,7 @@ class ModelScanner:
return None
async def get_top_tags(self, limit: int = 20) -> List[Dict[str, any]]:
async def get_top_tags(self, limit: int = 20) -> List[Dict[str, Any]]:
"""Get top tags sorted by count. If limit is 0, return all tags."""
await self.get_cached_data()
@@ -1862,7 +2074,7 @@ class ModelScanner:
async def search_tags(
self, query: str, limit: int = 50
) -> List[Dict[str, any]]:
) -> List[Dict[str, Any]]:
"""Search tags by case-insensitive substring match, sorted by count.
If query is empty, behaves like get_top_tags (returns top ``limit``
@@ -1885,7 +2097,7 @@ class ModelScanner:
return matched
return matched[:limit]
async def get_base_models(self, limit: int = 20) -> List[Dict[str, any]]:
async def get_base_models(self, limit: int = 20) -> List[Dict[str, Any]]:
"""Get base models sorted by count. If limit is 0, return all."""
cache = await self.get_cached_data()
@@ -1966,7 +2178,7 @@ class ModelScanner:
await self._persist_current_cache()
return updated
async def bulk_delete_models(self, file_paths: List[str]) -> Dict:
async def bulk_delete_models(self, file_paths: List[str]) -> Dict[str, Any]:
"""Delete multiple models and update cache in a batch operation
Args:
@@ -2114,6 +2326,8 @@ class ModelScanner:
await self._persist_current_cache()
self.bump_cache_version()
return True
except Exception as e:
@@ -2164,7 +2378,7 @@ class ModelScanner:
logger.error(f"Error checking model version existence: {e}")
return False
async def get_model_versions_by_id(self, model_id: int) -> List[Dict]:
async def get_model_versions_by_id(self, model_id: int) -> List[Dict[str, Any]]:
"""Get all versions of a model by its ID
Args:
+6 -6
View File
@@ -6,13 +6,13 @@ logger = logging.getLogger(__name__)
class ModelServiceFactory:
"""Factory for managing model services and routes"""
_services: Dict[str, Type] = {}
_routes: Dict[str, Type] = {}
_services: Dict[str, Type[Any]] = {}
_routes: Dict[str, Type[Any]] = {}
_initialized_services: Dict[str, Any] = {}
_initialized_routes: Dict[str, Any] = {}
@classmethod
def register_model_type(cls, model_type: str, service_class: Type, route_class: Type):
def register_model_type(cls, model_type: str, service_class: Type[Any], route_class: Type[Any]):
"""Register a new model type with its service and route classes
Args:
@@ -24,7 +24,7 @@ class ModelServiceFactory:
cls._routes[model_type] = route_class
@classmethod
def get_service_class(cls, model_type: str) -> Type:
def get_service_class(cls, model_type: str) -> Type[Any]:
"""Get service class for a model type
Args:
@@ -41,7 +41,7 @@ class ModelServiceFactory:
return cls._services[model_type]
@classmethod
def get_route_class(cls, model_type: str) -> Type:
def get_route_class(cls, model_type: str) -> Type[Any]:
"""Get route class for a model type
Args:
@@ -87,7 +87,7 @@ class ModelServiceFactory:
logger.error(f"Failed to setup routes for {model_type}: {e}", exc_info=True)
@classmethod
def get_registered_types(cls) -> list:
def get_registered_types(cls) -> list[str]:
"""Get list of all registered model types
Returns:
+29 -24
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
"""Service for tracking remote model version updates."""
from __future__ import annotations
@@ -336,9 +340,9 @@ class ModelUpdateService:
return
try:
from .persistent_model_cache import get_persistent_cache
from .persistent_model_cache import PersistentModelCache
legacy_path = get_persistent_cache(self._library_name).get_database_path()
legacy_path = PersistentModelCache.get_default(self._library_name).get_database_path()
except Exception:
return
@@ -735,7 +739,7 @@ class ModelUpdateService:
)
results: Dict[int, ModelUpdateRecord] = {}
prefetched: Dict[int, Mapping] = {}
prefetched: Dict[int, Mapping[Any, Any]] = {}
fetch_targets: List[int] = []
if metadata_provider and local_versions:
@@ -834,7 +838,7 @@ class ModelUpdateService:
model_id: int,
version_ids: Sequence[int],
*,
version_info: Optional[Mapping] = None,
version_info: Optional[Mapping[str, Any]] = None,
) -> ModelUpdateRecord:
"""Persist a new set of in-library version identifiers."""
@@ -954,7 +958,11 @@ class ModelUpdateService:
records = self._get_records_bulk(model_type, normalized_ids)
return {
model_id: records.get(model_id).has_update(hide_early_access=hide_early_access) if records.get(model_id) else False
model_id: (
records[model_id].has_update(hide_early_access=hide_early_access)
if model_id in records
else False
)
for model_id in normalized_ids
}
@@ -980,7 +988,7 @@ class ModelUpdateService:
metadata_provider,
*,
force_refresh: bool = False,
prefetched_response: Optional[Mapping] = None,
prefetched_response: Optional[Mapping[str, Any]] = None,
all_local_version_ids: Optional[Sequence[int]] = None,
) -> Optional[ModelUpdateRecord]:
normalized_local = self._normalize_sequence(local_versions)
@@ -1010,7 +1018,7 @@ class ModelUpdateService:
fallback_attempted = False
fallback_error_message: Optional[str] = None
mark_model_as_ignored = False
response: Optional[Mapping] = None
response: Optional[Mapping[str, Any]] = None
if metadata_provider and should_fetch:
response = prefetched_response
if response is None:
@@ -1122,7 +1130,7 @@ class ModelUpdateService:
async def _enrich_version_entries(
self,
metadata_provider,
responses_by_model_id: Dict[int, Mapping],
responses_by_model_id: Dict[int, Mapping[Any, Any]],
) -> None:
"""Enrich version entries with ``usageControl`` via batch hash endpoint.
@@ -1151,7 +1159,7 @@ class ModelUpdateService:
all_hashes = list(version_ids_by_hash.keys())
BATCH_SIZE = 100
enrichment: Dict[int, Dict] = {}
enrichment: Dict[int, Dict[str, Any]] = {}
try:
for start in range(0, len(all_hashes), BATCH_SIZE):
batch = all_hashes[start : start + BATCH_SIZE]
@@ -1208,7 +1216,7 @@ class ModelUpdateService:
version["earlyAccessEndsAt"] = extra["earlyAccessEndsAt"]
@staticmethod
def _collect_hashes_from_response(response: Mapping) -> Dict[int, str]:
def _collect_hashes_from_response(response: Mapping[str, Any]) -> Dict[int, str]:
"""Extract ``{version_id: sha256}`` from a model-level API response.
Returns an empty dict if the response structure is unexpected.
@@ -1229,7 +1237,7 @@ class ModelUpdateService:
return result
@staticmethod
def _extract_sha256_from_version_entry(entry: Mapping) -> Optional[str]:
def _extract_sha256_from_version_entry(entry: Mapping[str, Any]) -> Optional[str]:
"""Return the SHA256 hash from the primary model file of a version entry."""
files = entry.get("files")
if not isinstance(files, list):
@@ -1253,22 +1261,19 @@ class ModelUpdateService:
self,
metadata_provider,
model_ids: Sequence[int],
) -> Dict[int, Mapping]:
) -> Dict[int, Mapping[Any, Any]]:
"""Fetch model metadata in batches of up to 100 ids."""
BATCH_SIZE = 100
normalized = self._normalize_sequence(model_ids)
if not normalized:
provider = metadata_provider
if not normalized or provider is None:
return {}
aggregated: Dict[int, Mapping] = {}
aggregated: Dict[int, Mapping[Any, Any]] = {}
total_ids = len(normalized)
total_batches = (total_ids + BATCH_SIZE - 1) // BATCH_SIZE
provider_name = (
metadata_provider.__class__.__name__
if metadata_provider is not None
else "unknown"
)
provider_name = provider.__class__.__name__
for batch_index, start in enumerate(range(0, total_ids, BATCH_SIZE), start=1):
chunk = normalized[start : start + BATCH_SIZE]
logger.info(
@@ -1279,7 +1284,7 @@ class ModelUpdateService:
provider_name,
)
try:
response = await metadata_provider.get_model_versions_bulk(chunk)
response = await provider.get_model_versions_bulk(chunk)
except RateLimitError:
raise
if response is None:
@@ -1356,7 +1361,7 @@ class ModelUpdateService:
model_type: Optional[str] = None,
model_id: Optional[int] = None,
last_checked_at: Optional[float] = None,
version_info: Optional[Mapping] = None,
version_info: Optional[Mapping[str, Any]] = None,
) -> ModelUpdateRecord:
local_set = set(normalized_local)
# When folder-filtering, also consider versions in other folders
@@ -1578,7 +1583,7 @@ class ModelUpdateService:
if not isinstance(files, Iterable):
return None
def parse_size(entry: Mapping) -> Optional[int]:
def parse_size(entry: Mapping[str, Any]) -> Optional[int]:
size_kb = entry.get("sizeKB")
if size_kb is None:
return None
@@ -1664,8 +1669,8 @@ class ModelUpdateService:
return {}
ids = list(model_ids)
status_rows: list = []
version_rows: list = []
status_rows: list[sqlite3.Row] = []
version_rows: list[sqlite3.Row] = []
with self._connect() as conn:
for start in range(0, len(ids), self._SQLITE_MAX_VARIABLES):
+159 -21
View File
@@ -3,8 +3,8 @@ import logging
import os
import sqlite3
import threading
from dataclasses import dataclass
from typing import Dict, List, Mapping, Optional, Sequence, Tuple
from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional, Sequence, Tuple
from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration
@@ -15,9 +15,10 @@ logger = logging.getLogger(__name__)
class PersistedCacheData:
"""Lightweight structure returned by the persistent cache."""
raw_data: List[Dict]
raw_data: List[Dict[str, Any]]
hash_rows: List[Tuple[str, str]]
excluded_models: List[str]
autov3_hash_rows: List[Tuple[str, str]] = field(default_factory=list)
DEFAULT_LICENSE_FLAGS = 127 # 127 (0b1111111) encodes default CivitAI permissions with all commercial modes enabled.
@@ -36,6 +37,7 @@ class PersistentModelCache:
"size",
"modified",
"sha256",
"autov3",
"base_model",
"preview_url",
"preview_nsfw_level",
@@ -68,8 +70,8 @@ class PersistentModelCache:
self._db_path = db_path or self._resolve_default_path(self._library_name)
self._db_lock = threading.Lock()
self._schema_initialized = False
directory = os.path.dirname(self._db_path)
try:
directory = os.path.dirname(self._db_path)
if directory:
os.makedirs(directory, exist_ok=True)
except Exception as exc: # pragma: no cover - defensive guard
@@ -118,6 +120,10 @@ class PersistentModelCache:
"SELECT sha256, file_path FROM hash_index WHERE model_type = ?",
(model_type,),
).fetchall()
autov3_rows = conn.execute(
"SELECT autov3, file_path FROM autov3_index WHERE model_type = ?",
(model_type,),
).fetchall()
excluded = conn.execute(
"SELECT file_path FROM excluded_models WHERE model_type = ?",
(model_type,),
@@ -128,7 +134,7 @@ class PersistentModelCache:
logger.warning("Failed to load persisted cache for %s: %s", model_type, exc)
return None
raw_data: List[Dict] = []
raw_data: List[Dict[str, Any]] = []
for row in rows:
file_path: str = row["file_path"]
trained_words = []
@@ -139,7 +145,7 @@ class PersistentModelCache:
trained_words = []
creator_username = row["civitai_creator_username"]
civitai: Optional[Dict] = None
civitai: Optional[Dict[str, Any]] = None
civitai_has_data = any(
row[col] is not None
for col in ("civitai_id", "civitai_model_id", "civitai_model_type", "civitai_name")
@@ -191,6 +197,8 @@ class PersistentModelCache:
"hash_status": row["hash_status"] or "completed",
"hf_url": row["hf_url"] or "",
}
if row["autov3"] is not None:
item["autov3"] = (row["autov3"] or "").lower()
raw_data.append(item)
hash_pairs = [(entry["sha256"].lower(), entry["file_path"]) for entry in hash_rows if entry["sha256"]]
@@ -201,10 +209,21 @@ class PersistentModelCache:
if sha_value:
hash_pairs.append((sha_value.lower(), item["file_path"]))
excluded_paths = [row["file_path"] for row in excluded]
return PersistedCacheData(raw_data=raw_data, hash_rows=hash_pairs, excluded_models=excluded_paths)
autov3_pairs = [
(entry["autov3"].lower(), entry["file_path"])
for entry in autov3_rows
if entry["autov3"]
]
def save_cache(self, model_type: str, raw_data: Sequence[Dict], hash_index: Dict[str, List[str]], excluded_models: Sequence[str]) -> None:
excluded_paths = [row["file_path"] for row in excluded]
return PersistedCacheData(
raw_data=raw_data,
hash_rows=hash_pairs,
excluded_models=excluded_paths,
autov3_hash_rows=autov3_pairs,
)
def save_cache(self, model_type: str, raw_data: Sequence[Dict[str, Any]], hash_index: Dict[str, List[str]], excluded_models: Sequence[str], autov3_hash_index: Optional[Dict[str, List[str]]] = None) -> None:
if not self.is_enabled():
return
if not self._schema_initialized:
@@ -219,7 +238,7 @@ class PersistentModelCache:
conn.execute("BEGIN")
model_rows = [self._prepare_model_row(model_type, item) for item in raw_data]
model_map: Dict[str, Tuple] = {
model_map: Dict[str, Tuple[Any, ...]] = {
row[1]: row for row in model_rows if row[1] # row[1] is file_path
}
@@ -251,13 +270,17 @@ class PersistentModelCache:
"DELETE FROM hash_index WHERE model_type = ? AND file_path = ?",
to_remove_models,
)
conn.executemany(
"DELETE FROM autov3_index WHERE model_type = ? AND file_path = ?",
to_remove_models,
)
conn.executemany(
"DELETE FROM excluded_models WHERE model_type = ? AND file_path = ?",
to_remove_models,
)
insert_rows: List[Tuple] = []
update_rows: List[Tuple] = []
insert_rows: List[Tuple[Any, ...]] = []
update_rows: List[Tuple[Any, ...]] = []
for file_path, row in model_map.items():
existing = existing_model_map.get(file_path)
@@ -289,11 +312,11 @@ class PersistentModelCache:
"SELECT file_path, tag FROM model_tags WHERE model_type = ?",
(model_type,),
).fetchall()
existing_tags: Dict[str, set] = {}
existing_tags: Dict[str, set[str]] = {}
for row in existing_tags_rows:
existing_tags.setdefault(row["file_path"], set()).add(row["tag"])
new_tags: Dict[str, set] = {}
new_tags: Dict[str, set[str]] = {}
for item in raw_data:
file_path = item.get("file_path")
if not file_path:
@@ -332,14 +355,14 @@ class PersistentModelCache:
"SELECT sha256, file_path FROM hash_index WHERE model_type = ?",
(model_type,),
).fetchall()
existing_hash_map: Dict[str, set] = {}
existing_hash_map: Dict[str, set[str]] = {}
for row in existing_hash_rows:
sha_value = (row["sha256"] or "").lower()
if not sha_value:
continue
existing_hash_map.setdefault(sha_value, set()).add(row["file_path"])
new_hash_map: Dict[str, set] = {}
new_hash_map: Dict[str, set[str]] = {}
for sha_value, paths in hash_index.items():
normalized_sha = (sha_value or "").lower()
if not normalized_sha:
@@ -373,6 +396,52 @@ class PersistentModelCache:
hash_inserts,
)
if autov3_hash_index is not None:
existing_autov3_rows = conn.execute(
"SELECT autov3, file_path FROM autov3_index WHERE model_type = ?",
(model_type,),
).fetchall()
existing_autov3_map: Dict[str, set[str]] = {}
for row in existing_autov3_rows:
autov3_value = (row["autov3"] or "").lower()
if not autov3_value:
continue
existing_autov3_map.setdefault(autov3_value, set()).add(row["file_path"])
new_autov3_map: Dict[str, set[str]] = {}
for autov3_value, paths in autov3_hash_index.items():
normalized_autov3 = (autov3_value or "").lower()
if not normalized_autov3:
continue
bucket = new_autov3_map.setdefault(normalized_autov3, set())
for path in paths:
if path:
bucket.add(path)
autov3_inserts: List[Tuple[str, str, str]] = []
autov3_deletes: List[Tuple[str, str, str]] = []
all_autov3 = set(existing_autov3_map.keys()) | set(new_autov3_map.keys())
for autov3_value in all_autov3:
existing_paths = existing_autov3_map.get(autov3_value, set())
new_paths = new_autov3_map.get(autov3_value, set())
for path in existing_paths - new_paths:
autov3_deletes.append((model_type, autov3_value, path))
for path in new_paths - existing_paths:
autov3_inserts.append((model_type, autov3_value, path))
if autov3_deletes:
conn.executemany(
"DELETE FROM autov3_index WHERE model_type = ? AND autov3 = ? AND file_path = ?",
autov3_deletes,
)
if autov3_inserts:
conn.executemany(
"INSERT OR IGNORE INTO autov3_index (model_type, autov3, file_path) VALUES (?, ?, ?)",
autov3_inserts,
)
existing_excluded_rows = conn.execute(
"SELECT file_path FROM excluded_models WHERE model_type = ?",
(model_type,),
@@ -435,6 +504,7 @@ class PersistentModelCache:
size INTEGER,
modified REAL,
sha256 TEXT,
autov3 TEXT,
base_model TEXT,
preview_url TEXT,
preview_nsfw_level INTEGER,
@@ -472,6 +542,13 @@ class PersistentModelCache:
PRIMARY KEY (model_type, sha256, file_path)
);
CREATE TABLE IF NOT EXISTS autov3_index (
model_type TEXT NOT NULL,
autov3 TEXT NOT NULL,
file_path TEXT NOT NULL,
PRIMARY KEY (model_type, autov3, file_path)
);
CREATE TABLE IF NOT EXISTS excluded_models (
model_type TEXT NOT NULL,
file_path TEXT NOT NULL,
@@ -504,6 +581,7 @@ class PersistentModelCache:
"license_flags": f"INTEGER DEFAULT {DEFAULT_LICENSE_FLAGS}",
"hash_status": "TEXT DEFAULT 'completed'",
"hf_url": "TEXT DEFAULT ''",
"autov3": "TEXT",
}
for column, definition in required_columns.items():
@@ -522,7 +600,7 @@ class PersistentModelCache:
conn.row_factory = sqlite3.Row
return conn
def _prepare_model_row(self, model_type: str, item: Dict) -> Tuple:
def _prepare_model_row(self, model_type: str, item: Dict[str, Any]) -> Tuple[Any, ...]:
civitai = item.get("civitai") or {}
trained_words = civitai.get("trainedWords")
if isinstance(trained_words, str):
@@ -549,6 +627,12 @@ class PersistentModelCache:
if license_flags is None:
license_flags = DEFAULT_LICENSE_FLAGS
autov3_value = item.get("autov3")
if autov3_value is None:
autov3_column = None
else:
autov3_column = (autov3_value or "").lower()
return (
model_type,
item.get("file_path"),
@@ -558,6 +642,7 @@ class PersistentModelCache:
int(item.get("size") or 0),
float(item.get("modified") or 0.0),
(item.get("sha256") or "").lower() or None,
autov3_column,
item.get("base_model") or "",
item.get("preview_url") or "",
int(item.get("preview_nsfw_level") or 0),
@@ -590,8 +675,8 @@ class PersistentModelCache:
def update_single_model(
self,
model_type: str,
new_item: Dict,
old_item: Optional[Dict] = None,
new_item: Dict[str, Any],
old_item: Optional[Dict[str, Any]] = None,
) -> None:
"""Update a single model row in the persistent cache.
@@ -630,8 +715,8 @@ class PersistentModelCache:
conn.execute(self._insert_model_sql(), row)
# --- tags ---
new_tags: set = set(new_item.get("tags") or [])
old_tags: set = set(old_item.get("tags") or []) if old_item else set()
new_tags: set[str] = set(new_item.get("tags") or [])
old_tags: set[str] = set(old_item.get("tags") or []) if old_item else set()
tags_to_delete = old_tags - new_tags
tags_to_insert = new_tags - old_tags
@@ -663,6 +748,25 @@ class PersistentModelCache:
(model_type, new_sha, file_path),
)
# --- autov3_index ---
new_autov3: Optional[str] = new_item.get("autov3")
if new_autov3 is not None:
new_autov3 = (new_autov3 or "").lower()
old_autov3: Optional[str] = (old_item.get("autov3") if old_item else None)
if old_autov3 is not None:
old_autov3 = (old_autov3 or "").lower()
if new_autov3 != old_autov3:
if old_autov3:
conn.execute(
"DELETE FROM autov3_index WHERE model_type = ? AND autov3 = ? AND file_path = ?",
(model_type, old_autov3, file_path),
)
if new_autov3:
conn.execute(
"INSERT OR IGNORE INTO autov3_index (model_type, autov3, file_path) VALUES (?, ?, ?)",
(model_type, new_autov3, file_path),
)
conn.execute("COMMIT")
except Exception:
conn.execute("ROLLBACK")
@@ -676,6 +780,40 @@ class PersistentModelCache:
exc,
)
def get_models_missing_autov3(self, model_type: str) -> List[str]:
"""Return file paths whose models lack an AutoV3 checked state.
Only rows with a completed sha256 and a NULL autov3 column qualify
rows with '' (checked-unavailable) or a value are never returned, so
the backfill query self-terminates.
"""
if not self.is_enabled():
return []
if not self._schema_initialized:
self._initialize_schema()
if not self._schema_initialized:
return []
try:
with self._db_lock:
conn = self._connect(readonly=True)
try:
rows = conn.execute(
"SELECT file_path FROM models "
"WHERE model_type = ? AND autov3 IS NULL "
"AND sha256 IS NOT NULL AND sha256 != ''",
(model_type,),
).fetchall()
finally:
conn.close()
return [row["file_path"] for row in rows]
except Exception as exc:
logger.warning(
"Failed to query models missing autov3 for %s: %s",
model_type,
exc,
)
return []
def _load_tags(self, conn: sqlite3.Connection, model_type: str) -> Dict[str, List[str]]:
tag_rows = conn.execute(
"SELECT file_path, tag FROM model_tags WHERE model_type = ?",
+12 -8
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
"""SQLite-based persistent cache for recipe metadata.
This module provides fast recipe cache persistence using SQLite, enabling
@@ -13,7 +17,7 @@ import os
import sqlite3
import threading
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Set, Tuple
from typing import Any, Dict, List, Optional, Set, Tuple
from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration
@@ -24,7 +28,7 @@ logger = logging.getLogger(__name__)
class PersistedRecipeData:
"""Lightweight structure returned by the persistent recipe cache."""
raw_data: List[Dict]
raw_data: List[Dict[str, Any]]
file_stats: Dict[str, Tuple[float, int]] # json_path -> (mtime, size)
image_id_map: Dict[str, str] = field(default_factory=dict)
"""Precomputed mapping of civitai image_id → recipe_id."""
@@ -63,8 +67,8 @@ class PersistentRecipeCache:
self._db_path = db_path or self._resolve_default_path(self._library_name)
self._db_lock = threading.Lock()
self._schema_initialized = False
directory = os.path.dirname(self._db_path)
try:
directory = os.path.dirname(self._db_path)
if directory:
os.makedirs(directory, exist_ok=True)
except Exception as exc:
@@ -140,7 +144,7 @@ class PersistentRecipeCache:
logger.warning("Failed to load persisted recipe cache: %s", exc)
return None
raw_data: List[Dict] = []
raw_data: List[Dict[str, Any]] = []
file_stats: Dict[str, Tuple[float, int]] = {}
for row in rows:
@@ -162,7 +166,7 @@ class PersistentRecipeCache:
def save_cache(
self,
recipes: List[Dict],
recipes: List[Dict[str, Any]],
json_paths: Optional[Dict[str, str]] = None,
image_id_map: Optional[Dict[str, str]] = None,
) -> None:
@@ -251,7 +255,7 @@ class PersistentRecipeCache:
except Exception:
return {}
def update_recipe(self, recipe: Dict, json_path: Optional[str] = None) -> None:
def update_recipe(self, recipe: Dict[str, Any], json_path: Optional[str] = None) -> None:
"""Update or insert a single recipe in the cache.
Args:
@@ -439,7 +443,7 @@ class PersistentRecipeCache:
conn.row_factory = sqlite3.Row
return conn
def _prepare_recipe_row(self, recipe: Dict, json_path: str) -> Tuple:
def _prepare_recipe_row(self, recipe: Dict[str, Any], json_path: str) -> Tuple[Any, ...]:
"""Convert a recipe dict to a row tuple for SQLite insertion."""
loras = recipe.get("loras")
loras_json = json.dumps(loras) if loras else None
@@ -486,7 +490,7 @@ class PersistentRecipeCache:
tags_json,
)
def _row_to_recipe(self, row: sqlite3.Row) -> Dict:
def _row_to_recipe(self, row: sqlite3.Row) -> Dict[str, Any]:
"""Convert a SQLite row to a recipe dictionary."""
loras = []
if row["loras_json"]:
+3 -1
View File
@@ -22,7 +22,7 @@ class PreviewAssetService:
self,
*,
metadata_manager,
downloader_factory: Callable[[], Awaitable],
downloader_factory: Callable[[], Awaitable[Any]],
exif_utils,
) -> None:
self._metadata_manager = metadata_manager
@@ -69,6 +69,8 @@ class PreviewAssetService:
if not preview_url:
return
preview_url = str(preview_url)
def extension_from_url(url: str, fallback: str) -> str:
try:
parsed = urlparse(url)
+13 -12
View File
@@ -1,5 +1,5 @@
import asyncio
from typing import Iterable, List, Dict, Optional
from typing import Any, Iterable, List, Dict, Optional
from dataclasses import dataclass, field
from natsort import natsorted
@@ -8,12 +8,13 @@ from natsort import natsorted
class RecipeCache:
"""Cache structure for Recipe data"""
raw_data: List[Dict]
sorted_by_name: List[Dict]
sorted_by_date: List[Dict]
raw_data: List[Dict[str, Any]]
sorted_by_name: List[Dict[str, Any]]
sorted_by_date: List[Dict[str, Any]]
folders: List[str] | None = None
folder_tree: Dict | None = None
folder_tree: Dict[str, Any] | None = None
image_id_map: Dict[str, str] = field(default_factory=dict)
_lock: Any = field(init=False, repr=False, default=None)
"""Mapping of civitai image_id → recipe_id, precomputed at cache build time.
Built once during cache initialization (O(n)) so that
@@ -40,7 +41,7 @@ class RecipeCache:
)
async def update_recipe_metadata(
self, recipe_id: str, metadata: Dict, *, resort: bool = True
self, recipe_id: str, metadata: Dict[str, Any], *, resort: bool = True
) -> bool:
"""Update metadata for a specific recipe in all cached data
@@ -60,7 +61,7 @@ class RecipeCache:
return True
return False # Recipe not found
async def add_recipe(self, recipe_data: Dict, *, resort: bool = False) -> None:
async def add_recipe(self, recipe_data: Dict[str, Any], *, resort: bool = False) -> None:
"""Add a new recipe to the cache."""
async with self._lock:
@@ -70,7 +71,7 @@ class RecipeCache:
async def remove_recipe(
self, recipe_id: str, *, resort: bool = False
) -> Optional[Dict]:
) -> Optional[Dict[str, Any]]:
"""Remove a recipe from the cache by ID.
Args:
@@ -91,7 +92,7 @@ class RecipeCache:
async def bulk_remove(
self, recipe_ids: Iterable[str], *, resort: bool = False
) -> List[Dict]:
) -> List[Dict[str, Any]]:
"""Remove multiple recipes from the cache."""
id_set = {str(recipe_id) for recipe_id in recipe_ids}
@@ -111,7 +112,7 @@ class RecipeCache:
return removed
async def replace_recipe(
self, recipe_id: str, new_data: Dict, *, resort: bool = False
self, recipe_id: str, new_data: Dict[str, Any], *, resort: bool = False
) -> bool:
"""Replace cached data for a recipe."""
@@ -124,7 +125,7 @@ class RecipeCache:
return True
return False
async def get_recipe(self, recipe_id: str) -> Optional[Dict]:
async def get_recipe(self, recipe_id: str) -> Optional[Dict[str, Any]]:
"""Return a shallow copy of a cached recipe."""
async with self._lock:
@@ -133,7 +134,7 @@ class RecipeCache:
return dict(recipe)
return None
async def snapshot(self) -> List[Dict]:
async def snapshot(self) -> List[Dict[str, Any]]:
"""Return a copy of all cached recipes."""
async with self._lock:
+2 -2
View File
@@ -58,8 +58,8 @@ class RecipeFTSIndex:
self._warned_not_ready = False
# Ensure directory exists
directory = os.path.dirname(self._db_path)
try:
directory = os.path.dirname(self._db_path)
if directory:
os.makedirs(directory, exist_ok=True)
except Exception as exc:
@@ -509,7 +509,7 @@ class RecipeFTSIndex:
(recipe_id,)
)
def _prepare_fts_row(self, recipe: Dict[str, Any]) -> tuple:
def _prepare_fts_row(self, recipe: Dict[str, Any]) -> tuple[str, str, str, str, str, str, str]:
"""Prepare a row tuple for FTS insertion."""
recipe_id = str(recipe.get('id', ''))
title = str(recipe.get('title', ''))
+142 -63
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
from __future__ import annotations
import asyncio
@@ -5,28 +9,20 @@ import json
import logging
import os
import time
from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Tuple
from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Tuple, Union, cast
from ..config import config
from .recipe_cache import RecipeCache
from .recipe_fts_index import RecipeFTSIndex
from .persistent_recipe_cache import (
PersistentRecipeCache,
get_persistent_recipe_cache,
PersistedRecipeData,
)
from .service_registry import ServiceRegistry
from .lora_scanner import LoraScanner
from .metadata_service import get_default_metadata_provider
from .checkpoint_scanner import CheckpointScanner
from .settings_manager import get_settings_manager
from .recipes.errors import RecipeNotFoundError
from ..utils.civitai_utils import extract_civitai_image_id
from ..utils.utils import calculate_recipe_fingerprint
from natsort import natsorted
import sys
import re
from ..recipes.merger import GenParamsMerger
from ..recipes.enrichment import RecipeEnricher
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from .lora_scanner import LoraScanner
from .checkpoint_scanner import CheckpointScanner
from .recipe_fts_index import RecipeFTSIndex
from .persistent_recipe_cache import PersistentRecipeCache, PersistedRecipeData
logger = logging.getLogger(__name__)
@@ -48,8 +44,12 @@ class RecipeScanner:
if cls._instance is None:
if not lora_scanner:
# Get lora scanner from service registry if not provided
from .service_registry import ServiceRegistry
lora_scanner = await ServiceRegistry.get_lora_scanner()
if not checkpoint_scanner:
from .service_registry import ServiceRegistry
checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner()
cls._instance = cls(lora_scanner, checkpoint_scanner)
return cls._instance
@@ -77,24 +77,74 @@ class RecipeScanner:
if not hasattr(self, "_initialized"):
self._cache: Optional[RecipeCache] = None
self._initialization_lock = asyncio.Lock()
self._initialization_task: Optional[asyncio.Task] = None
self._initialization_task: Optional[asyncio.Task[Any]] = None
self._is_initializing = False
self._mutation_lock = asyncio.Lock()
self._post_scan_task: Optional[asyncio.Task] = None
self._resort_tasks: Set[asyncio.Task] = set()
self._post_scan_task: Optional[asyncio.Task[Any]] = None
self._resort_tasks: Set[asyncio.Task[Any]] = set()
self._cancel_requested = False
# FTS index for fast search
self._fts_index: Optional[RecipeFTSIndex] = None
self._fts_index_task: Optional[asyncio.Task] = None
self._fts_index_task: Optional[asyncio.Task[Any]] = None
# Persistent cache for fast startup
self._persistent_cache: Optional[PersistentRecipeCache] = None
self._civitai_client: Any = None # Lazily initialized from registry
self._json_path_map: Dict[str, str] = {} # recipe_id -> json_path
if lora_scanner:
self._lora_scanner = lora_scanner
if checkpoint_scanner:
self._checkpoint_scanner = checkpoint_scanner
# Local hash cache (sha256 / autov2 / stored autov3 -> cache item),
# rebuilt only when either model scanner's cache_version changes.
self._local_hash_cache: dict[str, dict[str, Any]] | None = None
self._local_hash_cache_versions: tuple[int, int] | None = None
self._local_hash_cache_lock = asyncio.Lock()
self._initialized = True
async def build_local_hash_cache(self) -> dict[str, dict[str, Any]]:
"""Build a version-cached map of local model hashes to cache items.
Keys are the lowercase full sha256, the first 10 chars of the sha256
(autov2), and the stored lowercase autov3 value when present. An empty
autov3 is the "checked but unavailable" state and never produces a key.
Items without a sha256 are skipped. The dict is reused while both
scanners' cache_version values are unchanged; concurrent callers share
a single build via the lock.
"""
async with self._local_hash_cache_lock:
lora_scanner = self._lora_scanner
checkpoint_scanner = self._checkpoint_scanner
versions = (
lora_scanner.cache_version if lora_scanner is not None else 0,
checkpoint_scanner.cache_version
if checkpoint_scanner is not None
else 0,
)
if (
self._local_hash_cache is not None
and self._local_hash_cache_versions == versions
):
return self._local_hash_cache
cache: dict[str, dict[str, Any]] = {}
for scanner in (lora_scanner, checkpoint_scanner):
if scanner is None:
continue
data = await scanner.get_cached_data()
for item in data.raw_data:
sha256 = (item.get("sha256") or "").lower()
if not sha256:
continue
cache[sha256] = item
cache[sha256[:10]] = item
autov3 = (item.get("autov3") or "").lower()
if autov3:
cache[autov3] = item
self._local_hash_cache = cache
self._local_hash_cache_versions = versions
return cache
def on_library_changed(self) -> None:
"""Reset cached state when the active library changes."""
@@ -123,6 +173,8 @@ class RecipeScanner:
# Reset persistent cache instance for new library
self._persistent_cache = None
self._json_path_map = {}
from .persistent_recipe_cache import PersistentRecipeCache
PersistentRecipeCache.clear_instances()
self._cache = None
@@ -140,6 +192,8 @@ class RecipeScanner:
async def _get_civitai_client(self):
"""Lazily initialize CivitaiClient from registry"""
if self._civitai_client is None:
from .service_registry import ServiceRegistry
self._civitai_client = await ServiceRegistry.get_civitai_client()
return self._civitai_client
@@ -157,7 +211,7 @@ class RecipeScanner:
return self._cancel_requested
async def repair_all_recipes(
self, progress_callback: Optional[Callable[[Dict], Any]] = None
self, progress_callback: Optional[Callable[[Dict[str, Any]], Any]] = None
) -> Dict[str, Any]:
"""Repair all recipes by enrichment with Civitai and embedded metadata.
@@ -339,6 +393,8 @@ class RecipeScanner:
# 3. Use Enricher to repair/enrich
try:
from ..recipes.enrichment import RecipeEnricher
updated = await RecipeEnricher.enrich_recipe(recipe, civitai_client)
except Exception as e:
logger.error(f"Error enriching recipe {recipe.get('id')}: {e}")
@@ -490,6 +546,7 @@ class RecipeScanner:
3. Fall back to full directory scan if cache miss or reconciliation fails
4. Persist results for next startup
"""
loop = None
try:
# Ensure cache exists to avoid None reference errors
if self._cache is None:
@@ -507,6 +564,8 @@ class RecipeScanner:
# Initialize persistent cache
if self._persistent_cache is None:
from .persistent_recipe_cache import get_persistent_recipe_cache
self._persistent_cache = get_persistent_recipe_cache()
recipes_dir = self.recipes_dir
@@ -592,13 +651,14 @@ class RecipeScanner:
return self._cache if hasattr(self, "_cache") else None
finally:
# Clean up the event loop
loop.close()
if loop is not None:
loop.close()
def _reconcile_recipe_cache(
self,
persisted: PersistedRecipeData,
recipes_dir: str,
) -> Tuple[List[Dict], bool, Dict[str, str]]:
) -> Tuple[List[Dict[str, Any]], bool, Dict[str, str]]:
"""Reconcile persisted cache with current filesystem state.
Args:
@@ -608,7 +668,7 @@ class RecipeScanner:
Returns:
Tuple of (recipes list, changed flag, json_paths dict).
"""
recipes: List[Dict] = []
recipes: List[Dict[str, Any]] = []
json_paths: Dict[str, str] = {}
changed = False
@@ -625,12 +685,12 @@ class RecipeScanner:
continue
# Build recipe_id -> recipe lookup (O(n) instead of O(n²))
recipe_by_id: Dict[str, Dict] = {
recipe_by_id: Dict[str, Dict[str, Any]] = {
str(r.get("id", "")): r for r in persisted.raw_data if r.get("id")
}
# Build json_path -> recipe lookup from file_stats (O(m))
persisted_by_path: Dict[str, Dict] = {}
persisted_by_path: Dict[str, Dict[str, Any]] = {}
for json_path in persisted.file_stats.keys():
basename = os.path.basename(json_path)
if basename.lower().endswith(".recipe.json"):
@@ -696,7 +756,7 @@ class RecipeScanner:
def _backfill_source_path_if_needed(
self,
recipes: List[Dict],
recipes: List[Dict[str, Any]],
json_paths: Dict[str, str],
) -> bool:
"""Backfill source_path from recipe JSON files if missing from cache.
@@ -724,7 +784,7 @@ class RecipeScanner:
def _full_directory_scan_sync(
self, recipes_dir: str
) -> Tuple[List[Dict], Dict[str, str]]:
) -> Tuple[List[Dict[str, Any]], Dict[str, str]]:
"""Perform a full synchronous directory scan for recipes.
Args:
@@ -733,7 +793,7 @@ class RecipeScanner:
Returns:
Tuple of (recipes list, json_paths dict).
"""
recipes: List[Dict] = []
recipes: List[Dict[str, Any]] = []
json_paths: Dict[str, str] = {}
# Get all recipe JSON files
@@ -756,7 +816,7 @@ class RecipeScanner:
return recipes, json_paths
def _load_recipe_file_sync(self, recipe_path: str) -> Optional[Dict]:
def _load_recipe_file_sync(self, recipe_path: str) -> Optional[Dict[str, Any]]:
"""Load a single recipe file synchronously.
Args:
@@ -835,6 +895,8 @@ class RecipeScanner:
def _sort_cache_sync(self) -> None:
"""Sort cache data synchronously."""
if self._cache is None:
return
try:
# Sort by name
self._cache.sorted_by_name = natsorted(
@@ -868,6 +930,8 @@ class RecipeScanner:
source = recipe.get("source_path")
if not source:
continue
from ..utils.civitai_utils import extract_civitai_image_id
image_id = extract_civitai_image_id(source)
if image_id and image_id not in mapping:
recipe_id = recipe.get("id")
@@ -950,6 +1014,8 @@ class RecipeScanner:
return
try:
from .recipe_fts_index import RecipeFTSIndex
self._fts_index = RecipeFTSIndex()
# Check if existing index is valid
@@ -987,7 +1053,7 @@ class RecipeScanner:
_build_fts(), name="recipe_fts_index_build"
)
def _search_with_fts(self, search: str, search_options: Dict) -> Optional[Set[str]]:
def _search_with_fts(self, search: str, search_options: Dict[str, Any]) -> Optional[Set[str]]:
"""Search recipes using FTS index if available.
Args:
@@ -1002,7 +1068,7 @@ class RecipeScanner:
return None
# Build the set of fields to search based on search_options
fields: Set[str] = set()
fields: Optional[Set[str]] = set()
if search_options.get("title", True):
fields.add("title")
if search_options.get("tags", True):
@@ -1033,12 +1099,12 @@ class RecipeScanner:
return None
def _update_fts_index_for_recipe(
self, recipe: Dict[str, Any], operation: str = "add"
self, recipe: Union[Dict[str, Any], str], operation: str = "add"
) -> None:
"""Update FTS index for a single recipe (add, update, or remove).
Args:
recipe: The recipe dictionary.
recipe: The recipe dictionary, or a recipe ID string for removal.
operation: One of 'add', 'update', or 'remove'.
"""
if not self._fts_index or not self._fts_index.is_ready():
@@ -1053,7 +1119,7 @@ class RecipeScanner:
)
self._fts_index.remove_recipe(recipe_id)
elif operation in ("add", "update"):
self._fts_index.update_recipe(recipe)
self._fts_index.update_recipe(cast(Dict[str, Any], recipe))
except Exception as exc:
logger.debug("Failed to update FTS index for recipe: %s", exc)
@@ -1071,6 +1137,8 @@ class RecipeScanner:
if value in (None, ""):
continue
from ..recipes.merger import GenParamsMerger
normalized_key = GenParamsMerger.NORMALIZATION_MAPPING.get(key, key)
if normalized_key not in GenParamsMerger.ALLOWED_KEYS:
continue
@@ -1130,7 +1198,8 @@ class RecipeScanner:
def _schedule_resort(self, *, name_only: bool = False) -> None:
"""Schedule a background resort of the recipe cache."""
if not self._cache:
cache = self._cache
if not cache:
return
# Keep folder metadata up to date alongside sort order
@@ -1138,7 +1207,7 @@ class RecipeScanner:
async def _resort_wrapper() -> None:
try:
await self._cache.resort(name_only=name_only)
await cache.resort(name_only=name_only)
except Exception as exc: # pragma: no cover - defensive logging
logger.error(
"Recipe Scanner: error resorting cache: %s", exc, exc_info=True
@@ -1164,10 +1233,10 @@ class RecipeScanner:
except Exception:
return ""
def _build_folder_tree(self, folders: list[str]) -> dict:
def _build_folder_tree(self, folders: list[str]) -> Dict[str, Any]:
"""Build a nested folder tree structure from relative folder paths."""
tree: dict[str, dict] = {}
tree: dict[str, Dict[str, Any]] = {}
for folder in folders:
if not folder:
continue
@@ -1208,18 +1277,20 @@ class RecipeScanner:
cache = await self.get_cached_data()
self._update_folder_metadata(cache)
return cache.folders
return cache.folders or []
async def get_folder_tree(self) -> dict:
async def get_folder_tree(self) -> Dict[str, Any]:
"""Return a hierarchical tree of recipe folders for sidebar navigation."""
cache = await self.get_cached_data()
self._update_folder_metadata(cache)
return cache.folder_tree
return cache.folder_tree or {}
@property
def recipes_dir(self) -> str:
"""Get path to recipes directory"""
from .settings_manager import get_settings_manager
custom_recipes_dir = get_settings_manager().get("recipes_path", "")
if isinstance(custom_recipes_dir, str) and custom_recipes_dir.strip():
recipes_dir = os.path.abspath(
@@ -1242,7 +1313,7 @@ class RecipeScanner:
# If cache is already initialized and no refresh is needed, return it immediately
if self._cache is not None and not force_refresh:
self._update_folder_metadata()
return self._cache
return cast(RecipeCache, self._cache)
# If another initialization is already in progress, wait for it to complete
if self._is_initializing and not force_refresh:
@@ -1293,7 +1364,7 @@ class RecipeScanner:
self._schedule_post_scan_enrichment()
self._schedule_fts_index_build()
return self._cache
return cast(RecipeCache, self._cache)
except Exception as e:
logger.error(
@@ -1344,6 +1415,8 @@ class RecipeScanner:
source = recipe_data.get("source_path")
if source:
from ..utils.civitai_utils import extract_civitai_image_id
image_id = extract_civitai_image_id(source)
if image_id:
recipe_id_value = recipe_data.get("id")
@@ -1410,7 +1483,7 @@ class RecipeScanner:
self._persistent_cache.save_image_id_map(cache.image_id_map)
return len(removed)
async def scan_all_recipes(self) -> List[Dict]:
async def scan_all_recipes(self) -> List[Dict[str, Any]]:
"""Scan all recipe JSON files and return metadata"""
recipes = []
recipes_dir = self.recipes_dir
@@ -1436,7 +1509,7 @@ class RecipeScanner:
return recipes
async def _load_recipe_file(self, recipe_path: str) -> Optional[Dict]:
async def _load_recipe_file(self, recipe_path: str) -> Optional[Dict[str, Any]]:
"""Load recipe data from a JSON file"""
try:
with open(recipe_path, "r", encoding="utf-8") as f:
@@ -1517,6 +1590,8 @@ class RecipeScanner:
# Calculate and update fingerprint if missing
if "loras" in recipe_data and "fingerprint" not in recipe_data:
from ..utils.utils import calculate_recipe_fingerprint
fingerprint = calculate_recipe_fingerprint(recipe_data["loras"])
recipe_data["fingerprint"] = fingerprint
@@ -1548,7 +1623,7 @@ class RecipeScanner:
with open(recipe_path, "w", encoding="utf-8") as file_obj:
json.dump(recipe_data, file_obj, indent=4, ensure_ascii=False)
async def _update_lora_information(self, recipe_data: Dict) -> bool:
async def _update_lora_information(self, recipe_data: Dict[str, Any]) -> bool:
"""Update LoRA information with hash and file_name
Returns:
@@ -1575,14 +1650,14 @@ class RecipeScanner:
if isinstance(model_version_id, int) and model_version_id > 0:
# Try to find in lora cache first
hash_from_cache = await self._find_hash_in_lora_cache(
model_version_id
str(model_version_id)
)
if hash_from_cache:
lora["hash"] = hash_from_cache
metadata_updated = True
else:
# If not in cache, fetch from Civitai
result = await self._get_hash_from_civitai(model_version_id)
result = await self._get_hash_from_civitai(str(model_version_id))
if isinstance(result, tuple):
hash_from_civitai, is_deleted = result
if hash_from_civitai:
@@ -1645,14 +1720,16 @@ class RecipeScanner:
logger.error(f"Error finding hash in lora cache: {e}")
return None
async def _get_hash_from_civitai(self, model_version_id: str) -> Optional[str]:
async def _get_hash_from_civitai(self, model_version_id: str) -> Tuple[Optional[str], bool]:
"""Get hash from Civitai API"""
try:
# Get metadata provider instead of civitai client directly
from .metadata_service import get_default_metadata_provider
metadata_provider = await get_default_metadata_provider()
if not metadata_provider:
logger.error("Failed to get metadata provider")
return None
return None, False
version_info, error_msg = await metadata_provider.get_model_version_info(
model_version_id
@@ -1733,7 +1810,7 @@ class RecipeScanner:
return version_index.get(normalized_id)
async def _determine_base_model(self, loras: List[Dict]) -> Optional[str]:
async def _determine_base_model(self, loras: List[Dict[str, Any]]) -> Optional[str]:
"""Determine the most common base model among LoRAs"""
base_models = {}
@@ -1956,11 +2033,11 @@ class RecipeScanner:
page: int,
page_size: int,
sort_by: str = "date",
search: str = None,
filters: dict = None,
search_options: dict = None,
lora_hash: str = None,
checkpoint_hash: str = None,
search: Optional[str] = None,
filters: Optional[Dict[str, Any]] = None,
search_options: Optional[Dict[str, Any]] = None,
lora_hash: Optional[str] = None,
checkpoint_hash: Optional[str] = None,
bypass_filters: bool = True,
folder: str | None = None,
recursive: bool = True,
@@ -2220,7 +2297,7 @@ class RecipeScanner:
return result
async def get_recipe_by_id(self, recipe_id: str) -> dict:
async def get_recipe_by_id(self, recipe_id: str) -> Optional[Dict[str, Any]]:
"""Get a single recipe by ID with all metadata and formatted URLs
Args:
@@ -2312,7 +2389,7 @@ class RecipeScanner:
return self._normalize_recipe_gen_params(recipe_data)
def _format_file_url(self, file_path: str) -> str:
def _format_file_url(self, file_path: Optional[str]) -> str:
"""Format file path as URL for serving in web UI"""
if not file_path:
return "/loras_static/images/no-preview.png"
@@ -2360,7 +2437,7 @@ class RecipeScanner:
return None
async def update_recipe_metadata(self, recipe_id: str, metadata: dict) -> bool:
async def update_recipe_metadata(self, recipe_id: str, metadata: Dict[str, Any]) -> bool:
"""Update recipe metadata (like title and tags) in both file system and cache
Args:
@@ -2465,6 +2542,8 @@ class RecipeScanner:
lora_entry["modelVersionName"] = civitai_info.get("name", "")
lora_entry["modelVersionId"] = civitai_info.get("id")
from ..utils.utils import calculate_recipe_fingerprint
recipe_data["fingerprint"] = calculate_recipe_fingerprint(
recipe_data.get("loras", [])
)
@@ -2696,7 +2775,7 @@ class RecipeScanner:
return file_updated_count, cache_updated_count
async def find_recipes_by_fingerprint(self, fingerprint: str) -> list:
async def find_recipes_by_fingerprint(self, fingerprint: str) -> List[Dict[str, Any]]:
"""Find recipes with a matching fingerprint
Args:
@@ -2727,7 +2806,7 @@ class RecipeScanner:
return matching_recipes
async def find_all_duplicate_recipes(self) -> dict:
async def find_all_duplicate_recipes(self) -> Dict[str, List[Any]]:
"""Find all recipe duplicates based on fingerprints
Returns:
@@ -2753,7 +2832,7 @@ class RecipeScanner:
return duplicate_groups
async def find_duplicate_recipes_by_source(self) -> dict:
async def find_duplicate_recipes_by_source(self) -> Dict[str, List[Any]]:
"""Find all recipe duplicates based on source_path (Civitai image URLs)
Returns:
+17 -3
View File
@@ -101,6 +101,7 @@ class RecipeAnalysisService:
temp_path = None
metadata: Optional[dict[str, Any]] = None
image_info: Optional[dict[str, Any]] = None
is_video = False
extension = ".jpg" # Default
@@ -413,14 +414,27 @@ class RecipeAnalysisService:
error_msg = "This image does not contain any generation metadata (prompt, models, or parameters)"
else:
error_msg = "No parser found for this image"
payload = {"error": error_msg, "loras": []}
payload: dict[str, Any] = {"error": error_msg, "loras": []}
if include_image_base64 and image_path:
payload["image_base64"] = self._encode_file(image_path)
payload["is_video"] = is_video
payload["extension"] = extension
return AnalysisResult(payload)
result = await parser.parse_metadata(metadata, recipe_scanner=recipe_scanner)
# Only the Civitai image parser accepts a local_cache parameter;
# passing it to other parsers would raise TypeError. Lazy import
# mirrors the repo style used in recipe_handlers.
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
if isinstance(parser, CivitaiApiMetadataParser):
local_cache = await recipe_scanner.build_local_hash_cache()
result = await parser.parse_metadata(
metadata, recipe_scanner=recipe_scanner, local_cache=local_cache
)
else:
result = await parser.parse_metadata(
metadata, recipe_scanner=recipe_scanner
)
if include_image_base64 and image_path:
result["image_base64"] = self._encode_file(image_path)
@@ -494,7 +508,7 @@ class RecipeAnalysisService:
getattr(tensor_image, "dtype", None),
)
import torch # type: ignore[import-not-found]
import torch # pyright: ignore[reportMissingImports]
if isinstance(tensor_image, torch.Tensor):
image_np = tensor_image.cpu().numpy()
+6 -2
View File
@@ -9,7 +9,7 @@ import shutil
import time
import uuid
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Optional
from typing import Any, Awaitable, Dict, Iterable, Optional, cast
from ...config import config
from ...recipes.constants import GEN_PARAM_KEYS
@@ -72,6 +72,8 @@ class RecipePersistenceService:
f"Missing required fields: {', '.join(missing_fields)}"
)
assert metadata is not None
resolved_image_bytes = self._resolve_image_bytes(image_bytes, image_base64)
recipes_dir = target_dir or recipe_scanner.recipes_dir
os.makedirs(recipes_dir, exist_ok=True)
@@ -650,7 +652,9 @@ class RecipePersistenceService:
for candidate in candidates:
try:
checkpoint_info = await lookup(candidate)
checkpoint_info = await cast(
Awaitable[Any], lookup(candidate)
)
except Exception as exc:
self._logger.debug(
"Failed to lookup checkpoint %s while saving widget recipe: %s",
+2 -2
View File
@@ -55,7 +55,7 @@ class ServerI18nManager:
logger.warning(f"Locale {locale} not found, using 'en'")
self.current_locale = 'en'
def get_translation(self, key: str, params: Dict[str, Any] = None, **kwargs) -> str:
def get_translation(self, key: str, params: Dict[str, Any] | None = None, **kwargs) -> str:
"""Get translation for a key with optional parameters (supports both dict and keyword args)"""
# Merge kwargs into params for convenience
if params is None:
@@ -100,7 +100,7 @@ class ServerI18nManager:
return value
def get_available_locales(self) -> list:
def get_available_locales(self) -> list[str]:
"""Get list of available locales"""
return list(self.translations.keys())
+4
View File
@@ -1,3 +1,7 @@
# pyright: reportImportCycles=false
# Lazy (function-local) imports still count as static edges in basedpyright's
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
# import cycles. Breaking them would require an architectural refactor.
import asyncio
import logging
from typing import Optional, Dict, Any, TypeVar, Type
+22 -12
View File
@@ -12,6 +12,7 @@ from threading import Lock
from typing import (
Any,
Awaitable,
Coroutine,
Dict,
Iterable,
List,
@@ -92,6 +93,7 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
"mature_blur_level": "R",
"autoplay_on_hover": False,
"display_density": "default",
"recipes_layout": "grid",
"card_info_display": "always",
"include_trigger_words": False,
"compact_mode": False,
@@ -410,7 +412,12 @@ class SettingsManager:
needs_library_bootstrap = not isinstance(libraries, dict) or not libraries
if not needs_library_bootstrap and top_level_has_paths and len(libraries) == 1:
if (
not needs_library_bootstrap
and top_level_has_paths
and isinstance(libraries, Mapping)
and len(libraries) == 1
):
only_library_payload = next(iter(libraries.values()))
if isinstance(only_library_payload, Mapping):
folder_payload = only_library_payload.get("folder_paths")
@@ -454,6 +461,9 @@ class SettingsManager:
):
seed_library_name = target_name
if not isinstance(libraries, dict) or not libraries:
return
sanitized_libraries: Dict[str, Dict[str, Any]] = {}
changed = False
for name, data in libraries.items():
@@ -593,7 +603,7 @@ class SettingsManager:
return payload
def _normalize_folder_paths(
self, folder_paths: Mapping[str, Iterable[str]]
self, folder_paths: Mapping[str, Any]
) -> Dict[str, List[str]]:
normalized: Dict[str, List[str]] = {}
for key, values in folder_paths.items():
@@ -622,7 +632,7 @@ class SettingsManager:
candidate_values = [values]
else:
try:
candidate_values = list(values) # type: ignore[arg-type]
candidate_values = list(values) # pyright: ignore[reportArgumentType]
except TypeError:
continue
@@ -655,7 +665,7 @@ class SettingsManager:
def _validate_folder_paths(
self,
library_name: str,
folder_paths: Mapping[str, Iterable[str]],
folder_paths: Mapping[str, Any],
) -> None:
"""Ensure folder paths do not overlap with other libraries.
@@ -1118,7 +1128,7 @@ class SettingsManager:
return []
if isinstance(value, str):
candidates: Iterable[str] = (
candidates: Iterable[Any] = (
value.replace("\n", ",").replace(";", ",").split(",")
)
elif isinstance(value, Sequence) and not isinstance(
@@ -1166,7 +1176,7 @@ class SettingsManager:
return []
if isinstance(value, str):
candidates: Iterable[str] = (
candidates: Iterable[Any] = (
value.replace("\n", ",").replace(";", ",").split(",")
)
elif isinstance(value, Sequence) and not isinstance(
@@ -1206,7 +1216,7 @@ class SettingsManager:
return []
if isinstance(value, str):
candidates: Iterable[str] = (
candidates: Iterable[Any] = (
value.replace("\n", ",").replace(";", ",").split(",")
)
elif isinstance(value, Sequence) and not isinstance(
@@ -1594,11 +1604,11 @@ class SettingsManager:
if key == "folder_paths" and isinstance(value, Mapping):
active_name = self.get_active_library_name()
self._validate_folder_paths(active_name, value)
self._update_active_library_entry(folder_paths=value) # type: ignore[arg-type]
self._update_active_library_entry(folder_paths=value) # pyright: ignore[reportArgumentType]
elif key == "extra_folder_paths" and isinstance(value, Mapping):
active_name = self.get_active_library_name()
self._validate_folder_paths(active_name, value)
self._update_active_library_entry(extra_folder_paths=value) # type: ignore[arg-type]
self._update_active_library_entry(extra_folder_paths=value) # pyright: ignore[reportArgumentType]
elif key == "default_lora_root":
self._update_active_library_entry(default_lora_root=str(value))
elif key == "default_checkpoint_root":
@@ -1751,12 +1761,12 @@ class SettingsManager:
"""Trigger cache resorting when the model name display preference updates."""
try:
from .service_registry import ServiceRegistry # type: ignore
from .service_registry import ServiceRegistry # pyright: ignore[reportImportCycles]
except Exception: # pragma: no cover - registry optional in some contexts
return
display_mode = value if isinstance(value, str) else "model_name"
pending: List[Tuple[Optional[asyncio.AbstractEventLoop], Awaitable[Any]]] = []
pending: List[Tuple[Optional[asyncio.AbstractEventLoop], Coroutine[Any, Any, Any]]] = []
def _resolve_service_loop(service: Any) -> Optional[asyncio.AbstractEventLoop]:
loop = getattr(service, "loop", None)
@@ -2117,7 +2127,7 @@ class SettingsManager:
logger.debug("Failed to apply library settings to config: %s", exc)
try:
from .service_registry import ServiceRegistry # type: ignore
from .service_registry import ServiceRegistry # pyright: ignore[reportImportCycles]
for service_name in (
"lora_scanner",
+6 -5
View File
@@ -18,7 +18,7 @@ import sqlite3
import threading
import time
from pathlib import Path
from typing import Dict, List, Optional, Set
from typing import Any, Dict, List, Optional, Set
from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration
@@ -87,10 +87,11 @@ class TagFTSIndex:
self._indexing_in_progress = False
self._schema_initialized = False
self._warned_not_ready = False
self._needs_rebuild = False
# Ensure directory exists
directory = os.path.dirname(self._db_path)
try:
directory = os.path.dirname(self._db_path)
if directory:
os.makedirs(directory, exist_ok=True)
except Exception as exc:
@@ -358,7 +359,7 @@ class TagFTSIndex:
finally:
self._indexing_in_progress = False
def _insert_batch(self, conn: sqlite3.Connection, rows: List[tuple]) -> None:
def _insert_batch(self, conn: sqlite3.Connection, rows: List[tuple[str, int, int, str]]) -> None:
"""Insert a batch of rows into the database.
Each row is a tuple of (tag_name, category, post_count, aliases).
@@ -443,7 +444,7 @@ class TagFTSIndex:
categories: Optional[List[int]] = None,
limit: int = 20,
offset: int = 0,
) -> List[Dict]:
) -> List[Dict[str, Any]]:
"""Search tags using FTS5 with prefix matching.
Supports alias search: if the query matches an alias rather than
@@ -530,7 +531,7 @@ class TagFTSIndex:
categories: Optional[List[int]],
limit: int,
offset: int,
) -> tuple[str, list[object]]:
) -> tuple[str, list[int | str]]:
"""Build the SQL statement and params for a tag search."""
# Escape special LIKE characters and add wildcard
query_escaped = (
+2 -1
View File
@@ -28,7 +28,8 @@ class TagUpdateService:
metadata_path = f"{base}.metadata.json"
metadata = await metadata_loader(metadata_path)
existing_tags = list(metadata.get("tags", []))
raw_tags = metadata.get("tags", [])
existing_tags = list(raw_tags) if isinstance(raw_tags, list) else []
existing_lower = [tag.lower() for tag in existing_tags]
tags_added: List[str] = []
@@ -13,9 +13,11 @@ class AutoOrganizeLockProvider(Protocol):
def is_auto_organize_running(self) -> bool:
"""Return ``True`` when an auto-organize operation is in-flight."""
...
async def get_auto_organize_lock(self) -> asyncio.Lock:
"""Return the asyncio lock guarding auto-organize operations."""
...
class AutoOrganizeInProgressError(RuntimeError):

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