diff --git a/js/sidebar/civitaiSearchTab.js b/js/sidebar/civitaiSearchTab.js index 0208c66..9b216c6 100644 --- a/js/sidebar/civitaiSearchTab.js +++ b/js/sidebar/civitaiSearchTab.js @@ -16,6 +16,7 @@ import { import { createButton, BUTTON_VARIANTS } from '../components/buttons.js'; import { loadSidebarStyle } from './sidebarStyles.js'; +import { loadAndCreateHtmlTemplate } from '../utils/htmlTemplateLoader.js'; import { escapeHtml, formatFileSize } from '../reports/reportGenerator.js'; @@ -35,7 +36,7 @@ const SEARCH_CONFIG = { * Creates the Civitai search tab content * @param {HTMLElement} container - Container element for the tab content */ -export function createCivitaiSearchTab(container) { +export async function createCivitaiSearchTab(container) { container.innerHTML = ''; loadSidebarStyle('civitai-search-styles', 'extensions/comfyui_sageutils/sidebar/civitaiSearchTab.css'); @@ -45,22 +46,23 @@ export function createCivitaiSearchTab(container) { searchContainer.className = 'civitai-search-tab'; // Create search form - const searchForm = createSearchForm(); + const searchForm = await createSearchForm(); // Create results container const resultsContainer = document.createElement('div'); resultsContainer.id = 'civitai-results'; resultsContainer.className = 'civitai-results-container'; - + // Initial message - resultsContainer.innerHTML = ` -
-

🔍 Civitai Model Search

-

Search for models on Civitai to discover and download new content.

-

Use the search form above to find models by name, type, or creator.

-
- `; - + const initialStatus = await createSearchStatusPanel({ + icon: '🔍', + title: 'Civitai Model Search', + messageLine1: 'Search for models on Civitai to discover and download new content.', + messageLine2: 'Use the search form above to find models by name, type, or creator.', + note: '' + }); + resultsContainer.appendChild(initialStatus); + searchContainer.appendChild(searchForm); searchContainer.appendChild(resultsContainer); @@ -74,83 +76,68 @@ export function createCivitaiSearchTab(container) { * Creates the search form with filters * @returns {HTMLElement} Search form element */ -function createSearchForm() { - const form = document.createElement('div'); - form.className = 'civitai-search-form'; - - // Search input - const searchGroup = document.createElement('div'); - searchGroup.className = 'civitai-search-group'; - searchGroup.innerHTML = ` - - - `; - - // Type filter - const typeGroup = document.createElement('div'); - typeGroup.className = 'civitai-search-group'; - typeGroup.innerHTML = ` - - - `; - - // Sort filter - const sortGroup = document.createElement('div'); - sortGroup.className = 'civitai-search-group'; - sortGroup.innerHTML = ` - - - `; - - // Search button row - const buttonRow = document.createElement('div'); - buttonRow.className = 'civitai-search-button-row'; - - // Search button +async function createSearchForm() { + const form = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiSearchForm.html' + ); + + const typeFilter = form.querySelector('#civitai-type-filter'); + if (typeFilter) { + SEARCH_CONFIG.SUPPORTED_TYPES.forEach((type) => { + const option = document.createElement('option'); + option.value = type; + option.textContent = type; + typeFilter.appendChild(option); + }); + } + + const sortFilter = form.querySelector('#civitai-sort-filter'); + if (sortFilter) { + SEARCH_CONFIG.SORT_OPTIONS.forEach((sort) => { + const option = document.createElement('option'); + option.value = sort; + option.textContent = sort; + sortFilter.appendChild(option); + }); + } + const searchButton = createButton('Search', { id: 'civitai-search-button', variant: BUTTON_VARIANTS.SUCCESS, className: 'civitai-search-button' }); - - // NSFW toggle - const nsfwToggle = document.createElement('label'); - nsfwToggle.className = 'civitai-nsfw-toggle'; - nsfwToggle.innerHTML = ` - - Include NSFW Results - `; - - // Results per page - const limitGroup = document.createElement('div'); - limitGroup.className = 'civitai-limit-group'; - limitGroup.innerHTML = ` - - - `; - - buttonRow.appendChild(searchButton); - buttonRow.appendChild(nsfwToggle); - buttonRow.appendChild(limitGroup); - - form.appendChild(searchGroup); - form.appendChild(typeGroup); - form.appendChild(sortGroup); - form.appendChild(buttonRow); - + + const buttonSpot = form.querySelector('#civitai-search-button-spot'); + if (buttonSpot) { + buttonSpot.appendChild(searchButton); + } + return form; } +async function createSearchStatusPanel(options = {}) { + const { + icon = '', + title = '', + messageLine1 = '', + messageLine2 = '', + note = '', + extraClass = '' + } = options; + + return await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiSearchStatus.html', + { + icon, + title, + messageLine1, + messageLine2, + note, + extraClass + } + ); +} + /** * Sets up event handlers for search functionality * @param {HTMLElement} searchForm - Search form element @@ -205,14 +192,15 @@ function setupSearchHandlers(searchForm, resultsContainer) { loadingButton.textContent = cursor ? 'Loading...' : 'Searching...'; if (isNewSearch) { - resultsContainer.innerHTML = ` -
-
🔍
-

Civitai Model Search

-

Search for models on Civitai to discover and download new content.

-

Use the search form above to find models by name, type, or creator.

-
- `; + resultsContainer.innerHTML = ''; + const searchStatus = await createSearchStatusPanel({ + icon: '🔍', + title: 'Civitai Model Search', + messageLine1: 'Search for models on Civitai to discover and download new content.', + messageLine2: 'Use the search form above to find models by name, type, or creator.', + note: '' + }); + resultsContainer.appendChild(searchStatus); } try { @@ -228,7 +216,7 @@ function setupSearchHandlers(searchForm, resultsContainer) { currentMetadata = response.metadata; } - displaySearchResults(resultsContainer, currentResults, { + await displaySearchResults(resultsContainer, currentResults, { includeNsfw, metadata: currentMetadata, onNextPage: currentMetadata && currentMetadata.nextCursor ? @@ -237,13 +225,16 @@ function setupSearchHandlers(searchForm, resultsContainer) { } catch (error) { console.error('Civitai search error:', error); - resultsContainer.innerHTML = ` -
-

❌ Search Failed

-

Failed to search Civitai: ${escapeHtml(error.message)}

-

Please check your internet connection and try again.

-
- `; + resultsContainer.innerHTML = ''; + const errorStatus = await createSearchStatusPanel({ + icon: '❌', + title: 'Search Failed', + messageLine1: `Failed to search Civitai: ${escapeHtml(error.message)}`, + messageLine2: '', + note: 'Please check your internet connection and try again.', + extraClass: ' civitai-search-status-error' + }); + resultsContainer.appendChild(errorStatus); } finally { loadingButton.disabled = false; loadingButton.textContent = originalText; @@ -319,45 +310,38 @@ async function searchCivitaiModels(options) { * @param {Array} results - Search results from Civitai API * @param {Object} options - Display options */ -function displaySearchResults(container, results, options = {}) { +async function displaySearchResults(container, results, options = {}) { const { includeNsfw = false, metadata = null, onNextPage = null } = options; if (!results || results.length === 0) { - container.innerHTML = ` -
-

No Results Found

-

No models found matching your search criteria.

-

Try adjusting your search terms or filters.

-
- `; + container.innerHTML = ''; + const noResultsStatus = await createSearchStatusPanel({ + icon: '', + title: 'No Results Found', + messageLine1: 'No models found matching your search criteria.', + messageLine2: 'Try adjusting your search terms or filters.', + note: '' + }); + container.appendChild(noResultsStatus); return; } - const header = document.createElement('div'); - header.className = 'civitai-results-header'; - - const headerText = document.createElement('div'); - headerText.innerHTML = ` -

Found ${results.length} models${metadata && metadata.totalItems ? ` (${metadata.totalItems} total)` : ''}

-

Click on any model to view details and download options

- `; - header.appendChild(headerText); + const header = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiResultsHeader.html', + { + resultCount: results.length, + totalItemsText: metadata && metadata.totalItems ? ` (${metadata.totalItems} total)` : '' + } + ); if (metadata) { - const paginationInfo = document.createElement('div'); - paginationInfo.className = 'civitai-results-pagination-info'; - - if (metadata.currentPage) { - paginationInfo.innerHTML = ` -
Page ${metadata.currentPage}${metadata.totalPages ? ` of ${metadata.totalPages}` : ''}
-
${metadata.pageSize || results.length} per page
- `; - } else if (metadata.nextCursor) { - paginationInfo.innerHTML = ` -
Cursor-based pagination
-
More results available
- `; - } + const paginationInfo = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiPaginationInfo.html', + { + pageLine: metadata.currentPage ? `Page ${metadata.currentPage}${metadata.totalPages ? ` of ${metadata.totalPages}` : ''}` : 'Cursor-based pagination', + detailText: metadata.currentPage ? `${metadata.pageSize || results.length} per page` : 'More results available' + } + ); header.appendChild(paginationInfo); } @@ -365,10 +349,10 @@ function displaySearchResults(container, results, options = {}) { const resultsGrid = document.createElement('div'); resultsGrid.className = 'civitai-results-grid'; - results.forEach(model => { - const modelCard = createModelCard(model, { includeNsfw }); + for (const model of results) { + const modelCard = await createModelCard(model, { includeNsfw }); resultsGrid.appendChild(modelCard); - }); + } const paginationControls = document.createElement('div'); paginationControls.className = 'civitai-pagination-controls'; @@ -400,11 +384,82 @@ function displaySearchResults(container, results, options = {}) { container.appendChild(paginationControls); } -function createModelCard(model, options = {}) { - const { includeNsfw = false } = options; +function createTextSpan(text) { + const span = document.createElement('span'); + span.textContent = text; + return span; +} - const card = document.createElement('div'); - card.className = 'civitai-model-card'; +function createTriggerGroup(trainedWords) { + const group = document.createElement('div'); + group.className = 'civitai-model-card-triggers'; + group.textContent = `🏷️ ${trainedWords.length} triggers`; + return group; +} + +function createVersionTriggerGroup(trainedWords) { + const group = document.createElement('div'); + group.className = 'civitai-trigger-words-group'; + + const title = document.createElement('strong'); + title.className = 'civitai-trigger-words-title'; + title.textContent = 'Trigger Words:'; + + const list = document.createElement('div'); + list.className = 'civitai-trigger-words'; + trainedWords.forEach((word) => { + const tag = document.createElement('span'); + tag.className = 'civitai-trigger-word'; + tag.textContent = word; + list.appendChild(tag); + }); + + group.appendChild(title); + group.appendChild(list); + return group; +} + +function createFileInfo(primaryFile) { + const fileInfo = document.createElement('div'); + fileInfo.className = 'civitai-version-file-info'; + + const sizeSpan = document.createElement('span'); + sizeSpan.textContent = `File Size: ${formatFileSize((primaryFile.sizeKB || 0) * 1024)}`; + fileInfo.appendChild(sizeSpan); + + if (primaryFile.metadata) { + if (primaryFile.metadata.format) { + const formatSpan = document.createElement('span'); + formatSpan.textContent = `Format: ${primaryFile.metadata.format}`; + fileInfo.appendChild(formatSpan); + } + + if (primaryFile.metadata.fp) { + const precisionSpan = document.createElement('span'); + precisionSpan.textContent = `Precision: ${primaryFile.metadata.fp}`; + fileInfo.appendChild(precisionSpan); + } + } + + return fileInfo; +} + +function createStatsContainer(stats) { + const fragment = document.createDocumentFragment(); + if (stats.downloadCount) { + fragment.appendChild(createTextSpan(`📥 ${stats.downloadCount.toLocaleString()}`)); + } + if (stats.favoriteCount) { + fragment.appendChild(createTextSpan(`❤️ ${stats.favoriteCount.toLocaleString()}`)); + } + if (stats.rating) { + fragment.appendChild(createTextSpan(`⭐ ${stats.rating.toFixed(1)}`)); + } + return fragment; +} + +async function createModelCard(model, options = {}) { + const { includeNsfw = false } = options; const modelName = model.name || 'Unknown Model'; const creator = model.creator ? model.creator.username : 'Unknown'; @@ -424,38 +479,65 @@ function createModelCard(model, options = {}) { } } - card.innerHTML = ` -
-
- ${imageUrl ? - `Model preview` : - `
No Image
` - } -
-
-

${escapeHtml(modelName)}

-

by ${escapeHtml(creator)} • ${escapeHtml(type)}

-

${escapeHtml(description)}

+ const card = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiModelCard.html', + { + modelName, + creator, + type, + description + } + ); -
-
- ${stats.downloadCount ? `📥 ${stats.downloadCount.toLocaleString()}` : ''} - ${stats.favoriteCount ? `❤️ ${stats.favoriteCount.toLocaleString()}` : ''} - ${stats.rating ? `⭐ ${stats.rating.toFixed(1)}` : ''} -
- ${trainedWords.length > 0 ? `
🏷️ ${trainedWords.length} triggers
` : ''} -
+ const imageContainer = card.querySelector('.civitai-model-card-image'); + if (imageContainer) { + if (imageUrl) { + const img = document.createElement('img'); + img.src = imageUrl; + img.className = 'civitai-model-card-image-img'; + img.alt = 'Model preview'; + img.loading = 'lazy'; + imageContainer.appendChild(img); + } else { + const noImage = document.createElement('div'); + noImage.className = 'civitai-model-card-no-image'; + noImage.textContent = 'No Image'; + imageContainer.appendChild(noImage); + } + } -
-
- - ${latestVersion ? `` : ''} -
- ${latestVersion && latestVersion.files && latestVersion.files[0] ? `
${formatFileSize((latestVersion.files[0].sizeKB || 0) * 1024)}
` : ''} -
-
-
- `; + const statsContainer = card.querySelector('.civitai-model-card-stats'); + if (statsContainer) { + statsContainer.appendChild(createStatsContainer(stats)); + } + + if (trainedWords.length > 0) { + const triggerContainer = card.querySelector('.civitai-model-card-triggers'); + if (triggerContainer) { + triggerContainer.appendChild(createTriggerGroup(trainedWords)); + } + } + + const buttonGroup = card.querySelector('.civitai-card-btn-group'); + if (buttonGroup && latestVersion) { + const downloadBtn = document.createElement('button'); + downloadBtn.className = 'civitai-card-btn civitai-download-btn'; + downloadBtn.dataset.versionId = latestVersion.id; + downloadBtn.dataset.versionName = latestVersion.name || 'Latest Version'; + downloadBtn.dataset.modelId = model.id; + downloadBtn.textContent = '📥 Download Latest'; + buttonGroup.appendChild(downloadBtn); + } + + if (latestVersion && latestVersion.files && latestVersion.files[0]) { + const fileSizeElement = document.createElement('div'); + fileSizeElement.className = 'civitai-card-stats-small'; + fileSizeElement.textContent = formatFileSize((latestVersion.files[0].sizeKB || 0) * 1024); + const fileSizeContainer = card.querySelector('.civitai-card-stats-small'); + if (fileSizeContainer) { + fileSizeContainer.appendChild(fileSizeElement); + } + } const detailsBtn = card.querySelector('.civitai-details-btn'); const downloadBtn = card.querySelector('.civitai-download-btn'); @@ -488,7 +570,7 @@ function createModelCard(model, options = {}) { return card; } -function showModelDetails(model, options = {}) { +async function showModelDetails(model, options = {}) { const { includeNsfw = false } = options; const modelName = model.name || 'Unknown Model'; @@ -498,52 +580,42 @@ function showModelDetails(model, options = {}) { const stats = model.stats || {}; const versions = model.modelVersions || []; - const content = document.createElement('div'); - content.className = 'civitai-dialog-content'; + const content = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiDetailsDialog.html', + { + modelName, + creator, + type, + modelUrl: getModelUrl(model.id), + descriptionText: stripHtml(description) + } + ); - const header = document.createElement('div'); - header.className = 'civitai-dialog-header'; - header.innerHTML = ` -

${escapeHtml(modelName)}

-
-
- by ${escapeHtml(creator)} - Type: ${escapeHtml(type)} -
- View on Civitai → -
-
- ${stats.downloadCount ? `📥 ${stats.downloadCount.toLocaleString()} downloads` : ''} - ${stats.favoriteCount ? `❤️ ${stats.favoriteCount.toLocaleString()} favorites` : ''} - ${stats.rating ? `⭐ ${stats.rating.toFixed(1)} rating` : ''} - ${stats.commentCount ? `💬 ${stats.commentCount.toLocaleString()} comments` : ''} -
- `; - - const descriptionSection = document.createElement('div'); - descriptionSection.className = 'civitai-dialog-description'; - descriptionSection.innerHTML = ` -

Description

-
${escapeHtml(stripHtml(description))}
- `; - - const versionsSection = document.createElement('div'); - versionsSection.className = 'civitai-versions-section'; - - if (versions.length > 0) { - versionsSection.innerHTML = ` -

Available Versions (${versions.length})

- `; - - versions.forEach((version, index) => { - const versionCard = createVersionCard(version, model.id, { includeNsfw, isLatest: index === 0 }); - versionsSection.appendChild(versionCard); - }); + const statsContainer = content.querySelector('.civitai-dialog-stats'); + if (statsContainer) { + statsContainer.appendChild(createStatsContainer(stats)); } - content.appendChild(header); - content.appendChild(descriptionSection); - content.appendChild(versionsSection); + let versionsSection = null; + + if (versions.length > 0) { + versionsSection = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiVersionSection.html', + { + versionCount: versions.length + } + ); + + const versionList = versionsSection.querySelector('.civitai-version-list'); + for (const [index, version] of versions.entries()) { + const versionCard = await createVersionCard(version, model.id, { includeNsfw, isLatest: index === 0 }); + versionList.appendChild(versionCard); + } + } + + if (versionsSection) { + content.appendChild(versionsSection); + } const dialog = createDialog({ title: `${modelName} - Civitai Model Details`, @@ -555,12 +627,9 @@ function showModelDetails(model, options = {}) { dialog.show(); } -function createVersionCard(version, modelId, options = {}) { +async function createVersionCard(version, modelId, options = {}) { const { includeNsfw = false, isLatest = false } = options; - const card = document.createElement('div'); - card.className = `civitai-version-card${isLatest ? ' latest' : ''}`; - const versionName = version.name || 'Unnamed Version'; const createdAt = version.createdAt ? new Date(version.createdAt).toLocaleDateString() : 'Unknown date'; const trainedWords = version.trainedWords || []; @@ -570,41 +639,55 @@ function createVersionCard(version, modelId, options = {}) { const images = version.images || []; const appropriateImages = images.filter(img => includeNsfw || (img.nsfwLevel || 0) <= SEARCH_CONFIG.MAX_NSFW_LEVEL).slice(0, 3); - card.innerHTML = ` -
-
-

${escapeHtml(versionName)}${isLatest ? ' LATEST' : ''}

-

Released: ${createdAt}

- ${version.description ? `

${escapeHtml(version.description)}

` : ''} - ${trainedWords.length > 0 ? ` -
- Trigger Words: -
${trainedWords.map(word => `${escapeHtml(word)}`).join('')}
-
- ` : ''} - ${primaryFile ? ` -
- File Size: ${formatFileSize(primaryFile.sizeKB * 1024)} - ${primaryFile.metadata ? ` - ${primaryFile.metadata.format ? `Format: ${primaryFile.metadata.format}` : ''} - ${primaryFile.metadata.fp ? `Precision: ${primaryFile.metadata.fp}` : ''} - ` : ''} -
- ` : ''} -
-
- -
-
- ${appropriateImages.length > 0 ? ` -
- ${appropriateImages.map(img => `Example image`).join('')} -
- ` : ''} - `; + const latestClass = isLatest ? ' latest' : ''; + + const card = await loadAndCreateHtmlTemplate( + 'extensions/comfyui_sageutils/sidebar/partials/civitaiVersionCard.html', + { + latestClass, + versionName, + createdAt + } + ); + + if (isLatest) { + const titleEl = card.querySelector('.civitai-version-title'); + if (titleEl) { + const badge = document.createElement('span'); + badge.className = 'civitai-version-badge'; + badge.textContent = 'LATEST'; + titleEl.appendChild(badge); + } + } + + const descriptionContainer = card.querySelector('.civitai-version-description'); + if (descriptionContainer && version.description) { + const descriptionParagraph = document.createElement('p'); + descriptionParagraph.className = 'civitai-version-description'; + descriptionParagraph.textContent = escapeHtml(version.description); + descriptionContainer.appendChild(descriptionParagraph); + } + + const triggerContainer = card.querySelector('.civitai-trigger-words-group'); + if (triggerContainer && trainedWords.length > 0) { + triggerContainer.appendChild(createVersionTriggerGroup(trainedWords)); + } + + const fileInfoContainer = card.querySelector('.civitai-version-file-info'); + if (fileInfoContainer && primaryFile) { + fileInfoContainer.appendChild(createFileInfo(primaryFile)); + } + + const downloadContainer = card.querySelector('.civitai-version-download-container'); + if (downloadContainer) { + const downloadBtn = document.createElement('button'); + downloadBtn.className = 'civitai-version-download-button'; + downloadBtn.dataset.versionId = version.id; + downloadBtn.dataset.versionName = escapeHtml(versionName); + downloadBtn.dataset.modelId = modelId; + downloadBtn.textContent = '📥 Download'; + downloadContainer.appendChild(downloadBtn); - const downloadBtn = card.querySelector('.civitai-version-download-button'); - if (downloadBtn) { downloadBtn.addEventListener('click', (e) => { e.stopPropagation(); const versionId = downloadBtn.dataset.versionId; @@ -614,10 +697,19 @@ function createVersionCard(version, modelId, options = {}) { }); } - const imageElements = card.querySelectorAll('.civitai-version-image'); - imageElements.forEach(img => { - img.addEventListener('click', () => showImageExpanded(img.dataset.imageUrl)); - }); + const imagesContainer = card.querySelector('.civitai-version-images'); + if (imagesContainer && appropriateImages.length > 0) { + appropriateImages.forEach(img => { + const imageElement = document.createElement('img'); + imageElement.src = img.url; + imageElement.className = 'civitai-version-image'; + imageElement.dataset.imageUrl = img.url; + imageElement.alt = 'Example image'; + imageElement.loading = 'lazy'; + imageElement.addEventListener('click', () => showImageExpanded(img.url)); + imagesContainer.appendChild(imageElement); + }); + } return card; } diff --git a/js/sidebar/partials/civitaiDetailsDialog.html b/js/sidebar/partials/civitaiDetailsDialog.html new file mode 100644 index 0000000..5521ef0 --- /dev/null +++ b/js/sidebar/partials/civitaiDetailsDialog.html @@ -0,0 +1,19 @@ +
+
+

{{modelName}}

+
+
+ by + {{creator}} + Type: + {{type}} +
+ View on Civitai → +
+
+
+
+

Description

+
{{descriptionText}}
+
+
diff --git a/js/sidebar/partials/civitaiModelCard.html b/js/sidebar/partials/civitaiModelCard.html new file mode 100644 index 0000000..59c10aa --- /dev/null +++ b/js/sidebar/partials/civitaiModelCard.html @@ -0,0 +1,22 @@ +
+
+
+
+

{{modelName}}

+

by {{creator}} • {{type}}

+

{{description}}

+ +
+
+
+
+ +
+
+ +
+
+
+
+
+
diff --git a/js/sidebar/partials/civitaiPaginationInfo.html b/js/sidebar/partials/civitaiPaginationInfo.html new file mode 100644 index 0000000..16965d4 --- /dev/null +++ b/js/sidebar/partials/civitaiPaginationInfo.html @@ -0,0 +1,4 @@ +
+
{{pageLine}}
+
{{detailText}}
+
diff --git a/js/sidebar/partials/civitaiResultsHeader.html b/js/sidebar/partials/civitaiResultsHeader.html new file mode 100644 index 0000000..08be065 --- /dev/null +++ b/js/sidebar/partials/civitaiResultsHeader.html @@ -0,0 +1,6 @@ +
+
+

Found {{resultCount}} models{{{totalItemsText}}}

+

Click on any model to view details and download options

+
+
diff --git a/js/sidebar/partials/civitaiSearchForm.html b/js/sidebar/partials/civitaiSearchForm.html new file mode 100644 index 0000000..5f9d8b6 --- /dev/null +++ b/js/sidebar/partials/civitaiSearchForm.html @@ -0,0 +1,40 @@ +
+
+ + +
+ +
+ + +
+ +
+ + +
+ +
+
+ + + +
+ + +
+
+
diff --git a/js/sidebar/partials/civitaiSearchStatus.html b/js/sidebar/partials/civitaiSearchStatus.html new file mode 100644 index 0000000..2e14efe --- /dev/null +++ b/js/sidebar/partials/civitaiSearchStatus.html @@ -0,0 +1,7 @@ +
+
{{icon}}
+

{{title}}

+

{{messageLine1}}

+

{{messageLine2}}

+

{{note}}

+
diff --git a/js/sidebar/partials/civitaiVersionCard.html b/js/sidebar/partials/civitaiVersionCard.html new file mode 100644 index 0000000..6550dbc --- /dev/null +++ b/js/sidebar/partials/civitaiVersionCard.html @@ -0,0 +1,13 @@ +
+
+
+

{{versionName}}

+

Released: {{createdAt}}

+
+
+
+
+
+
+
+
diff --git a/js/sidebar/partials/civitaiVersionSection.html b/js/sidebar/partials/civitaiVersionSection.html new file mode 100644 index 0000000..05829aa --- /dev/null +++ b/js/sidebar/partials/civitaiVersionSection.html @@ -0,0 +1,4 @@ +
+

Available Versions ({{versionCount}})

+
+
diff --git a/js/utils/htmlTemplateLoader.js b/js/utils/htmlTemplateLoader.js new file mode 100644 index 0000000..20de686 --- /dev/null +++ b/js/utils/htmlTemplateLoader.js @@ -0,0 +1,81 @@ +const htmlTemplateCache = new Map(); + +function escapeHtml(value) { + return value.replace(/[&<>"]|\'/g, (char) => { + switch (char) { + case '&': return '&'; + case '<': return '<'; + case '>': return '>'; + case '"': return '"'; + case "'": return '''; + default: return char; + } + }); +} + +function getTemplateValue(key, data) { + return key.split('.').reduce((current, part) => { + return current && current[part] !== undefined ? current[part] : undefined; + }, data); +} + +export async function loadHtmlTemplate(templatePath) { + if (typeof fetch === 'undefined') { + return ''; + } + + if (htmlTemplateCache.has(templatePath)) { + return htmlTemplateCache.get(templatePath); + } + + const response = await fetch(templatePath, { cache: 'force-cache' }); + if (!response.ok) { + throw new Error(`Unable to load HTML template: ${templatePath} (${response.status})`); + } + + const text = await response.text(); + htmlTemplateCache.set(templatePath, text); + return text; +} + +export function renderHtmlTemplate(templateString, data = {}) { + return templateString + .replace(/\{\{\{\s*([\w.]+)\s*\}\}\}/g, (_, key) => { + const value = getTemplateValue(key, data); + return value == null ? '' : String(value); + }) + .replace(/\{\{\s*([\w.]+)\s*\}\}/g, (_, key) => { + const value = getTemplateValue(key, data); + return value == null ? '' : escapeHtml(String(value)); + }); +} + +export function createElementFromTemplate(templateString, data = {}) { + if (typeof document === 'undefined') { + return null; + } + + const rendered = renderHtmlTemplate(templateString, data).trim(); + const template = document.createElement('template'); + template.innerHTML = rendered; + + if (template.content.childElementCount === 1) { + return template.content.firstElementChild; + } + + return template.content; +} + +export async function loadAndCreateHtmlTemplate(templatePath, data = {}) { + const templateString = await loadHtmlTemplate(templatePath); + return createElementFromTemplate(templateString, data); +} + +export function clearHtmlTemplateCache(templatePath) { + if (templatePath) { + htmlTemplateCache.delete(templatePath); + return; + } + + htmlTemplateCache.clear(); +} diff --git a/nodes/lmstudio_v3.py b/nodes/lmstudio_v3.py index 24a28ac..786b789 100644 --- a/nodes/lmstudio_v3.py +++ b/nodes/lmstudio_v3.py @@ -6,10 +6,10 @@ from comfy_api.latest import io from comfy.utils import ProgressBar from ..utils.llm.service import ( - lmstudio_load_model, - lmstudio_generate_with_model, - lmstudio_generate_vision_with_model, - lmstudio_unload_model, + lmstudio_rest_load_model, + lmstudio_rest_generate, + lmstudio_rest_generate_vision, + lmstudio_rest_unload_model, ) from ..utils.performance_fix import ( get_cached_lmstudio_models_for_input_types, @@ -70,22 +70,22 @@ class Sage_LMStudioLLMPromptText(io.ComfyNode): options = {"seed": seed} pbar = ProgressBar(2) - lms_model = None + model_loaded = False try: # Step 1: load model - lms_model = lmstudio_load_model(model=model, keep_alive=load_for_seconds) + model_loaded = lmstudio_rest_load_model(model=model, keep_alive=load_for_seconds) pbar.update(1) # Step 2: generate - response = lmstudio_generate_with_model(lms_model, prompt=prompt, options=options) + response = lmstudio_rest_generate(model, prompt=prompt, options=options) if load_for_seconds < 1: - lmstudio_unload_model(lms_model) - lms_model = None + lmstudio_rest_unload_model(model) + model_loaded = False pbar.update(1) except Exception: logger.exception('LM Studio text node failed during load or generation') - if lms_model is not None and load_for_seconds < 1: + if model_loaded and load_for_seconds < 1: try: - lmstudio_unload_model(lms_model) + lmstudio_rest_unload_model(model) except Exception: pass if _should_reraise_llm_node_errors(): @@ -132,22 +132,22 @@ class Sage_LMStudioLLMPromptVision(io.ComfyNode): options = {"seed": seed} pbar = ProgressBar(2) - lms_model = None + model_loaded = False try: # Step 1: load model - lms_model = lmstudio_load_model(model=model, keep_alive=load_for_seconds) + model_loaded = lmstudio_rest_load_model(model=model, keep_alive=load_for_seconds) pbar.update(1) # Step 2: generate vision response - response = lmstudio_generate_vision_with_model(lms_model, prompt=prompt, images=image, options=options) + response = lmstudio_rest_generate_vision(model, prompt=prompt, images=image, options=options) if load_for_seconds < 1: - lmstudio_unload_model(lms_model) - lms_model = None + lmstudio_rest_unload_model(model) + model_loaded = False pbar.update(1) except Exception: logger.exception('LM Studio vision node failed during load or generation') - if lms_model is not None and load_for_seconds < 1: + if model_loaded and load_for_seconds < 1: try: - lmstudio_unload_model(lms_model) + lmstudio_rest_unload_model(model) except Exception: pass if _should_reraise_llm_node_errors(): @@ -208,26 +208,30 @@ class Sage_LMStudioLLMPromptVisionRefine(io.ComfyNode): return io.NodeOutput("", "") pbar = ProgressBar(2) - lms_model = None + model_loaded = False + loaded_model_name = '' try: # Step 1: load vision model - lms_model = lmstudio_load_model(model=model, keep_alive=0) + model_loaded = lmstudio_rest_load_model(model=model, keep_alive=0) + loaded_model_name = model pbar.update(1) # Step 2: generate initial vision response, then refine - initial = lmstudio_generate_vision_with_model(lms_model, prompt=prompt, images=image, options={"seed": seed}) + initial = lmstudio_rest_generate_vision(model, prompt=prompt, images=image, options={"seed": seed}) if refine_model != model: - lmstudio_unload_model(lms_model) - lms_model = lmstudio_load_model(model=refine_model, keep_alive=0) + lmstudio_rest_unload_model(model) + model_loaded = lmstudio_rest_load_model(model=refine_model, keep_alive=0) + loaded_model_name = refine_model combined_refine_prompt = f'{actual_refine_prompt}\n{initial}' - refined = lmstudio_generate_with_model(lms_model, prompt=combined_refine_prompt, options={"seed": refine_seed}) - lmstudio_unload_model(lms_model) - lms_model = None + refined = lmstudio_rest_generate(refine_model, prompt=combined_refine_prompt, options={"seed": refine_seed}) + lmstudio_rest_unload_model(loaded_model_name) + model_loaded = False + loaded_model_name = '' pbar.update(1) except Exception: logger.exception('LM Studio vision refine node failed during load or generation') - if lms_model is not None: + if model_loaded and loaded_model_name: try: - lmstudio_unload_model(lms_model) + lmstudio_rest_unload_model(loaded_model_name) except Exception: pass if _should_reraise_llm_node_errors(): diff --git a/nodes/ollama_v3.py b/nodes/ollama_v3.py index 9aeec8f..8c3bfd8 100644 --- a/nodes/ollama_v3.py +++ b/nodes/ollama_v3.py @@ -7,10 +7,8 @@ from comfy_api.latest import io from comfy.utils import ProgressBar from ..utils.llm.service import ( - ollama_generate_vision, - ollama_generate_vision_refine, - ollama_preload_model, - ollama_generate_preloaded, + ollama_rest_generate, + ollama_rest_generate_vision, ) from ..utils.performance_fix import ( get_cached_ollama_models_for_input_types, @@ -40,6 +38,15 @@ def _should_reraise_llm_node_errors() -> bool: return bool(get_setting('llm_raise_node_exceptions', False)) +def _normalize_ollama_keep_alive(value, default: str = '5m') -> str: + """Normalize keep_alive values for Ollama REST generation calls.""" + if isinstance(value, str) and value.strip(): + return value.strip() + if isinstance(value, (int, float)) and value > 0: + return f'{int(value)}s' + return default + + class Sage_OllamaAdvancedOptions(io.ComfyNode): """Get advanced options for Ollama LLMs.""" @classmethod @@ -110,14 +117,19 @@ class Sage_OllamaLLMPromptText(io.ComfyNode): options = options or {} options["seed"] = seed - pbar = ProgressBar(2) + keep_alive_value = _normalize_ollama_keep_alive(keep_alive) + pbar = ProgressBar(1) try: - ollama_preload_model(model=model, keep_alive=keep_alive) - pbar.update(1) - response = ollama_generate_preloaded(model=model, prompt=prompt, keep_alive=keep_alive, options=options, system_prompt=system_prompt) + response = ollama_rest_generate( + model=model, + prompt=prompt, + keep_alive=keep_alive_value, + options=options, + system_prompt=system_prompt, + ) pbar.update(1) except Exception: - logger.exception('Ollama text node failed during preload or generation') + logger.exception('Ollama text node failed during generation') if _should_reraise_llm_node_errors(): raise response = "" @@ -166,14 +178,20 @@ class Sage_OllamaLLMPromptVision(io.ComfyNode): options = options or {} options["seed"] = seed - pbar = ProgressBar(2) + keep_alive_value = _normalize_ollama_keep_alive(keep_alive) + pbar = ProgressBar(1) try: - ollama_preload_model(model=model, keep_alive=keep_alive) - pbar.update(1) - response = ollama_generate_vision(model=model, prompt=prompt, system_prompt=system_prompt, images=image, keep_alive=keep_alive, options=options) + response = ollama_rest_generate_vision( + model=model, + prompt=prompt, + system_prompt=system_prompt, + images=image, + keep_alive=keep_alive_value, + options=options, + ) pbar.update(1) except Exception: - logger.exception('Ollama vision node failed during preload or generation') + logger.exception('Ollama vision node failed during generation') if _should_reraise_llm_node_errors(): raise response = "" @@ -225,24 +243,26 @@ class Sage_OllamaLLMPromptVisionRefine(io.ComfyNode): if model == "(No Ollama vision models available)" or refine_model == "(Ollama not available)" or not OLLAMA_AVAILABLE: return io.NodeOutput("", "") - pbar = ProgressBar(2) + pbar = ProgressBar(1) try: - # Step 1: pre-warm the first model - ollama_preload_model(model=model, keep_alive=0) - pbar.update(1) - # Step 2: generate initial vision response then refine - initial, refined = ollama_generate_vision_refine( + # Step 1: generate initial vision response then refine + initial = ollama_rest_generate_vision( model=model, prompt=prompt, images=image, options={"seed": seed}, - refine_model=refine_model, - refine_prompt=refine_prompt, - refine_options={"seed": refine_seed} ) + if not refine_model or not refine_prompt: + refined = "" + else: + refined = ollama_rest_generate( + model=refine_model, + prompt=refine_prompt, + options={"seed": refine_seed}, + ) pbar.update(1) except Exception: - logger.exception('Ollama vision refine node failed during preload or generation') + logger.exception('Ollama vision refine node failed during generation') if _should_reraise_llm_node_errors(): raise initial, refined = "", "" diff --git a/routes/LLM_ROUTES_README.md b/routes/LLM_ROUTES_README.md index 7e6aa60..9266f0f 100644 --- a/routes/LLM_ROUTES_README.md +++ b/routes/LLM_ROUTES_README.md @@ -1,898 +1,96 @@ -# LLM Routes - Phase 1 & 2 Documentation +# LLM Routes Documentation ## Overview -Phase 1 provides core backend API endpoints for interacting with Ollama and LM Studio language models. Phase 2 adds real-time streaming and vision capabilities. -**Phase 1 Features:** -- Service status checking -- Model discovery (text and vision) -- Prompt template loading -- Basic text generation +LLM routes expose model discovery, text generation, vision generation, streaming, prompts, presets, and model load flows. -**Phase 2 Features:** -- Server-Sent Events (SSE) streaming for real-time token generation -- Vision generation (single image + multi-image) -- Vision streaming (SSE) -- Error handling for disconnects and timeouts +Current route implementation lives in: -## Endpoints Summary +- `routes/llm_routes.py` +- helper normalization/validation in `utils/llm/routes_helpers.py` -| Endpoint | Method | Phase | Description | -|----------|--------|-------|-------------| -| `/sage_llm/status` | GET | 1 | Check service availability | -| `/sage_llm/models` | GET | 1 | List text models | -| `/sage_llm/vision_models` | GET | 1 | List vision models | -| `/sage_llm/prompts` | GET | 1 | Get prompt templates | -| `/sage_llm/generate` | POST | 1 | Generate text (non-streaming) | -| `/sage_llm/generate_stream` | POST | 2 | Generate text (SSE streaming) | -| `/sage_llm/vision_generate` | POST | 2 | Generate with vision (non-streaming) | -| `/sage_llm/vision_generate_stream` | POST | 2 | Generate with vision (SSE streaming) | +## Provider Naming ---- +Canonical backend provider keys used in route responses and internal dispatch: -## Endpoints +- `lmstudio_rest` +- `ollama_rest` +- `openai` +- `native` -### 1. GET /sage_llm/status -**Description:** Check availability and status of Ollama and LM Studio services. +Compatibility aliases accepted in request payloads: + +- `lmstudio` is normalized to `lmstudio_rest` +- `ollama` is normalized to `ollama_rest` + +Normalization is handled by `routes_helpers.normalize_provider(...)`. + +## Core Endpoints + +- `GET /sage_llm/status` +- `GET /sage_llm/models` +- `GET /sage_llm/vision_models` +- `GET /sage_llm/prompts` +- `GET /sage_llm/integration_profiles` +- `POST /sage_llm/generate` +- `POST /sage_llm/generate_stream` +- `POST /sage_llm/vision_generate` +- `POST /sage_llm/vision_generate_stream` +- `POST /sage_llm/load_model` +- `POST /sage_llm/generate_only` + +Additional preset/system prompt endpoints are also registered in `routes/llm_routes.py`. + +## Status Response Shape + +`GET /sage_llm/status` returns canonical keys: -**Response:** ```json { "success": true, "data": { - "ollama": { - "available": true, - "enabled": true, - "url": "http://custom-url:11434" // if using custom URL - }, - "lmstudio": { - "available": false, - "enabled": true - } + "lmstudio_rest": {"available": true, "enabled": true}, + "ollama_rest": {"available": true, "enabled": true}, + "openai": {"available": false, "enabled": false}, + "native": {"available": true, "enabled": true} } } ``` -**Use Case:** Frontend can check which services are available before showing model selection. +## Models Response Shape ---- +`GET /sage_llm/models` and `GET /sage_llm/vision_models` return provider-model maps keyed by canonical provider names. -### 2. GET /sage_llm/models -**Description:** Get list of available text-only models from both providers. +`status` subfields also use canonical naming, for example: -**Response:** -```json -{ - "success": true, - "data": { - "models": { - "ollama": ["llama3", "mistral", "codellama"], - "lmstudio": ["llama-2-7b-chat"] - }, - "status": { - "ollama_available": true, - "lmstudio_available": true - } - } -} -``` +- `lmstudio_rest_available` +- `ollama_rest_available` +- `openai_available` +- `native_available` -**Notes:** -- Uses existing model cache from `utils/llm/cache.py` -- Filters out placeholder messages like "(Ollama not available)" -- Returns empty arrays if service is disabled or unavailable +## Generation Request Notes ---- +For `POST /sage_llm/generate` and streaming/vision variants: -### 3. GET /sage_llm/vision_models -**Description:** Get list of available vision-capable models from both providers. +- Required fields vary by endpoint, but always include `provider`, `model`, and `prompt`. +- Vision endpoints also require `images`. +- Provider is normalized before dispatch, so legacy aliases continue to work. -**Response:** -```json -{ - "success": true, - "data": { - "models": { - "ollama": ["llava", "bakllava", "llava-phi3"], - "lmstudio": ["llava-v1.6-mistral"] - }, - "status": { - "ollama_available": true, - "lmstudio_available": true - } - } -} -``` +## Streaming Notes -**Notes:** -- Similar to `/sage_llm/models` but for multimodal models -- Vision generation will be implemented in Phase 2 +Streaming routes use SSE (`text/event-stream`) and emit chunks in this form: ---- +```text +data: {"chunk":"...","done":false} -### 4. GET /sage_llm/prompts -**Description:** Get available LLM prompt templates from `assets/llm_prompts.json`. - -**Response:** -```json -{ - "success": true, - "data": { - "prompts": { - "base": { - "descriptive": { - "name": "Descriptive Prompt", - "category": "description", - "input_type": "image", - "prompt": "Write a detailed description for this image." - }, - "danbooru": { - "name": "Danbooru Tag List", - "category": "tagging", - "input_type": "image", - "prompt": "Generate only comma-separated Danbooru tags..." - } - // ... more prompts - }, - "extra": { - "be_specific": { - "name": "Be Specific", - "category": "content_focus", - "type": "boolean", - "prompt": "Use precise, unambiguous language..." - } - // ... more extra options - } - } - } -} -``` - -**Notes:** -- Serves the entire `llm_prompts.json` file -- Frontend can use this to populate prompt template dropdowns -- Includes both base prompts and extra instruction options - ---- - -### 5. POST /sage_llm/generate -**Description:** Generate a text response from a language model (non-streaming in Phase 1). - -**Request Body:** -```json -{ - "provider": "ollama", // "ollama" or "lmstudio" - "model": "llama3", - "prompt": "Explain quantum computing in simple terms.", - "system_prompt": "You are a helpful assistant.", // optional - "options": { - "temperature": 0.7, - "seed": 42, - "max_tokens": 1024, - // Ollama-specific options - "num_predict": 256, - "top_k": 40, - "top_p": 0.9, - "repeat_penalty": 1.1 - } -} -``` - -**Response (Success):** -```json -{ - "success": true, - "data": { - "response": "Quantum computing is a revolutionary approach to computation...", - "provider": "ollama", - "model": "llama3" - } -} -``` - -**Response (Error - Missing Fields):** -```json -{ - "success": false, - "error": "Missing required fields: model, prompt" -} -``` - -**Response (Error - Service Unavailable):** -```json -{ - "success": false, - "error": "Ollama is not available" -} -``` - -**Notes:** -- Non-streaming in Phase 1 (full response returned at once) -- **Phase 2 adds SSE streaming** (see `/sage_llm/generate_stream`) -- Validates required fields and provider type -- Uses existing `utils/llm/service.py` functions -- System prompt optional (useful for context/persona) - ---- - -### 6. POST /sage_llm/generate_stream (Phase 2) -**Description:** Generate text with Server-Sent Events (SSE) for real-time token streaming. - -**Request Body:** -```json -{ - "provider": "ollama", // "ollama" or "lmstudio" - "model": "llama3", - "prompt": "Write a haiku about coding in Python.", - "system_prompt": "You are a poetic assistant.", // optional - "options": { - "temperature": 0.8, - "seed": 42, - "max_tokens": 1024 - } -} -``` - -**SSE Stream Format:** -``` -data: {"chunk": "Quantum", "done": false} - -data: {"chunk": " computing", "done": false} - -data: {"chunk": " is", "done": false} - -data: {"chunk": "", "done": true, "full_response": "Quantum computing is..."} +data: {"chunk":"","done":true,"full_response":"..."} ``` -**Response Structure:** -- Each line: `data: {json}\n\n` -- **chunk**: Text fragment to append -- **done**: Boolean indicating completion -- **full_response**: Complete text (only in final chunk) -- **error**: Error message (only if error occurred) +Error chunks include standardized metadata fields such as `error_code`, and when available, `provider`, `operation`, and `cause`. -**JavaScript Example:** -```javascript -const eventSource = new EventSource('/sage_llm/generate_stream?' + new URLSearchParams({ - provider: 'ollama', - model: 'llama3', - prompt: 'Write a haiku about coding', - options: JSON.stringify({ temperature: 0.8 }) -})); +## Phase 6 Direction -let fullText = ''; - -eventSource.onmessage = (event) => { - const data = JSON.parse(event.data); - - if (data.chunk) { - fullText += data.chunk; - document.getElementById('output').textContent = fullText; - } - - if (data.done) { - console.log('Generation complete:', fullText); - eventSource.close(); - } - - if (data.error) { - console.error('Error:', data.error); - eventSource.close(); - } -}; - -eventSource.onerror = () => { - console.error('Connection error'); - eventSource.close(); -}; -``` - -**curl Example:** -```bash -curl -N -X POST http://localhost:8188/sage_llm/generate_stream \ - -H "Content-Type: application/json" \ - -d '{ - "provider": "ollama", - "model": "llama3", - "prompt": "Write a haiku about coding", - "options": {"temperature": 0.8} - }' -``` - -**Notes:** -- Uses HTTP chunked transfer encoding -- Content-Type: `text/event-stream` -- Client must support SSE (browsers, curl -N, Python SSE libraries) -- Ollama: True streaming from API -- LM Studio: Simulated streaming (5-character chunks) - ---- - -### 7. POST /sage_llm/vision_generate (Phase 2) -**Description:** Generate text from images (non-streaming vision generation). - -**Request Body:** -```json -{ - "provider": "ollama", // "ollama" or "lmstudio" - "model": "llava", - "prompt": "Describe this image in detail.", - "images": [ - "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ...", // base64 PNG/JPEG - "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAA..." // data URI also supported - ], - "system_prompt": "You are a detailed image analyst.", // optional - "options": { - "temperature": 0.7, - "seed": 42 - } -} -``` - -**Response (Success):** -```json -{ - "success": true, - "data": { - "response": "This image shows a red square on a white background...", - "provider": "ollama", - "model": "llava" - } -} -``` - -**Image Format Notes:** -- **Ollama**: Accepts base64 strings directly (with or without data URI prefix) -- **LM Studio**: Requires temp file conversion (handled automatically) -- Supported formats: PNG, JPEG, GIF, WebP -- Multiple images supported (multi-modal models only) -- Images automatically cleaned up after generation - -**Python Example (Encoding Image):** -```python -import base64 - -# Read and encode image -with open('image.png', 'rb') as f: - img_b64 = base64.b64encode(f.read()).decode('utf-8') - -# Use in request -response = requests.post('http://localhost:8188/sage_llm/vision_generate', json={ - 'provider': 'ollama', - 'model': 'llava', - 'prompt': 'What is in this image?', - 'images': [img_b64] -}) -``` - ---- - -### 8. POST /sage_llm/vision_generate_stream (Phase 2) -**Description:** Generate text from images with SSE streaming for real-time token generation. - -**Request Body:** -```json -{ - "provider": "ollama", - "model": "llava", - "prompt": "Describe what you see in this image.", - "images": [ - "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ..." - ], - "system_prompt": "You are a visual analyst.", // optional - "options": { - "temperature": 0.7, - "seed": 42 - } -} -``` - -**SSE Stream Format:** -``` -data: {"chunk": "This", "done": false} - -data: {"chunk": " image", "done": false} - -data: {"chunk": " shows", "done": false} - -data: {"chunk": "", "done": true, "full_response": "This image shows..."} - -``` - -**JavaScript Example:** -```javascript -// Encode image to base64 -async function encodeImage(file) { - return new Promise((resolve) => { - const reader = new FileReader(); - reader.onload = (e) => { - const base64 = e.target.result.split(',')[1]; // Remove data URI prefix - resolve(base64); - }; - reader.readAsDataURL(file); - }); -} - -// Stream vision generation -async function streamVisionGeneration(imageFile, prompt) { - const imageBase64 = await encodeImage(imageFile); - - const response = await fetch('/sage_llm/vision_generate_stream', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ - provider: 'ollama', - model: 'llava', - prompt: prompt, - images: [imageBase64], - options: { temperature: 0.7 } - }) - }); - - const reader = response.body.getReader(); - const decoder = new TextDecoder(); - let fullText = ''; - - while (true) { - const { done, value } = await reader.read(); - if (done) break; - - const chunk = decoder.decode(value); - const lines = chunk.split('\n'); - - for (const line of lines) { - if (line.startsWith('data: ')) { - const data = JSON.parse(line.slice(6)); - - if (data.chunk) { - fullText += data.chunk; - updateDisplay(fullText); // Your UI update function - } - - if (data.done) { - console.log('Complete:', fullText); - return fullText; - } - - if (data.error) { - throw new Error(data.error); - } - } - } - } -} -``` - -**curl Example:** -```bash -# Encode image -IMAGE_B64=$(base64 -w 0 image.png) - -# Stream vision generation -curl -N -X POST http://localhost:8188/sage_llm/vision_generate_stream \ - -H "Content-Type: application/json" \ - -d "{ - \"provider\": \"ollama\", - \"model\": \"llava\", - \"prompt\": \"Describe this image\", - \"images\": [\"$IMAGE_B64\"], - \"options\": {\"temperature\": 0.7} - }" -``` - -**Notes:** -- Combines vision capabilities with real-time streaming -- Base64 encoding required for all images -- Temp files auto-cleaned after generation (LM Studio) -- Multiple images supported (model-dependent) -- SSE format identical to text streaming - ---- - -## Next Steps (Phase 2) - -Phase 2 will add: -1. **Streaming responses** using Server-Sent Events (SSE) -2. **Vision generation** endpoints -3. **Stop generation** capability -4. **Progress indicators** -5. **Connection error handling** - -The non-streaming endpoint will remain available for compatibility. - -## Error Handling - -All endpoints use consistent error format: - -```json -{ - "success": false, - "error": "Error message description" -} -``` - -**SSE Error Format:** -``` -data: {"chunk": "", "done": true, "error": "Error message"} - -``` - -Common HTTP status codes: -- `200` - Success (including SSE streams) -- `400` - Bad request (invalid input, missing fields) -- `500` - Internal server error -- `503` - Service unavailable (LLM service not running) - -**Common Errors:** - -| Error Message | Cause | Solution | -|---------------|-------|----------| -| "Missing required fields: model, prompt" | Request missing parameters | Include all required fields | -| "Ollama is not available" | Ollama not running | Start Ollama service | -| "LM Studio is not available" | LM Studio not running | Start LM Studio | -| "Provider must be 'ollama' or 'lmstudio'" | Invalid provider | Use valid provider name | -| "No images provided" | Vision endpoint called without images | Include base64 images array | -| "Model [name] not found" | Model not installed | Install model or use different model | - ---- - -## Testing - -### Phase 1 Tests (Basic Endpoints): -```bash -# Bash version -bash test_llm_routes.sh - -# Python version -python test_llm_routes.py -``` - -### Phase 2 Tests (Streaming & Vision): -```bash -# Bash version (uses curl with SSE) -bash test_llm_routes_phase2.sh - -# Python version (uses requests with streaming) -python test_llm_routes_phase2.py -``` - -### Manual Testing with curl: - -**Test Status:** -```bash -curl http://localhost:8188/sage_llm/status -``` - -**Test Models:** -```bash -curl http://localhost:8188/sage_llm/models -``` - -**Test Text Generation:** -```bash -curl -X POST http://localhost:8188/sage_llm/generate \ - -H "Content-Type: application/json" \ - -d '{"provider":"ollama","model":"llama3","prompt":"Hello!"}' -``` - -**Test Streaming (SSE):** -```bash -curl -N -X POST http://localhost:8188/sage_llm/generate_stream \ - -H "Content-Type: application/json" \ - -d '{"provider":"ollama","model":"llama3","prompt":"Write a haiku"}' -``` - -**Test Vision:** -```bash -# Encode image -IMAGE_B64=$(base64 -w 0 test_image.png) - -# Non-streaming -curl -X POST http://localhost:8188/sage_llm/vision_generate \ - -H "Content-Type: application/json" \ - -d "{\"provider\":\"ollama\",\"model\":\"llava\",\"prompt\":\"Describe this\",\"images\":[\"$IMAGE_B64\"]}" - -# Streaming -curl -N -X POST http://localhost:8188/sage_llm/vision_generate_stream \ - -H "Content-Type: application/json" \ - -d "{\"provider\":\"ollama\",\"model\":\"llava\",\"prompt\":\"Describe this\",\"images\":[\"$IMAGE_B64\"]}" -``` - ---- - -## Testing - -### Using curl: -```bash -# Test status -curl http://localhost:8188/sage_llm/status - -# Test models -curl http://localhost:8188/sage_llm/models - -# Test generation -curl -X POST http://localhost:8188/sage_llm/generate \ - -H "Content-Type: application/json" \ - -d '{"provider":"ollama","model":"llama3","prompt":"Hello!"}' -``` - -### Using the test script: -```bash -# Bash version -bash test_llm_routes.sh - -# Python version (better output) -python test_llm_routes.py -``` - ---- - -## Testing - -### Using curl: -```bash -# Test status -curl http://localhost:8188/sage_llm/status - -# Test models -curl http://localhost:8188/sage_llm/models - -# Test generation -curl -X POST http://localhost:8188/sage_llm/generate \ - -H "Content-Type: application/json" \ - -d '{"provider":"ollama","model":"llama3","prompt":"Hello!"}' -``` - -### Using the test script: -```bash -# Bash version -bash test_llm_routes.sh - -# Python version (better output) -python test_llm_routes.py -``` - ---- - -## Integration with Existing Code - -### LLM Service Integration -Routes use the existing `utils/llm/service.py` module: - -**Phase 1 Functions:** -- `ensure_llm_initialized()` - Initialize services lazily -- `get_ollama_models()` - Get Ollama text models -- `get_ollama_vision_models()` - Get Ollama vision models -- `get_lmstudio_models()` - Get LM Studio text models -- `get_lmstudio_vision_models()` - Get LM Studio vision models -- `ollama_generate()` - Generate with Ollama (non-streaming) -- `lmstudio_generate()` - Generate with LM Studio (non-streaming) - -**Phase 2 Functions (Streaming):** -- `ollama_generate_stream()` - Streaming text generation (Ollama) -- `ollama_generate_vision_stream()` - Streaming vision generation (Ollama) -- `lmstudio_generate_stream()` - Simulated streaming text (LM Studio) -- `lmstudio_generate_vision_stream()` - Simulated streaming vision (LM Studio) - -### Model Cache Integration -Routes leverage the existing LLM model cache: -- `utils/llm/cache.py` - Thread-safe provider-aware caching (5-minute TTL) -- Reduces API calls during node initialization -- Shared between routes and Python nodes - -### Settings Integration -Routes respect existing settings: -- `enable_ollama` - Enable/disable Ollama -- `enable_lmstudio` - Enable/disable LM Studio -- `ollama_use_custom_url` - Use custom Ollama URL -- `ollama_custom_url` - Custom Ollama URL -- `lmstudio_use_custom_url` - Use custom LM Studio URL -- `lmstudio_custom_url` - Custom LM Studio URL - ---- - -## Next Steps (Phase 3+) - -**Phase 3: Frontend Core** -- Build LLM tab UI in sidebar -- Model selection dropdown -- Chat interface with message history -- Streaming token display -- Send/stop buttons - -**Phase 4: Settings Integration** -- Model parameter controls (temperature, seed, etc.) -- Provider selection -- Custom system prompts -- Save/load chat history - -**Phase 5: Vision UI** -- Image upload area -- Gallery integration hooks -- Multi-image support -- Preview thumbnails - -**Phase 8: Cross-Tab Integration** -- Gallery → LLM: Attach images to chat -- LLM → Prompts: Send generated text -- Prompts → LLM: Use prompt templates - ---- - -## Security Considerations - -- All endpoints validate input data -- Model names validated against available models -- Provider type restricted to "ollama" or "lmstudio" -- File paths sanitized (for future vision endpoints) -- Error messages don't expose sensitive information - ---- - -## Performance Notes - -**Caching:** -- Model lists are cached (5-minute TTL) -- First request may be slower (initializes services) -- Subsequent requests use cached model lists - -**Generation Performance:** -- Depends on model size, prompt length, hardware -- Streaming provides perceived faster response (tokens appear immediately) -- Non-streaming waits for complete response - -**Resource Management:** -- Temp files cleaned up automatically (LM Studio vision) -- SSE connections closed after completion -- Graceful error handling prevents resource leaks - -**Optimization Tips:** -- Use smaller models for faster responses -- Lower `max_tokens` for shorter outputs -- Adjust `temperature` for creativity vs consistency -- Use caching for repeated model list requests - ---- - -## Troubleshooting - -### "Ollama is not available" -- Check if Ollama is running: `ollama list` -- Verify `enable_ollama` setting is true -- Check custom URL if configured - -### "LM Studio is not available" -- Verify LM Studio is running -- Check if server is enabled in LM Studio settings -- Verify `enable_lmstudio` setting is true - -### Empty model lists -- Ensure at least one model is installed -- Check service logs for errors -- Try refreshing cache (restart ComfyUI) - -### "Missing required fields" error -- Verify request body includes: `provider`, `model`, `prompt` -- For vision endpoints, also include `images` array -- Check JSON syntax is valid -- Ensure Content-Type header is `application/json` - -### SSE Stream Not Connecting -- Check browser supports EventSource API -- For curl, use `-N` flag for no buffering -- Verify Content-Type is `text/event-stream` -- Check firewall/proxy settings - -### Vision Generation Errors -- Ensure model supports vision (llava, bakllava, etc.) -- Verify image is valid base64-encoded PNG/JPEG -- Check image size (very large images may fail) -- Remove data URI prefix if present (or keep it, both work) - -### "Stream ended without completion" -- Check provider logs for errors -- Verify model is installed and accessible -- Check disk space (temp files for LM Studio) -- Ensure stable network connection to provider - ---- - -## Example Usage - -### JavaScript Fetch Example: -```javascript -// Get status -const status = await fetch('/sage_llm/status').then(r => r.json()); -console.log('Ollama available:', status.data.ollama.available); - -// Get models -const models = await fetch('/sage_llm/models').then(r => r.json()); -console.log('Ollama models:', models.data.models.ollama); - -// Generate text (non-streaming) -const response = await fetch('/sage_llm/generate', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ - provider: 'ollama', - model: 'llama3', - prompt: 'Write a haiku about coding', - options: { temperature: 0.9, seed: 42 } - }) -}).then(r => r.json()); -console.log('Response:', response.data.response); - -// Generate text (streaming SSE) -const eventSource = new EventSource('/sage_llm/generate_stream?' + new URLSearchParams({ - provider: 'ollama', - model: 'llama3', - prompt: 'Write a story', - options: JSON.stringify({ temperature: 0.8 }) -})); - -eventSource.onmessage = (event) => { - const data = JSON.parse(event.data); - if (data.chunk) { - appendToOutput(data.chunk); // Your UI update function - } - if (data.done) { - console.log('Complete'); - eventSource.close(); - } -}; - -// Vision generation (with image upload) -async function generateWithImage(imageFile) { - // Convert to base64 - const base64 = await new Promise((resolve) => { - const reader = new FileReader(); - reader.onload = (e) => resolve(e.target.result.split(',')[1]); - reader.readAsDataURL(imageFile); - }); - - // Generate - const response = await fetch('/sage_llm/vision_generate', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ - provider: 'ollama', - model: 'llava', - prompt: 'Describe this image', - images: [base64] - }) - }).then(r => r.json()); - - console.log('Description:', response.data.response); -} -``` - ---- - -## Route Registration - -Routes are automatically registered via the modular route system: -1. `routes/__init__.py` imports `llm_routes` -2. `register_routes()` called on ComfyUI startup -3. Routes added to PromptServer instance -4. Available immediately at `http://localhost:8188/sage_llm/*` - -Logging confirms registration: -``` -SageUtils: Registered 9 LLM routes (Phase 1: 5, Phase 2: 4) -``` - -**Phase 1 Routes (5):** -- GET `/sage_llm/status` -- GET `/sage_llm/models` -- GET `/sage_llm/vision_models` -- GET `/sage_llm/prompts` -- POST `/sage_llm/generate` - -**Phase 2 Routes (4):** -- POST `/sage_llm/generate_stream` -- POST `/sage_llm/vision_generate` -- POST `/sage_llm/vision_generate_stream` -- *(Future: POST `/sage_llm/stop`)* - -```` +- Keep aliases supported only at route boundary. +- Keep all internal service/provider dispatch canonical. +- Add new providers through shared provider-key and registry/service integration paths documented in `docs/LLM_ARCHITECTURE_NOTE.md`. diff --git a/routes/llm_routes.py b/routes/llm_routes.py index ea25a3b..f5ac592 100644 --- a/routes/llm_routes.py +++ b/routes/llm_routes.py @@ -21,12 +21,25 @@ from ..utils.settings import get_setting from ..utils.llm import clean_response, routes_helpers, set_llm_error_reporter from ..utils.llm import service as llm from ..utils.llm.routes_helpers import get_compatible_models, get_available_presets_full +from ..utils.llm.provider_keys import ( + LMSTUDIO_REST_KEY, + NATIVE_KEY, + OLLAMA_REST_KEY, + OPENAI_KEY, + ROUTE_PROVIDER_KEYS, +) from ..utils.config_manager import llm_prompts from .base import route_error_handler, success_response, error_response, validate_json_body logger = get_logger('routes.llm') +STATUS_KEY_SUFFIX = '_available' + + +def _status_key(provider_key: str) -> str: + return f'{provider_key}{STATUS_KEY_SUFFIX}' + # Route list for documentation and registration tracking _route_list = [] _last_llm_error = ContextVar('sageutils_last_llm_error', default=None) @@ -276,10 +289,10 @@ def register_routes(routes_instance): } return success_response(data={ - "lmstudio_rest": lmstudio_rest_info, - "ollama_rest": ollama_rest_info, - "openai": openai_info, - "native": native_info, + LMSTUDIO_REST_KEY: lmstudio_rest_info, + OLLAMA_REST_KEY: ollama_rest_info, + OPENAI_KEY: openai_info, + NATIVE_KEY: native_info, }) except Exception as e: @@ -349,9 +362,9 @@ def register_routes(routes_instance): if lmstudio_rest_enabled and llm.LMSTUDIO_REST_AVAILABLE: try: models = llm.get_lmstudio_rest_models() - lmstudio_rest_models = get_compatible_models('lmstudio_rest', models) - lmstudio_rest_tool_models = get_compatible_models('lmstudio_rest', llm.get_lmstudio_rest_tool_models()) - lmstudio_rest_reasoning_models = get_compatible_models('lmstudio_rest', llm.get_lmstudio_rest_reasoning_models()) + lmstudio_rest_models = get_compatible_models(LMSTUDIO_REST_KEY, models) + lmstudio_rest_tool_models = get_compatible_models(LMSTUDIO_REST_KEY, llm.get_lmstudio_rest_tool_models()) + lmstudio_rest_reasoning_models = get_compatible_models(LMSTUDIO_REST_KEY, llm.get_lmstudio_rest_reasoning_models()) lmstudio_rest_capabilities = llm.get_lmstudio_rest_model_capabilities_map() except Exception as e: logger.warning(f"Failed to get LM Studio REST models: {e}") @@ -362,9 +375,9 @@ def register_routes(routes_instance): if ollama_rest_enabled and llm.OLLAMA_REST_AVAILABLE: try: models = llm.get_ollama_rest_models() - ollama_rest_models = get_compatible_models('ollama_rest', models) - ollama_rest_tool_models = get_compatible_models('ollama_rest', llm.get_ollama_rest_tool_models()) - ollama_rest_reasoning_models = get_compatible_models('ollama_rest', llm.get_ollama_rest_reasoning_models()) + ollama_rest_models = get_compatible_models(OLLAMA_REST_KEY, models) + ollama_rest_tool_models = get_compatible_models(OLLAMA_REST_KEY, llm.get_ollama_rest_tool_models()) + ollama_rest_reasoning_models = get_compatible_models(OLLAMA_REST_KEY, llm.get_ollama_rest_reasoning_models()) ollama_rest_capabilities = llm.get_ollama_rest_model_capabilities_map() except Exception as e: logger.warning(f"Failed to get Ollama REST models: {e}") @@ -372,9 +385,9 @@ def register_routes(routes_instance): if openai_enabled and llm.OPENAI_AVAILABLE: try: models = llm.get_openai_models() - openai_models = get_compatible_models('openai', models) - openai_tool_models = get_compatible_models('openai', llm.get_openai_tool_models()) - openai_reasoning_models = get_compatible_models('openai', llm.get_openai_reasoning_models()) + openai_models = get_compatible_models(OPENAI_KEY, models) + openai_tool_models = get_compatible_models(OPENAI_KEY, llm.get_openai_tool_models()) + openai_reasoning_models = get_compatible_models(OPENAI_KEY, llm.get_openai_reasoning_models()) openai_capabilities = llm.get_openai_model_capabilities_map() except Exception as e: logger.warning(f"Failed to get OpenAI models: {e}") @@ -382,7 +395,7 @@ def register_routes(routes_instance): native_capabilities = { model_name: { 'name': model_name, - 'provider': 'native', + 'provider': NATIVE_KEY, 'vision': False, 'tool_use': False, 'reasoning': False, @@ -395,34 +408,34 @@ def register_routes(routes_instance): return success_response(data={ "models": { - "lmstudio_rest": lmstudio_rest_models, - "ollama_rest": ollama_rest_models, - "openai": openai_models, - "native": native_models, + LMSTUDIO_REST_KEY: lmstudio_rest_models, + OLLAMA_REST_KEY: ollama_rest_models, + OPENAI_KEY: openai_models, + NATIVE_KEY: native_models, }, "capabilities": { - "lmstudio_rest": lmstudio_rest_capabilities, - "ollama_rest": ollama_rest_capabilities, - "openai": openai_capabilities, - "native": native_capabilities, + LMSTUDIO_REST_KEY: lmstudio_rest_capabilities, + OLLAMA_REST_KEY: ollama_rest_capabilities, + OPENAI_KEY: openai_capabilities, + NATIVE_KEY: native_capabilities, }, "tool_models": { - "lmstudio_rest": lmstudio_rest_tool_models, - "ollama_rest": ollama_rest_tool_models, - "openai": openai_tool_models, - "native": [], + LMSTUDIO_REST_KEY: lmstudio_rest_tool_models, + OLLAMA_REST_KEY: ollama_rest_tool_models, + OPENAI_KEY: openai_tool_models, + NATIVE_KEY: [], }, "reasoning_models": { - "lmstudio_rest": lmstudio_rest_reasoning_models, - "ollama_rest": ollama_rest_reasoning_models, - "openai": openai_reasoning_models, - "native": [], + LMSTUDIO_REST_KEY: lmstudio_rest_reasoning_models, + OLLAMA_REST_KEY: ollama_rest_reasoning_models, + OPENAI_KEY: openai_reasoning_models, + NATIVE_KEY: [], }, "status": { - "lmstudio_rest_available": len(lmstudio_rest_models) > 0, - "ollama_rest_available": len(ollama_rest_models) > 0, - "openai_available": len(openai_models) > 0, - "native_available": len(native_models) > 0, + _status_key(LMSTUDIO_REST_KEY): len(lmstudio_rest_models) > 0, + _status_key(OLLAMA_REST_KEY): len(ollama_rest_models) > 0, + _status_key(OPENAI_KEY): len(openai_models) > 0, + _status_key(NATIVE_KEY): len(native_models) > 0, } }) @@ -493,9 +506,9 @@ def register_routes(routes_instance): if lmstudio_rest_enabled and llm.LMSTUDIO_REST_AVAILABLE: try: models = llm.get_lmstudio_rest_vision_models() - lmstudio_rest_models = get_compatible_models('lmstudio_rest', models) - lmstudio_rest_tool_models = get_compatible_models('lmstudio_rest', llm.get_lmstudio_rest_tool_models()) - lmstudio_rest_reasoning_models = get_compatible_models('lmstudio_rest', llm.get_lmstudio_rest_reasoning_models()) + lmstudio_rest_models = get_compatible_models(LMSTUDIO_REST_KEY, models) + lmstudio_rest_tool_models = get_compatible_models(LMSTUDIO_REST_KEY, llm.get_lmstudio_rest_tool_models()) + lmstudio_rest_reasoning_models = get_compatible_models(LMSTUDIO_REST_KEY, llm.get_lmstudio_rest_reasoning_models()) lmstudio_rest_capabilities = llm.get_lmstudio_rest_model_capabilities_map() except Exception as e: logger.warning(f"Failed to get LM Studio REST vision models: {e}") @@ -506,9 +519,9 @@ def register_routes(routes_instance): if ollama_rest_enabled and llm.OLLAMA_REST_AVAILABLE: try: models = llm.get_ollama_rest_vision_models() - ollama_rest_models = get_compatible_models('ollama_rest', models) - ollama_rest_tool_models = get_compatible_models('ollama_rest', llm.get_ollama_rest_tool_models()) - ollama_rest_reasoning_models = get_compatible_models('ollama_rest', llm.get_ollama_rest_reasoning_models()) + ollama_rest_models = get_compatible_models(OLLAMA_REST_KEY, models) + ollama_rest_tool_models = get_compatible_models(OLLAMA_REST_KEY, llm.get_ollama_rest_tool_models()) + ollama_rest_reasoning_models = get_compatible_models(OLLAMA_REST_KEY, llm.get_ollama_rest_reasoning_models()) ollama_rest_capabilities = llm.get_ollama_rest_model_capabilities_map() except Exception as e: logger.warning(f"Failed to get Ollama REST vision models: {e}") @@ -516,9 +529,9 @@ def register_routes(routes_instance): if openai_enabled and llm.OPENAI_AVAILABLE: try: models = llm.get_openai_vision_models() - openai_vision_models = get_compatible_models('openai', models) - openai_tool_models = get_compatible_models('openai', llm.get_openai_tool_models()) - openai_reasoning_models = get_compatible_models('openai', llm.get_openai_reasoning_models()) + openai_vision_models = get_compatible_models(OPENAI_KEY, models) + openai_tool_models = get_compatible_models(OPENAI_KEY, llm.get_openai_tool_models()) + openai_reasoning_models = get_compatible_models(OPENAI_KEY, llm.get_openai_reasoning_models()) openai_capabilities = llm.get_openai_model_capabilities_map() except Exception as e: logger.warning(f"Failed to get OpenAI vision models: {e}") @@ -527,34 +540,34 @@ def register_routes(routes_instance): return success_response(data={ "models": { - "lmstudio_rest": lmstudio_rest_models, - "ollama_rest": ollama_rest_models, - "openai": openai_vision_models, - "native": native_models, + LMSTUDIO_REST_KEY: lmstudio_rest_models, + OLLAMA_REST_KEY: ollama_rest_models, + OPENAI_KEY: openai_vision_models, + NATIVE_KEY: native_models, }, "capabilities": { - "lmstudio_rest": lmstudio_rest_capabilities, - "ollama_rest": ollama_rest_capabilities, - "openai": openai_capabilities, - "native": native_capabilities, + LMSTUDIO_REST_KEY: lmstudio_rest_capabilities, + OLLAMA_REST_KEY: ollama_rest_capabilities, + OPENAI_KEY: openai_capabilities, + NATIVE_KEY: native_capabilities, }, "tool_models": { - "lmstudio_rest": lmstudio_rest_tool_models, - "ollama_rest": ollama_rest_tool_models, - "openai": openai_tool_models, - "native": [], + LMSTUDIO_REST_KEY: lmstudio_rest_tool_models, + OLLAMA_REST_KEY: ollama_rest_tool_models, + OPENAI_KEY: openai_tool_models, + NATIVE_KEY: [], }, "reasoning_models": { - "lmstudio_rest": lmstudio_rest_reasoning_models, - "ollama_rest": ollama_rest_reasoning_models, - "openai": openai_reasoning_models, - "native": [], + LMSTUDIO_REST_KEY: lmstudio_rest_reasoning_models, + OLLAMA_REST_KEY: ollama_rest_reasoning_models, + OPENAI_KEY: openai_reasoning_models, + NATIVE_KEY: [], }, "status": { - "lmstudio_rest_available": len(lmstudio_rest_models) > 0, - "ollama_rest_available": len(ollama_rest_models) > 0, - "openai_available": len(openai_vision_models) > 0, - "native_available": False, + _status_key(LMSTUDIO_REST_KEY): len(lmstudio_rest_models) > 0, + _status_key(OLLAMA_REST_KEY): len(ollama_rest_models) > 0, + _status_key(OPENAI_KEY): len(openai_vision_models) > 0, + _status_key(NATIVE_KEY): False, } }) @@ -667,13 +680,13 @@ def register_routes(routes_instance): # Generate response based on provider response_text = "" - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: return _error_response_with_metadata( 'LM Studio REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='generate', ) @@ -685,13 +698,13 @@ def register_routes(routes_instance): system_prompt=system_prompt, ) - elif provider == "ollama_rest": + elif provider == OLLAMA_REST_KEY: if not llm.OLLAMA_REST_AVAILABLE: return _error_response_with_metadata( 'Ollama REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='ollama_rest', + provider=OLLAMA_REST_KEY, operation='generate', ) @@ -702,13 +715,13 @@ def register_routes(routes_instance): system_prompt=system_prompt, ) - elif provider == "openai": + elif provider == OPENAI_KEY: if not llm.OPENAI_AVAILABLE: return _error_response_with_metadata( 'OpenAI provider is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='openai', + provider=OPENAI_KEY, operation='generate', ) @@ -719,14 +732,14 @@ def register_routes(routes_instance): system_prompt=system_prompt, ) - elif provider == "native": + elif provider == NATIVE_KEY: available_native = _get_native_clip_models() if model not in available_native: return _error_response_with_metadata( f"Native CLIP model '{model}' is not available", status=404, error_code='LLM_MODEL_NOT_FOUND', - provider='native', + provider=NATIVE_KEY, operation='generate', ) @@ -804,7 +817,7 @@ def register_routes(routes_instance): options = routes_helpers.build_generation_payload_options(provider, data) # Validate provider - if provider not in ["lmstudio_rest", "ollama_rest", "openai", "native"]: + if provider not in ROUTE_PROVIDER_KEYS: return _error_response_with_metadata( f"Invalid provider: {provider}. Must be 'lmstudio', 'ollama', 'openai', or 'native'", status=400, @@ -823,12 +836,12 @@ def register_routes(routes_instance): try: # Generate streaming response based on provider - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: error_chunk = _sse_error_chunk( 'LM Studio REST is not available', error_code='LLM_PROVIDER_UNAVAILABLE', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -848,12 +861,12 @@ def register_routes(routes_instance): if chunk_data.get("done", False): break - elif provider == "ollama_rest": + elif provider == OLLAMA_REST_KEY: if not llm.OLLAMA_REST_AVAILABLE: error_chunk = _sse_error_chunk( 'Ollama REST is not available', error_code='LLM_PROVIDER_UNAVAILABLE', - provider='ollama_rest', + provider=OLLAMA_REST_KEY, operation='generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -872,12 +885,12 @@ def register_routes(routes_instance): if chunk_data.get("done", False): break - elif provider == "openai": + elif provider == OPENAI_KEY: if not llm.OPENAI_AVAILABLE: error_chunk = _sse_error_chunk( 'OpenAI provider is not available', error_code='LLM_PROVIDER_UNAVAILABLE', - provider='openai', + provider=OPENAI_KEY, operation='generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -896,13 +909,13 @@ def register_routes(routes_instance): if chunk_data.get("done", False): break - elif provider == "native": + elif provider == NATIVE_KEY: available_native = _get_native_clip_models() if model not in available_native: error_chunk = _sse_error_chunk( f"Native CLIP model '{model}' is not available", error_code='LLM_MODEL_NOT_FOUND', - provider='native', + provider=NATIVE_KEY, operation='generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -1019,13 +1032,13 @@ def register_routes(routes_instance): response_text = "" - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: return _error_response_with_metadata( 'LM Studio REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='vision_generate', ) @@ -1038,13 +1051,13 @@ def register_routes(routes_instance): system_prompt=system_prompt, ) - elif provider == "ollama_rest": + elif provider == OLLAMA_REST_KEY: if not llm.OLLAMA_REST_AVAILABLE: return _error_response_with_metadata( 'Ollama REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='ollama_rest', + provider=OLLAMA_REST_KEY, operation='vision_generate', ) @@ -1056,13 +1069,13 @@ def register_routes(routes_instance): system_prompt=system_prompt, ) - elif provider == "openai": + elif provider == OPENAI_KEY: if not llm.OPENAI_AVAILABLE: return _error_response_with_metadata( 'OpenAI provider is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='openai', + provider=OPENAI_KEY, operation='vision_generate', ) @@ -1152,12 +1165,12 @@ def register_routes(routes_instance): await response.prepare(request) try: - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: error_chunk = _sse_error_chunk( 'LM Studio REST is not available', error_code='LLM_PROVIDER_UNAVAILABLE', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='vision_generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -1178,12 +1191,12 @@ def register_routes(routes_instance): if chunk_data.get('done', False): break - elif provider == "ollama_rest": + elif provider == OLLAMA_REST_KEY: if not llm.OLLAMA_REST_AVAILABLE: error_chunk = _sse_error_chunk( 'Ollama REST is not available', error_code='LLM_PROVIDER_UNAVAILABLE', - provider='ollama_rest', + provider=OLLAMA_REST_KEY, operation='vision_generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -1203,12 +1216,12 @@ def register_routes(routes_instance): if chunk_data.get('done', False): break - elif provider == "openai": + elif provider == OPENAI_KEY: if not llm.OPENAI_AVAILABLE: error_chunk = _sse_error_chunk( 'OpenAI provider is not available', error_code='LLM_PROVIDER_UNAVAILABLE', - provider='openai', + provider=OPENAI_KEY, operation='vision_generate_stream', ) await response.write(error_chunk.encode('utf-8')) @@ -1690,7 +1703,7 @@ def register_routes(routes_instance): # Get all presets builtin_presets = { 'descriptive_prompt': { - 'provider': 'lmstudio_rest', + 'provider': LMSTUDIO_REST_KEY, 'model': 'gemma3:12b', 'promptTemplate': 'description/Descriptive Prompt', 'systemPrompt': 'e621_prompt_generator', @@ -1702,7 +1715,7 @@ def register_routes(routes_instance): } }, 'e621_description': { - 'provider': 'lmstudio_rest', + 'provider': LMSTUDIO_REST_KEY, 'model': 'gemma3:12b', 'promptTemplate': 'description/Descriptive Prompt', 'systemPrompt': 'e621_prompt_generator', @@ -1714,7 +1727,7 @@ def register_routes(routes_instance): } }, 'casual_chat': { - 'provider': 'lmstudio_rest', + 'provider': LMSTUDIO_REST_KEY, 'model': None, 'promptTemplate': '', 'systemPrompt': 'default', @@ -1740,7 +1753,7 @@ def register_routes(routes_instance): return error_response(f"Preset '{preset_id}' not found", status=404) # Get provider and model - provider = routes_helpers.normalize_provider(preset.get('provider', 'lmstudio_rest')) + provider = routes_helpers.normalize_provider(preset.get('provider', LMSTUDIO_REST_KEY)) model = preset.get('model') if not model: @@ -1804,7 +1817,7 @@ def register_routes(routes_instance): } # Add provider-specific options - if provider == 'ollama_rest': + if provider == OLLAMA_REST_KEY: if 'top_k' in settings: options['top_k'] = settings['top_k'] if 'top_p' in settings: @@ -1818,7 +1831,7 @@ def register_routes(routes_instance): # Generate response response_text = "" - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: return error_response("LM Studio REST is not available", status=503) @@ -1831,7 +1844,7 @@ def register_routes(routes_instance): system_prompt=system_prompt_text, ) - elif provider == "ollama_rest": + elif provider == OLLAMA_REST_KEY: if not llm.OLLAMA_REST_AVAILABLE: return error_response("Ollama REST is not available", status=503) @@ -1844,7 +1857,7 @@ def register_routes(routes_instance): keep_alive=str(settings.get('keepAlive', '5m') or '5m'), ) - elif provider == "openai": + elif provider == OPENAI_KEY: if not llm.OPENAI_AVAILABLE: return error_response("OpenAI provider is not available", status=503) @@ -1918,7 +1931,7 @@ def register_routes(routes_instance): model = str(data.get("model", "")) keep_alive = data.get("keep_alive", 60) - if not provider or provider not in ("lmstudio_rest", "ollama_rest", "openai", "native"): + if not provider or provider not in ROUTE_PROVIDER_KEYS: return _error_response_with_metadata( f"Invalid or missing provider: {provider!r}. Must be 'lmstudio', 'ollama', 'openai', or 'native'", status=400, @@ -1935,13 +1948,13 @@ def register_routes(routes_instance): llm.ensure_llm_initialized(force=True) - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: return _error_response_with_metadata( 'LM Studio REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='load_model', ) # LM Studio REST can return load failures even when the model is already @@ -1953,42 +1966,42 @@ def register_routes(routes_instance): f"Model '{model}' is not available in LM Studio REST", status=404, error_code='LLM_MODEL_NOT_FOUND', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='load_model', ) - elif provider == "ollama_rest": + elif provider == OLLAMA_REST_KEY: if not llm.OLLAMA_REST_AVAILABLE: return _error_response_with_metadata( 'Ollama REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='ollama_rest', + provider=OLLAMA_REST_KEY, operation='load_model', ) # No dedicated preload endpoint in Ollama REST mode here. # Keep behavior consistent with native/openai by validating selection. - elif provider == "openai": + elif provider == OPENAI_KEY: if not llm.OPENAI_AVAILABLE: return _error_response_with_metadata( 'OpenAI provider is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='openai', + provider=OPENAI_KEY, operation='load_model', ) # OpenAI-compatible providers do not expose a standard load API. # Treat this as a successful readiness check. - elif provider == "native": + elif provider == NATIVE_KEY: available_native = _get_native_clip_models() if model not in available_native: return _error_response_with_metadata( f"Native CLIP model '{model}' is not available", status=404, error_code='LLM_MODEL_NOT_FOUND', - provider='native', + provider=NATIVE_KEY, operation='load_model', ) # Native models are loaded on demand for generation; this validates selection. @@ -2063,13 +2076,13 @@ def register_routes(routes_instance): response_text = "" - if provider == "lmstudio_rest": + if provider == LMSTUDIO_REST_KEY: if not llm.LMSTUDIO_REST_AVAILABLE: return _error_response_with_metadata( 'LM Studio REST is not available', status=503, error_code='LLM_PROVIDER_UNAVAILABLE', - provider='lmstudio_rest', + provider=LMSTUDIO_REST_KEY, operation='generate_only', ) response_text = llm.lmstudio_rest_generate( @@ -2080,14 +2093,14 @@ def register_routes(routes_instance): system_prompt=system_prompt, ) - elif provider == "native": + elif provider == NATIVE_KEY: available_native = _get_native_clip_models() if model not in available_native: return _error_response_with_metadata( f"Native CLIP model '{model}' is not available", status=404, error_code='LLM_MODEL_NOT_FOUND', - provider='native', + provider=NATIVE_KEY, operation='generate_only', ) diff --git a/utils/llm/compat.py b/utils/llm/compat.py new file mode 100644 index 0000000..fa4a409 --- /dev/null +++ b/utils/llm/compat.py @@ -0,0 +1,4 @@ +"""Compatibility module retained for import stability during refactor migration.""" + + + diff --git a/utils/llm/init.py b/utils/llm/init.py index 09a3880..98847f8 100644 --- a/utils/llm/init.py +++ b/utils/llm/init.py @@ -1,7 +1,5 @@ """Initialization helpers for LLM provider clients.""" -from typing import Any - from ..logger import get_logger from .providers.lmstudio import client as lmstudio_rest_provider from .providers.ollama import client as ollama_rest_provider @@ -64,36 +62,3 @@ def init_openai_provider(enabled: bool) -> bool: return False -def init_ollama_client( - ollama_available: bool, - ollama_module: Any, - enabled: bool, - custom_url: str, -) -> tuple[Any, bool]: - """Compatibility alias for legacy callers. - - SDK initialization was removed. This now initializes the Ollama REST provider - and returns a tuple compatible with the old signature. - """ - _ = ollama_available - _ = ollama_module - _ = custom_url - initialized = init_ollama_rest(enabled) - return None, initialized - - -def init_lmstudio_client( - lmstudio_available: bool, - lms_module: Any, - enabled: bool, - custom_url: str, -) -> bool: - """Compatibility alias for legacy callers. - - SDK initialization was removed. This now initializes the LM Studio REST - provider and keeps the old function signature for compatibility. - """ - _ = lmstudio_available - _ = lms_module - _ = custom_url - return init_lmstudio_rest(enabled) diff --git a/utils/llm/provider_keys.py b/utils/llm/provider_keys.py new file mode 100644 index 0000000..8b3aee7 --- /dev/null +++ b/utils/llm/provider_keys.py @@ -0,0 +1,30 @@ +"""Canonical provider keys and alias normalization helpers.""" + +LMSTUDIO_REST_KEY = 'lmstudio_rest' +OLLAMA_REST_KEY = 'ollama_rest' +OPENAI_KEY = 'openai' +NATIVE_KEY = 'native' + +_PROVIDER_ALIAS_MAP = { + 'lmstudio': LMSTUDIO_REST_KEY, + 'ollama': OLLAMA_REST_KEY, +} + +SERVICE_PROVIDER_KEYS = ( + LMSTUDIO_REST_KEY, + OLLAMA_REST_KEY, + OPENAI_KEY, +) + +ROUTE_PROVIDER_KEYS = ( + LMSTUDIO_REST_KEY, + OLLAMA_REST_KEY, + OPENAI_KEY, + NATIVE_KEY, +) + + +def normalize_provider_key(provider: str) -> str: + """Normalize provider aliases to canonical backend keys.""" + normalized = (provider or '').strip().lower() + return _PROVIDER_ALIAS_MAP.get(normalized, normalized) diff --git a/utils/llm/providers/availability.py b/utils/llm/providers/availability.py new file mode 100644 index 0000000..0b7a594 --- /dev/null +++ b/utils/llm/providers/availability.py @@ -0,0 +1,84 @@ +"""Shared provider availability helpers.""" + +from typing import Any, Callable, Type + + +def is_provider_unavailable(enabled: bool) -> bool: + """Return True when a provider is disabled/unavailable by settings.""" + return not enabled + + +def unavailable_models_placeholder(message: str) -> list[str]: + """Return standard unavailable placeholder model list for provider clients.""" + return [message] + + +def report_fetch_error( + report_func: Callable[..., None], + message: str, + *, + provider: str, + operation: str, + cause: Any, + fallback, +): + """Report a provider fetch error and return caller-defined fallback value.""" + report_func(message, provider=provider, operation=operation, cause=cause) + return fallback + + +def raise_if_provider_unavailable( + enabled: bool, + raise_func: Callable[..., None], + *, + error_type: Type[Exception], + message: str, + provider: str, + operation: str, +) -> None: + """Raise provider-scoped unavailable errors through the shared availability gate.""" + if is_provider_unavailable(enabled): + raise_func(error_type, message, provider=provider, operation=operation) + + +def raise_if_model_unavailable( + model: str, + model_list: list[str], + raise_func: Callable[..., None], + *, + provider: str, + operation: str, +) -> None: + """Raise provider-scoped model-not-found errors using a shared message shape.""" + if model not in model_list: + raise_func( + ValueError, + f"Model '{model}' is not available. Available models: {model_list}", + provider=provider, + operation=operation, + ) + + +def raise_if_missing_images( + images, + raise_func: Callable[..., None], + *, + provider: str, + operation: str, + message: str = 'No images provided for vision model.', +) -> None: + """Raise provider-scoped image-required errors for vision endpoints.""" + if images is None: + raise_func(ValueError, message, provider=provider, operation=operation) + + +def stream_error_payload(error_text: str, *, include_full_response: bool = True) -> dict[str, object]: + """Build a standard stream terminal error payload for provider wrappers.""" + payload: dict[str, object] = { + 'chunk': '', + 'done': True, + 'error': error_text, + } + if include_full_response: + payload['full_response'] = '' + return payload \ No newline at end of file diff --git a/utils/llm/providers/lmstudio/client.py b/utils/llm/providers/lmstudio/client.py index d91a9a3..1f0f1ab 100644 --- a/utils/llm/providers/lmstudio/client.py +++ b/utils/llm/providers/lmstudio/client.py @@ -5,6 +5,15 @@ from ...common import clean_response from ...cache import get_llm_cache from ...rest import iter_sse_events, normalize_image_data_url from ...errors import llm_report, llm_raise, llm_stringify +from ..availability import ( + is_provider_unavailable, + unavailable_models_placeholder, + report_fetch_error, + raise_if_provider_unavailable, + raise_if_model_unavailable, + raise_if_missing_images, + stream_error_payload, +) from .requests import ( lmstudio_request_json_models, @@ -32,12 +41,6 @@ _PROVIDER_NAME = 'lmstudio_rest' _UNAVAILABLE_MESSAGE = '(LM Studio REST not available)' -def _unavailable_models() -> list[str]: - return [_UNAVAILABLE_MESSAGE] - -def _is_unavailable(enabled: bool) -> bool: - return not enabled - _LMSTUDIO_OPTION_KEY_MAP = ( ('temperature', 'temperature'), ('max_tokens', 'maxTokens'), @@ -81,6 +84,23 @@ def _build_chat_options(options: Optional[dict[str, Any]]) -> dict[str, Any]: return payload_options + +def _build_lmstudio_stream_payload( + model: str, + input_value, + *, + options: Optional[dict[str, Any]] = None, + system_prompt: str = '', +) -> dict[str, Any]: + payload = { + 'model': model, + 'input': input_value, + } + if system_prompt: + payload['system_prompt'] = system_prompt + payload.update(_build_chat_options(options)) + return payload + def build_lmstudio_config(options: dict) -> dict: """Build LM Studio configuration from options dictionary.""" config = {} @@ -131,7 +151,7 @@ def _build_progress_event_payload(event_type: str, event_data: Any) -> dict[str, def is_running(enabled: bool) -> bool: """Check if LM Studio REST server is reachable.""" - if _is_unavailable(enabled): + if is_provider_unavailable(enabled): return False try: @@ -142,8 +162,8 @@ def is_running(enabled: bool) -> bool: def get_models(enabled: bool) -> list[str]: """Retrieve a list of available models from LM Studio REST.""" - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) def _fetch_models() -> list[str]: try: @@ -156,8 +176,14 @@ def get_models(enabled: bool) -> list[str]: models.append(model_name) return models except Exception as e: - llm_report('Error retrieving models from LM Studio REST', provider='lmstudio_rest', operation='get_models', cause=e) - return _unavailable_models() + return report_fetch_error( + llm_report, + 'Error retrieving models from LM Studio REST', + provider=_PROVIDER_NAME, + operation='get_models', + cause=e, + fallback=unavailable_models_placeholder(_UNAVAILABLE_MESSAGE), + ) cache = get_llm_cache() return cache.get_model_list( @@ -169,8 +195,8 @@ def get_models(enabled: bool) -> list[str]: def get_vision_models(enabled: bool) -> list[str]: """Retrieve a list of available vision models from LM Studio REST.""" - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) def _fetch_vision_models(cache_instance) -> list[str]: try: @@ -190,8 +216,14 @@ def get_vision_models(enabled: bool) -> list[str]: return vision_models except Exception as e: - llm_report('Error retrieving vision models from LM Studio REST', provider='lmstudio_rest', operation='get_vision_models', cause=e) - return [] + return report_fetch_error( + llm_report, + 'Error retrieving vision models from LM Studio REST', + provider=_PROVIDER_NAME, + operation='get_vision_models', + cause=e, + fallback=[], + ) cache = get_llm_cache() return cache.get_model_list( @@ -203,23 +235,29 @@ def get_vision_models(enabled: bool) -> list[str]: ) def get_tool_models(enabled: bool) -> list[str]: - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) capabilities_map = get_model_capabilities_map(enabled) return sorted([name for name, capabilities in capabilities_map.items() if capabilities.tool_use]) def get_reasoning_models(enabled: bool) -> list[str]: - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) capabilities_map = get_model_capabilities_map(enabled) return sorted([name for name, capabilities in capabilities_map.items() if capabilities.reasoning]) def load_model(enabled: bool, model: str, keep_alive: int = 0) -> bool: """Ask LM Studio REST server to load a model.""" - if _is_unavailable(enabled): - llm_raise(RuntimeError, 'LM Studio REST is not enabled.', provider='lmstudio_rest', operation='load_model') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=RuntimeError, + message='LM Studio REST is not enabled.', + provider=_PROVIDER_NAME, + operation='load_model', + ) # keep_alive is not currently used by LM Studio REST. payload: dict[str, Any] = {'model': model} @@ -228,19 +266,19 @@ def load_model(enabled: bool, model: str, keep_alive: int = 0) -> bool: lmstudio_request_json_load(payload) return True except Exception as e: - llm_raise(RuntimeError, f"Failed to load model '{model}' via LM Studio REST", provider='lmstudio_rest', operation='load_model', cause=RuntimeError(llm_stringify(e))) + llm_raise(RuntimeError, f"Failed to load model '{model}' via LM Studio REST", provider=_PROVIDER_NAME, operation='load_model', cause=RuntimeError(llm_stringify(e))) return False def unload_model(enabled: bool, model: str) -> bool: """Ask LM Studio REST server to unload a model.""" - if _is_unavailable(enabled): + if is_provider_unavailable(enabled): return False try: lmstudio_request_json_unload({'model': model}) return True except Exception as e: - llm_report('Error unloading model via LM Studio REST', provider='lmstudio_rest', operation='unload_model', cause=e) + llm_report('Error unloading model via LM Studio REST', provider=_PROVIDER_NAME, operation='unload_model', cause=e) return False @@ -320,40 +358,65 @@ def _stream_chat_response(payload: dict[str, Any], operation: str): def generate_stream(enabled: bool, model: str, prompt: str, options=None, system_prompt: str = ''): """Generate a real streaming response from LM Studio REST.""" try: - if _is_unavailable(enabled): - llm_raise(ImportError, 'LM Studio REST is not enabled.', provider='lmstudio_rest', operation='generate_stream') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='LM Studio REST is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_stream', + ) model_list = get_models(enabled) - if model not in model_list: - llm_raise(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio_rest', operation='generate_stream') + raise_if_model_unavailable( + model, + model_list, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_stream', + ) - payload = { - 'model': model, - 'input': prompt, - } - if system_prompt: - payload['system_prompt'] = system_prompt - payload.update(_build_chat_options(options)) + payload = _build_lmstudio_stream_payload( + model, + prompt, + options=options, + system_prompt=system_prompt, + ) for chunk_data in _stream_chat_response(payload, 'generate_stream'): yield chunk_data except Exception as e: - llm_report('Error streaming response from LM Studio REST', provider='lmstudio_rest', operation='generate_stream', cause=e) - yield {'chunk': '', 'done': True, 'error': llm_stringify(e)} + llm_report('Error streaming response from LM Studio REST', provider=_PROVIDER_NAME, operation='generate_stream', cause=e) + yield stream_error_payload(llm_stringify(e), include_full_response=False) def generate_vision_stream(enabled: bool, model: str, prompt: str, images, options=None, system_prompt: str = ''): """Generate a real streaming vision response from LM Studio REST.""" try: - if _is_unavailable(enabled): - llm_raise(ImportError, 'LM Studio REST is not enabled.', provider='lmstudio_rest', operation='generate_vision_stream') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='LM Studio REST is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) - if images is None: - llm_raise(ValueError, 'No images provided for vision model.', provider='lmstudio_rest', operation='generate_vision_stream') + raise_if_missing_images( + images, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) model_list = get_vision_models(enabled) - if model not in model_list: - llm_raise(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio_rest', operation='generate_vision_stream') + raise_if_model_unavailable( + model, + model_list, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) image_entries = images if isinstance(images, list) else [images] input_items = [{'type': 'text', 'content': prompt}] @@ -365,19 +428,18 @@ def generate_vision_stream(enabled: bool, model: str, prompt: str, images, optio } ) - payload = { - 'model': model, - 'input': input_items, - } - if system_prompt: - payload['system_prompt'] = system_prompt - payload.update(_build_chat_options(options)) + payload = _build_lmstudio_stream_payload( + model, + input_items, + options=options, + system_prompt=system_prompt, + ) for chunk_data in _stream_chat_response(payload, 'generate_vision_stream'): yield chunk_data except Exception as e: - llm_report('Error streaming vision response from LM Studio REST', provider='lmstudio_rest', operation='generate_vision_stream', cause=e) - yield {'chunk': '', 'done': True, 'error': llm_stringify(e)} + llm_report('Error streaming vision response from LM Studio REST', provider=_PROVIDER_NAME, operation='generate_vision_stream', cause=e) + yield stream_error_payload(llm_stringify(e), include_full_response=False) def generate_with_stream(enabled: bool, model: str, prompt: str, options=None, system_prompt: str = '') -> str: """Generate a response using the streaming API and return the full result when finished or on error.""" @@ -385,7 +447,7 @@ def generate_with_stream(enabled: bool, model: str, prompt: str, options=None, s result = '' error = None for chunk in stream: - print(f"DEBUG: Received stream chunk: {chunk}") + #print(f"DEBUG: Received stream chunk: {chunk}") if chunk.get('chunk'): result += chunk['chunk'] if chunk.get('error'): @@ -394,7 +456,7 @@ def generate_with_stream(enabled: bool, model: str, prompt: str, options=None, s if chunk.get('done'): break if error: - llm_raise(RuntimeError, f"Streaming error: {error}", provider='lmstudio_rest', operation='generate_with_stream') + llm_raise(RuntimeError, f"Streaming error: {error}", provider=_PROVIDER_NAME, operation='generate_with_stream') return clean_response(result) def generate_vision_with_stream(enabled: bool, model: str, prompt: str, images, options=None, system_prompt: str = '') -> str: @@ -403,7 +465,7 @@ def generate_vision_with_stream(enabled: bool, model: str, prompt: str, images, result = '' error = None for chunk in stream: - print(f"DEBUG: Received stream chunk: {chunk}") + #print(f"DEBUG: Received stream chunk: {chunk}") if chunk.get('chunk'): result += chunk['chunk'] if chunk.get('error'): @@ -412,5 +474,5 @@ def generate_vision_with_stream(enabled: bool, model: str, prompt: str, images, if chunk.get('done'): break if error: - llm_raise(RuntimeError, f"Streaming error: {error}", provider='lmstudio_rest', operation='generate_vision_with_stream') + llm_raise(RuntimeError, f"Streaming error: {error}", provider=_PROVIDER_NAME, operation='generate_vision_with_stream') return clean_response(result) diff --git a/utils/llm/providers/ollama/client.py b/utils/llm/providers/ollama/client.py index 88eba6f..10d2ef0 100644 --- a/utils/llm/providers/ollama/client.py +++ b/utils/llm/providers/ollama/client.py @@ -13,18 +13,21 @@ from ...common import clean_response from ...cache import get_llm_cache from ...rest import iter_json_lines, normalize_raw_image_base64 from ...errors import llm_raise, llm_report, llm_stringify +from ..availability import ( + is_provider_unavailable, + unavailable_models_placeholder, + report_fetch_error, + raise_if_provider_unavailable, + raise_if_model_unavailable, + raise_if_missing_images, + stream_error_payload, +) from ...tensor import tensor_to_base64_safe _PROVIDER_NAME = 'ollama_rest' _UNAVAILABLE_MESSAGE = '(Ollama REST not available)' _MAX_TOOL_ITERATIONS = 4 -def _unavailable_models() -> list[str]: - return [_UNAVAILABLE_MESSAGE] - -def _is_unavailable(enabled: bool) -> bool: - return not enabled - def _get_request_timeout() -> float: """Return request timeout seconds for generation calls. @@ -111,6 +114,34 @@ def _build_ollama_payload(model: str, messages: list[dict[str, Any]], options=No return payload + +def _build_ollama_vision_payload( + model: str, + prompt: str, + images, + *, + options=None, + system_prompt: str = '', + keep_alive: str = '5m', + include_stream_field: bool = False, +) -> dict[str, Any]: + payload: dict[str, Any] = { + 'model': model, + 'messages': _build_messages(prompt, system_prompt, raw_images=_normalize_raw_images(images)), + } + if include_stream_field: + payload['stream'] = False + if keep_alive: + payload['keep_alive'] = keep_alive + + _apply_top_level_controls(payload, options) + + built_options = _build_options(options) + if built_options: + payload['options'] = built_options + + return payload + def _chat_once(payload: dict[str, Any], operation: str) -> dict[str, Any]: response = ollama_request_json_chat(payload, timeout=_get_request_timeout()) @@ -174,7 +205,7 @@ def _stream_chat_response(payload: dict[str, Any], operation: str): def is_running(enabled: bool) -> bool: """Check if the Ollama REST server is reachable.""" - if _is_unavailable(enabled): + if is_provider_unavailable(enabled): return False try: @@ -185,8 +216,8 @@ def is_running(enabled: bool) -> bool: def get_models(enabled: bool) -> list[str]: """Retrieve a list of available models from Ollama REST.""" - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) def _fetch_models() -> list[str]: try: @@ -199,8 +230,14 @@ def get_models(enabled: bool) -> list[str]: models.append(model_name) return models except Exception as e: - llm_report('Error retrieving models from Ollama REST', provider=_PROVIDER_NAME, operation='get_models', cause=e) - return _unavailable_models() + return report_fetch_error( + llm_report, + 'Error retrieving models from Ollama REST', + provider=_PROVIDER_NAME, + operation='get_models', + cause=e, + fallback=unavailable_models_placeholder(_UNAVAILABLE_MESSAGE), + ) cache = get_llm_cache() return cache.get_model_list( @@ -212,8 +249,8 @@ def get_models(enabled: bool) -> list[str]: def get_vision_models(enabled: bool) -> list[str]: """Retrieve a list of available vision models from Ollama REST.""" - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) def _fetch_vision_models(cache_instance) -> list[str]: try: @@ -233,8 +270,14 @@ def get_vision_models(enabled: bool) -> list[str]: return vision_models except Exception as e: - llm_report('Error retrieving vision models from Ollama REST', provider=_PROVIDER_NAME, operation='get_vision_models', cause=e) - return [] + return report_fetch_error( + llm_report, + 'Error retrieving vision models from Ollama REST', + provider=_PROVIDER_NAME, + operation='get_vision_models', + cause=e, + fallback=[], + ) cache = get_llm_cache() return cache.get_model_list( @@ -246,23 +289,29 @@ def get_vision_models(enabled: bool) -> list[str]: ) def get_tool_models(enabled: bool) -> list[str]: - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) capabilities_map = get_model_capabilities_map(enabled, model_names=get_models(enabled)) return sorted([name for name, capabilities in capabilities_map.items() if capabilities.tool_use]) def get_reasoning_models(enabled: bool) -> list[str]: - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) capabilities_map = get_model_capabilities_map(enabled, model_names=get_models(enabled)) return sorted([name for name, capabilities in capabilities_map.items() if capabilities.reasoning]) def load_model(enabled: bool, model: str, keep_alive: int = 60) -> bool: """Warm-load an Ollama model via /api/generate with an empty prompt.""" - if _is_unavailable(enabled): - llm_raise(ImportError, 'Ollama REST is not enabled.', provider=_PROVIDER_NAME, operation='load_model') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='Ollama REST is not enabled.', + provider=_PROVIDER_NAME, + operation='load_model', + ) try: keep_alive_seconds = max(0, int(keep_alive)) @@ -287,8 +336,14 @@ def load_model(enabled: bool, model: str, keep_alive: int = 60) -> bool: def generate(enabled: bool, model: str, prompt: str, options=None, system_prompt: str = '', keep_alive: str = '5m') -> str: """Generate a response from Ollama REST using /api/chat.""" - if _is_unavailable(enabled): - llm_raise(ImportError, 'Ollama REST is not enabled.', provider=_PROVIDER_NAME, operation='generate') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='Ollama REST is not enabled.', + provider=_PROVIDER_NAME, + operation='generate', + ) messages = _build_messages(prompt, system_prompt) @@ -320,27 +375,31 @@ def generate(enabled: bool, model: str, prompt: str, options=None, system_prompt def generate_vision(enabled: bool, model: str, prompt: str, images=None, options=None, system_prompt: str = '', keep_alive: str = '5m') -> str: """Generate a vision response from Ollama REST using /api/chat with inline images.""" - if _is_unavailable(enabled): - llm_raise(ImportError, 'Ollama REST is not enabled.', provider=_PROVIDER_NAME, operation='generate_vision') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='Ollama REST is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_vision', + ) - if images is None: - llm_raise(ValueError, 'No images provided for vision model.', provider=_PROVIDER_NAME, operation='generate_vision') + raise_if_missing_images( + images, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision', + ) - raw_images = _normalize_raw_images(images) - - payload: dict[str, Any] = { - 'model': model, - 'messages': _build_messages(prompt, system_prompt, raw_images=raw_images), - 'stream': False, - } - if keep_alive: - payload['keep_alive'] = keep_alive - - _apply_top_level_controls(payload, options) - - built_options = _build_options(options) - if built_options: - payload['options'] = built_options + payload = _build_ollama_vision_payload( + model, + prompt, + images, + options=options, + system_prompt=system_prompt, + keep_alive=keep_alive, + include_stream_field=True, + ) try: response = ollama_request_json_chat(payload, timeout=_get_request_timeout()) @@ -355,12 +414,23 @@ def generate_vision(enabled: bool, model: str, prompt: str, images=None, options def generate_stream(enabled: bool, model: str, prompt: str, options=None, system_prompt: str = '', keep_alive: str = '5m'): """Generate a real streaming response from Ollama REST.""" try: - if _is_unavailable(enabled): - llm_raise(ImportError, 'Ollama REST is not enabled.', provider=_PROVIDER_NAME, operation='generate_stream') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='Ollama REST is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_stream', + ) model_list = get_models(enabled) - if model not in model_list: - llm_raise(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider=_PROVIDER_NAME, operation='generate_stream') + raise_if_model_unavailable( + model, + model_list, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_stream', + ) messages = _build_messages(prompt, system_prompt) @@ -391,39 +461,51 @@ def generate_stream(enabled: bool, model: str, prompt: str, options=None, system yield chunk_data except Exception as e: llm_report('Error in Ollama REST streaming', provider=_PROVIDER_NAME, operation='generate_stream', cause=e) - yield {'chunk': '', 'done': True, 'full_response': '', 'error': llm_stringify(e)} + yield stream_error_payload(llm_stringify(e), include_full_response=True) def generate_vision_stream(enabled: bool, model: str, prompt: str, images=None, options=None, system_prompt: str = '', keep_alive: str = '5m'): """Generate a real streaming vision response from Ollama REST.""" try: - if _is_unavailable(enabled): - llm_raise(ImportError, 'Ollama REST is not enabled.', provider=_PROVIDER_NAME, operation='generate_vision_stream') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='Ollama REST is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) - if images is None: - llm_raise(ValueError, 'No images provided for vision model.', provider=_PROVIDER_NAME, operation='generate_vision_stream') + raise_if_missing_images( + images, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) model_list = get_vision_models(enabled) - if model not in model_list: - llm_raise(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider=_PROVIDER_NAME, operation='generate_vision_stream') + raise_if_model_unavailable( + model, + model_list, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) - payload: dict[str, Any] = { - 'model': model, - 'messages': _build_messages(prompt, system_prompt, raw_images=_normalize_raw_images(images)), - } - if keep_alive: - payload['keep_alive'] = keep_alive - - _apply_top_level_controls(payload, options) - - built_options = _build_options(options) - if built_options: - payload['options'] = built_options + payload = _build_ollama_vision_payload( + model, + prompt, + images, + options=options, + system_prompt=system_prompt, + keep_alive=keep_alive, + include_stream_field=False, + ) for chunk_data in _stream_chat_response(payload, 'generate_vision_stream'): yield chunk_data except Exception as e: llm_report('Error in Ollama REST vision streaming', provider=_PROVIDER_NAME, operation='generate_vision_stream', cause=e) - yield {'chunk': '', 'done': True, 'full_response': '', 'error': llm_stringify(e)} + yield stream_error_payload(llm_stringify(e), include_full_response=True) def build_response_parameters(model: str, prompt: str, keep_alive: float, options: dict, system_prompt: str, images) -> dict: """Build the response parameters for Ollama generate call.""" diff --git a/utils/llm/providers/openai/client.py b/utils/llm/providers/openai/client.py index 4e49b74..2f9e799 100644 --- a/utils/llm/providers/openai/client.py +++ b/utils/llm/providers/openai/client.py @@ -16,6 +16,14 @@ from ...cache import get_llm_cache from ....logger import get_logger from ...errors import llm_raise, llm_report, llm_stringify from ...rest import iter_sse_events, normalize_image_data_url +from ..availability import ( + is_provider_unavailable, + unavailable_models_placeholder, + report_fetch_error, + raise_if_provider_unavailable, + raise_if_missing_images, + stream_error_payload, +) logger = get_logger('llm.providers.openai') @@ -25,12 +33,6 @@ _UNAVAILABLE_MESSAGE = '(OpenAI not available)' # Progress event types tracked during OpenAI streaming _PROGRESS_EVENT_TYPES = {'generation.start', 'generation.progress', 'generation.end'} -def _unavailable_models() -> list[str]: - return [_UNAVAILABLE_MESSAGE] - -def _is_unavailable(enabled: bool) -> bool: - return not enabled - def _build_messages(prompt: str, system_prompt: str = '') -> list[dict[str, Any]]: messages = [] if system_prompt: @@ -54,6 +56,25 @@ def _build_vision_messages(prompt: str, images, system_prompt: str = '') -> list messages.append({'role': 'user', 'content': content_parts}) return messages + +def _build_openai_vision_payload( + model: str, + prompt: str, + images, + *, + options=None, + system_prompt: str = '', + include_stream_field: bool = False, +) -> dict[str, Any]: + payload: dict[str, Any] = { + 'model': model, + 'messages': _build_vision_messages(prompt, images, system_prompt), + } + if include_stream_field: + payload['stream'] = False + payload.update(_build_options(options)) + return payload + def _build_options(options: Optional[dict[str, Any]]) -> dict[str, Any]: input_options = options or {} result: dict[str, Any] = {} @@ -166,7 +187,7 @@ def _stream_chat_response(payload: dict[str, Any], operation: str): def is_running(enabled: bool) -> bool: """Check if the OpenAI-compatible server is reachable.""" - if _is_unavailable(enabled): + if is_provider_unavailable(enabled): return False try: @@ -177,8 +198,8 @@ def is_running(enabled: bool) -> bool: def get_models(enabled: bool) -> list[str]: """Retrieve a list of available models from the OpenAI-compatible endpoint.""" - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) def _fetch_models() -> list[str]: try: @@ -191,8 +212,14 @@ def get_models(enabled: bool) -> list[str]: models.append(model_id) return sorted(models) except Exception as e: - llm_report('Error retrieving models from OpenAI', provider=_PROVIDER_NAME, operation='get_models', cause=e) - return _unavailable_models() + return report_fetch_error( + llm_report, + 'Error retrieving models from OpenAI', + provider=_PROVIDER_NAME, + operation='get_models', + cause=e, + fallback=unavailable_models_placeholder(_UNAVAILABLE_MESSAGE), + ) cache = get_llm_cache() return cache.get_model_list( @@ -204,8 +231,8 @@ def get_models(enabled: bool) -> list[str]: def get_vision_models(enabled: bool) -> list[str]: """Retrieve a list of available vision models from the OpenAI-compatible endpoint.""" - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) def _fetch_vision_models(cache_instance) -> list[str]: try: @@ -225,8 +252,14 @@ def get_vision_models(enabled: bool) -> list[str]: return sorted(vision_models) except Exception as e: - llm_report('Error retrieving vision models from OpenAI', provider=_PROVIDER_NAME, operation='get_vision_models', cause=e) - return [] + return report_fetch_error( + llm_report, + 'Error retrieving vision models from OpenAI', + provider=_PROVIDER_NAME, + operation='get_vision_models', + cause=e, + fallback=[], + ) cache = get_llm_cache() return cache.get_model_list( @@ -238,23 +271,29 @@ def get_vision_models(enabled: bool) -> list[str]: ) def get_tool_models(enabled: bool) -> list[str]: - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) capabilities_map = get_model_capabilities_map(enabled) return sorted([name for name, capabilities in capabilities_map.items() if capabilities.tool_use]) def get_reasoning_models(enabled: bool) -> list[str]: - if _is_unavailable(enabled): - return _unavailable_models() + if is_provider_unavailable(enabled): + return unavailable_models_placeholder(_UNAVAILABLE_MESSAGE) capabilities_map = get_model_capabilities_map(enabled) return sorted([name for name, capabilities in capabilities_map.items() if capabilities.reasoning]) def generate(enabled: bool, model: str, prompt: str, options=None, system_prompt: str = '') -> str: """Generate a response from an OpenAI-compatible model using /v1/chat/completions.""" - if _is_unavailable(enabled): - llm_raise(ImportError, 'OpenAI provider is not enabled.', provider=_PROVIDER_NAME, operation='generate') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='OpenAI provider is not enabled.', + provider=_PROVIDER_NAME, + operation='generate', + ) payload: dict[str, Any] = { 'model': model, @@ -275,18 +314,30 @@ def generate(enabled: bool, model: str, prompt: str, options=None, system_prompt def generate_vision(enabled: bool, model: str, prompt: str, images=None, options=None, system_prompt: str = '') -> str: """Generate a vision response from an OpenAI-compatible model.""" - if _is_unavailable(enabled): - llm_raise(ImportError, 'OpenAI provider is not enabled.', provider=_PROVIDER_NAME, operation='generate_vision') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='OpenAI provider is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_vision', + ) - if images is None: - llm_raise(ValueError, 'No images provided for vision model.', provider=_PROVIDER_NAME, operation='generate_vision') + raise_if_missing_images( + images, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision', + ) - payload: dict[str, Any] = { - 'model': model, - 'messages': _build_vision_messages(prompt, images, system_prompt), - 'stream': False, - } - payload.update(_build_options(options)) + payload = _build_openai_vision_payload( + model, + prompt, + images, + options=options, + system_prompt=system_prompt, + include_stream_field=True, + ) try: response = openai_request_json_chat(payload) @@ -301,8 +352,14 @@ def generate_vision(enabled: bool, model: str, prompt: str, images=None, options def generate_stream(enabled: bool, model: str, prompt: str, options=None, system_prompt: str = ''): """Generate a real streaming response from an OpenAI-compatible model.""" try: - if _is_unavailable(enabled): - llm_raise(ImportError, 'OpenAI provider is not enabled.', provider=_PROVIDER_NAME, operation='generate_stream') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='OpenAI provider is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_stream', + ) payload: dict[str, Any] = { 'model': model, @@ -314,26 +371,39 @@ def generate_stream(enabled: bool, model: str, prompt: str, options=None, system yield chunk_data except Exception as e: llm_report('Error in OpenAI streaming', provider=_PROVIDER_NAME, operation='generate_stream', cause=e) - yield {'chunk': '', 'done': True, 'full_response': '', 'error': llm_stringify(e)} + yield stream_error_payload(llm_stringify(e), include_full_response=True) def generate_vision_stream(enabled: bool, model: str, prompt: str, images=None, options=None, system_prompt: str = ''): """Generate a real streaming vision response from an OpenAI-compatible model.""" try: - if _is_unavailable(enabled): - llm_raise(ImportError, 'OpenAI provider is not enabled.', provider=_PROVIDER_NAME, operation='generate_vision_stream') + raise_if_provider_unavailable( + enabled, + llm_raise, + error_type=ImportError, + message='OpenAI provider is not enabled.', + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) - if images is None: - llm_raise(ValueError, 'No images provided for vision model.', provider=_PROVIDER_NAME, operation='generate_vision_stream') + raise_if_missing_images( + images, + llm_raise, + provider=_PROVIDER_NAME, + operation='generate_vision_stream', + ) - payload: dict[str, Any] = { - 'model': model, - 'messages': _build_vision_messages(prompt, images, system_prompt), - } - payload.update(_build_options(options)) + payload = _build_openai_vision_payload( + model, + prompt, + images, + options=options, + system_prompt=system_prompt, + include_stream_field=False, + ) for chunk_data in _stream_chat_response(payload, 'generate_vision_stream'): yield chunk_data except Exception as e: llm_report('Error in OpenAI vision streaming', provider=_PROVIDER_NAME, operation='generate_vision_stream', cause=e) - yield {'chunk': '', 'done': True, 'full_response': '', 'error': llm_stringify(e)} + yield stream_error_payload(llm_stringify(e), include_full_response=True) diff --git a/utils/llm/providers/settings.py b/utils/llm/providers/settings.py index a712679..1dd5d64 100644 --- a/utils/llm/providers/settings.py +++ b/utils/llm/providers/settings.py @@ -17,16 +17,6 @@ def _is_provider_enabled(feature_key: str) -> bool: return True -def is_ollama_enabled() -> bool: - """Compatibility helper for removed Ollama SDK provider.""" - return False - - -def is_lmstudio_enabled() -> bool: - """Compatibility helper for removed LM Studio SDK provider.""" - return False - - def is_lmstudio_rest_enabled() -> bool: """Check if LM Studio REST provider is enabled in settings.""" return _is_provider_enabled('enable_lmstudio_rest') diff --git a/utils/llm/registry.py b/utils/llm/registry.py new file mode 100644 index 0000000..f955f30 --- /dev/null +++ b/utils/llm/registry.py @@ -0,0 +1,63 @@ +"""Provider registry/state helpers for LLM service initialization flow.""" + +from dataclasses import dataclass +from typing import Callable, Dict, Optional + + +@dataclass +class ProviderState: + """Runtime state for one provider.""" + + initialized: bool = False + available: bool = False + last_checked: Optional[float] = None + + +@dataclass(frozen=True) +class ProviderDescriptor: + """Static provider configuration used by the service facade.""" + + key: str + setting_key: str + display_name: str + initializer: Callable[[bool], bool] + + +class ProviderRegistry: + """Simple provider registry with per-provider runtime state.""" + + def __init__(self, retry_seconds: float = 60.0): + self.retry_seconds = retry_seconds + self._descriptors: Dict[str, ProviderDescriptor] = {} + self._state: Dict[str, ProviderState] = {} + + def register(self, descriptor: ProviderDescriptor) -> None: + self._descriptors[descriptor.key] = descriptor + self._state.setdefault(descriptor.key, ProviderState()) + + def descriptor(self, key: str) -> ProviderDescriptor: + return self._descriptors[key] + + def state(self, key: str) -> ProviderState: + return self._state[key] + + def set_state( + self, + key: str, + *, + initialized: Optional[bool] = None, + available: Optional[bool] = None, + last_checked: Optional[float] = None, + ) -> ProviderState: + state = self.state(key) + if initialized is not None: + state.initialized = initialized + if available is not None: + state.available = available + if last_checked is not None: + state.last_checked = last_checked + return state + + def reset(self) -> None: + for key in self._state: + self._state[key] = ProviderState() \ No newline at end of file diff --git a/utils/llm/routes_helpers.py b/utils/llm/routes_helpers.py index 3c19e9a..aee0c7f 100644 --- a/utils/llm/routes_helpers.py +++ b/utils/llm/routes_helpers.py @@ -10,6 +10,13 @@ from typing import Any, Optional from ..logger import get_logger from ..settings import get_setting from . import llm_raise +from .provider_keys import ( + LMSTUDIO_REST_KEY, + NATIVE_KEY, + OLLAMA_REST_KEY, + OPENAI_KEY, + normalize_provider_key, +) logger = get_logger('llm.routes_helpers') @@ -237,12 +244,7 @@ def _resolve_profile_entries( def normalize_provider(provider: str) -> str: """Normalize provider aliases to canonical backend keys.""" - normalized = (provider or '').strip().lower() - alias_map = { - 'lmstudio': 'lmstudio_rest', - 'ollama': 'ollama_rest', - } - return alias_map.get(normalized, normalized) + return normalize_provider_key(provider) def get_compatible_models(provider: str, model_list: list[str]) -> list[str]: @@ -677,7 +679,7 @@ def validate_provider(provider: str) -> tuple[bool, Optional[str]]: Tuple of (is_valid, error_message) """ provider = normalize_provider(provider) - if provider not in ['lmstudio_rest', 'ollama_rest', 'openai', 'native']: + if provider not in [LMSTUDIO_REST_KEY, OLLAMA_REST_KEY, OPENAI_KEY, NATIVE_KEY]: return False, f"Invalid provider: {provider}. Must be 'lmstudio', 'ollama', 'openai', or 'native'" return True, None diff --git a/utils/llm/service.py b/utils/llm/service.py index 8193893..13b5550 100644 --- a/utils/llm/service.py +++ b/utils/llm/service.py @@ -2,9 +2,7 @@ import time from ..logger import get_logger from .providers.settings import ( - is_lmstudio_enabled, is_lmstudio_rest_enabled, - is_ollama_enabled, is_ollama_rest_enabled, is_openai_enabled, ) @@ -16,6 +14,8 @@ from .init import ( from .providers.lmstudio import client as lmstudio_rest_provider from .providers.ollama import client as ollama_rest_provider from .providers.openai import client as openai_provider +from .provider_keys import LMSTUDIO_REST_KEY, OLLAMA_REST_KEY, OPENAI_KEY +from . import registry as llm_registry logger = get_logger('llm') @@ -34,196 +34,373 @@ _lmstudio_rest_last_checked = None _ollama_rest_last_checked = None _openai_last_checked = None +_PROVIDER_DESCRIPTORS = ( + llm_registry.ProviderDescriptor( + key=LMSTUDIO_REST_KEY, + setting_key='enable_lmstudio_rest', + display_name='LM Studio REST', + initializer=_init_lmstudio_rest, + ), + llm_registry.ProviderDescriptor( + key=OLLAMA_REST_KEY, + setting_key='enable_ollama_rest', + display_name='Ollama REST', + initializer=_init_ollama_rest, + ), + llm_registry.ProviderDescriptor( + key=OPENAI_KEY, + setting_key='enable_openai', + display_name='OpenAI provider', + initializer=_init_openai_provider, + ), +) -def _is_lmstudio_service_enabled() -> bool: - """Treat legacy LM Studio SDK toggle as an alias for REST service usage.""" - return is_lmstudio_rest_enabled() or is_lmstudio_enabled() +_PROVIDER_GLOBAL_BINDINGS = { + LMSTUDIO_REST_KEY: ('_lmstudio_rest_initialized', 'LMSTUDIO_REST_AVAILABLE', '_lmstudio_rest_last_checked'), + OLLAMA_REST_KEY: ('_ollama_rest_initialized', 'OLLAMA_REST_AVAILABLE', '_ollama_rest_last_checked'), + OPENAI_KEY: ('_openai_initialized', 'OPENAI_AVAILABLE', '_openai_last_checked'), +} + +_PROVIDER_CLIENT_ATTRS = { + LMSTUDIO_REST_KEY: 'lmstudio_rest_provider', + OLLAMA_REST_KEY: 'ollama_rest_provider', + OPENAI_KEY: 'openai_provider', +} + +_PROVIDER_ENABLED_CHECK_ATTRS = { + LMSTUDIO_REST_KEY: 'is_lmstudio_rest_enabled', + OLLAMA_REST_KEY: 'is_ollama_rest_enabled', + OPENAI_KEY: 'is_openai_enabled', +} + +_provider_registry = llm_registry.ProviderRegistry(retry_seconds=_LLM_INITIALIZATION_RETRY_SECONDS) +for _provider_descriptor in _PROVIDER_DESCRIPTORS: + _provider_registry.register(_provider_descriptor) -def _is_ollama_service_enabled() -> bool: - """Treat legacy Ollama SDK toggle as an alias for REST service usage.""" - return is_ollama_rest_enabled() or is_ollama_enabled() +def _sync_registry_from_legacy_globals() -> None: + global_values = globals() + for provider_key, (initialized_name, available_name, last_checked_name) in _PROVIDER_GLOBAL_BINDINGS.items(): + _provider_registry.set_state( + provider_key, + initialized=global_values[initialized_name], + available=global_values[available_name], + last_checked=global_values[last_checked_name], + ) -def _normalize_ollama_keep_alive(value, default: str = '5m') -> str: - """Normalize keep_alive values for Ollama REST calls.""" - if isinstance(value, str) and value.strip(): - return value.strip() - if isinstance(value, (int, float)) and value > 0: - return f'{int(value)}s' - return default +def _sync_legacy_globals_from_registry() -> None: + global_values = globals() + for provider_key, (initialized_name, available_name, last_checked_name) in _PROVIDER_GLOBAL_BINDINGS.items(): + provider_state = _provider_registry.state(provider_key) + global_values[initialized_name] = provider_state.initialized + global_values[available_name] = provider_state.available + global_values[last_checked_name] = provider_state.last_checked + + +def _init_provider(provider_key: str) -> bool: + from ..settings import get_setting + + descriptor = _provider_registry.descriptor(provider_key) + enabled = bool(get_setting(descriptor.setting_key, False)) + initialized = descriptor.initializer(enabled) + _provider_registry.set_state( + provider_key, + initialized=initialized, + available=initialized, + last_checked=time.monotonic(), + ) + _sync_legacy_globals_from_registry() + return initialized + + +def _ensure_provider_initialized(provider_key: str, force: bool = False, init_func=None) -> bool: + from ..settings import get_setting + + _sync_registry_from_legacy_globals() + descriptor = _provider_registry.descriptor(provider_key) + if not get_setting(descriptor.setting_key, False): + return False + + state = _provider_registry.state(provider_key) + if state.initialized: + return True + + if not force and state.last_checked is not None: + if time.monotonic() - state.last_checked < _provider_registry.retry_seconds: + return state.available + + logger.info(f'{descriptor.display_name} is enabled but not initialized, initializing now...') + if init_func is not None: + return init_func() + return _init_provider(provider_key) + + +def _provider_client(provider_key: str): + return globals()[_PROVIDER_CLIENT_ATTRS[provider_key]] + + +def _provider_enabled(provider_key: str) -> bool: + return globals()[_PROVIDER_ENABLED_CHECK_ATTRS[provider_key]]() + + +def _get_models_by_kind(provider_key: str, kind: str) -> list[str]: + method_map = { + 'text': 'get_models', + 'vision': 'get_vision_models', + 'tool': 'get_tool_models', + 'reasoning': 'get_reasoning_models', + } + method_name = method_map[kind] + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + return getattr(provider, method_name)(enabled) + + +def _get_model_capabilities_dict(provider_key: str) -> dict[str, dict[str, object]]: + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + capability_map = provider.get_model_capabilities_map(enabled) + return {model_name: capabilities.to_dict() for model_name, capabilities in capability_map.items()} + + +def _generate_non_streaming( + provider_key: str, + model: str, + prompt: str, + *, + options=None, + system_prompt: str = '', + keep_alive: str | None = None, +) -> str: + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + kwargs = { + 'options': options, + 'system_prompt': system_prompt, + } + if keep_alive is not None: + kwargs['keep_alive'] = keep_alive + return provider.generate(enabled, model, prompt, **kwargs) + + +def _generate_vision_non_streaming( + provider_key: str, + model: str, + prompt: str, + *, + images=None, + options=None, + system_prompt: str = '', + keep_alive: str | None = None, +) -> str: + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + kwargs = { + 'images': images, + 'options': options, + 'system_prompt': system_prompt, + } + if keep_alive is not None: + kwargs['keep_alive'] = keep_alive + return provider.generate_vision(enabled, model, prompt, **kwargs) + + +def _generate_streaming( + provider_key: str, + model: str, + prompt: str, + *, + options=None, + system_prompt: str = '', + keep_alive: str | None = None, +): + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + kwargs = { + 'options': options, + 'system_prompt': system_prompt, + } + if keep_alive is not None: + kwargs['keep_alive'] = keep_alive + return provider.generate_stream(enabled, model, prompt, **kwargs) + + +def _generate_vision_streaming( + provider_key: str, + model: str, + prompt: str, + *, + images=None, + options=None, + system_prompt: str = '', + keep_alive: str | None = None, +): + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + kwargs = { + 'images': images, + 'options': options, + 'system_prompt': system_prompt, + } + if keep_alive is not None: + kwargs['keep_alive'] = keep_alive + return provider.generate_vision_stream(enabled, model, prompt, **kwargs) + + +def _generate_lmstudio_non_streaming(model: str, prompt: str, *, options=None, system_prompt: str = '') -> str: + provider = _provider_client('lmstudio_rest') + enabled = _provider_enabled('lmstudio_rest') + return provider.generate_with_stream( + enabled, + model, + prompt, + options=options, + system_prompt=system_prompt, + ) + + +def _generate_lmstudio_vision_non_streaming(model: str, prompt: str, *, images=None, options=None, system_prompt: str = '') -> str: + provider = _provider_client('lmstudio_rest') + enabled = _provider_enabled('lmstudio_rest') + return provider.generate_vision_with_stream( + enabled, + model, + prompt, + images, + options=options, + system_prompt=system_prompt, + ) + + +def _generate_lmstudio_streaming(model: str, prompt: str, *, options=None, system_prompt: str = ''): + return _generate_streaming( + LMSTUDIO_REST_KEY, + model, + prompt, + options=options, + system_prompt=system_prompt, + ) + + +def _generate_lmstudio_vision_streaming(model: str, prompt: str, *, images=None, options=None, system_prompt: str = ''): + return _generate_vision_streaming( + LMSTUDIO_REST_KEY, + model, + prompt, + images=images, + options=options, + system_prompt=system_prompt, + ) + + +def _load_provider_model( + provider_key: str, + model: str, + *, + keep_alive: int = 0, +) -> bool: + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + return provider.load_model(enabled, model, keep_alive) + + +def _unload_provider_model( + provider_key: str, + model: str, +) -> bool: + provider = _provider_client(provider_key) + enabled = _provider_enabled(provider_key) + return provider.unload_model(enabled, model) # ============================================================================ # MODEL DISCOVERY # ============================================================================ - -def get_lmstudio_models() -> list[str]: - """Legacy alias: retrieve text models from LM Studio REST.""" - return lmstudio_rest_provider.get_models(_is_lmstudio_service_enabled()) - - -def get_lmstudio_vision_models() -> list[str]: - """Legacy alias: retrieve vision models from LM Studio REST.""" - return lmstudio_rest_provider.get_vision_models(_is_lmstudio_service_enabled()) - - -def get_ollama_models() -> list[str]: - """Legacy alias: retrieve text models from Ollama REST.""" - return ollama_rest_provider.get_models(_is_ollama_service_enabled()) - - -def get_ollama_vision_models() -> list[str]: - """Legacy alias: retrieve vision models from Ollama REST.""" - return ollama_rest_provider.get_vision_models(_is_ollama_service_enabled()) - - def get_lmstudio_rest_models() -> list[str]: """Retrieve text models from LM Studio REST.""" - return lmstudio_rest_provider.get_models(is_lmstudio_rest_enabled()) + return _get_models_by_kind(LMSTUDIO_REST_KEY, 'text') def get_lmstudio_rest_vision_models() -> list[str]: """Retrieve vision models from LM Studio REST.""" - return lmstudio_rest_provider.get_vision_models(is_lmstudio_rest_enabled()) + return _get_models_by_kind(LMSTUDIO_REST_KEY, 'vision') def get_ollama_rest_models() -> list[str]: """Retrieve text models from Ollama REST.""" - return ollama_rest_provider.get_models(is_ollama_rest_enabled()) + return _get_models_by_kind(OLLAMA_REST_KEY, 'text') def get_ollama_rest_vision_models() -> list[str]: """Retrieve vision models from Ollama REST.""" - return ollama_rest_provider.get_vision_models(is_ollama_rest_enabled()) + return _get_models_by_kind(OLLAMA_REST_KEY, 'vision') def get_openai_models() -> list[str]: """Retrieve text models from OpenAI-compatible provider.""" - return openai_provider.get_models(is_openai_enabled()) + return _get_models_by_kind(OPENAI_KEY, 'text') def get_openai_vision_models() -> list[str]: """Retrieve vision models from OpenAI-compatible provider.""" - return openai_provider.get_vision_models(is_openai_enabled()) - - -def get_lmstudio_tool_models() -> list[str]: - """Legacy alias: retrieve tool-capable models from LM Studio REST.""" - return lmstudio_rest_provider.get_tool_models(_is_lmstudio_service_enabled()) - - -def get_lmstudio_reasoning_models() -> list[str]: - """Legacy alias: retrieve reasoning-capable models from LM Studio REST.""" - return lmstudio_rest_provider.get_reasoning_models(_is_lmstudio_service_enabled()) - - -def get_ollama_tool_models() -> list[str]: - """Legacy alias: retrieve tool-capable models from Ollama REST.""" - return ollama_rest_provider.get_tool_models(_is_ollama_service_enabled()) - - -def get_ollama_reasoning_models() -> list[str]: - """Legacy alias: retrieve reasoning-capable models from Ollama REST.""" - return ollama_rest_provider.get_reasoning_models(_is_ollama_service_enabled()) + return _get_models_by_kind(OPENAI_KEY, 'vision') def get_lmstudio_rest_tool_models() -> list[str]: """Retrieve tool-capable models from LM Studio REST.""" - return lmstudio_rest_provider.get_tool_models(is_lmstudio_rest_enabled()) + return _get_models_by_kind('lmstudio_rest', 'tool') def get_lmstudio_rest_reasoning_models() -> list[str]: """Retrieve reasoning-capable models from LM Studio REST.""" - return lmstudio_rest_provider.get_reasoning_models(is_lmstudio_rest_enabled()) + return _get_models_by_kind('lmstudio_rest', 'reasoning') def get_ollama_rest_tool_models() -> list[str]: """Retrieve tool-capable models from Ollama REST.""" - return ollama_rest_provider.get_tool_models(is_ollama_rest_enabled()) + return _get_models_by_kind('ollama_rest', 'tool') def get_ollama_rest_reasoning_models() -> list[str]: """Retrieve reasoning-capable models from Ollama REST.""" - return ollama_rest_provider.get_reasoning_models(is_ollama_rest_enabled()) + return _get_models_by_kind('ollama_rest', 'reasoning') def get_openai_tool_models() -> list[str]: """Retrieve tool-capable models from OpenAI-compatible provider.""" - return openai_provider.get_tool_models(is_openai_enabled()) + return _get_models_by_kind('openai', 'tool') def get_openai_reasoning_models() -> list[str]: """Retrieve reasoning-capable models from OpenAI-compatible provider.""" - return openai_provider.get_reasoning_models(is_openai_enabled()) - - -def get_lmstudio_model_capabilities_map() -> dict[str, dict[str, object]]: - """Legacy alias: retrieve model capabilities map from LM Studio REST.""" - capability_map = lmstudio_rest_provider.get_model_capabilities_map(_is_lmstudio_service_enabled()) - return {model_name: capabilities.to_dict() for model_name, capabilities in capability_map.items()} - - -def get_ollama_model_capabilities_map() -> dict[str, dict[str, object]]: - """Legacy alias: retrieve model capabilities map from Ollama REST.""" - capability_map = ollama_rest_provider.get_model_capabilities_map(_is_ollama_service_enabled()) - return {model_name: capabilities.to_dict() for model_name, capabilities in capability_map.items()} + return _get_models_by_kind('openai', 'reasoning') def get_lmstudio_rest_model_capabilities_map() -> dict[str, dict[str, object]]: """Retrieve model capabilities map from LM Studio REST.""" - capability_map = lmstudio_rest_provider.get_model_capabilities_map(is_lmstudio_rest_enabled()) - return {model_name: capabilities.to_dict() for model_name, capabilities in capability_map.items()} + return _get_model_capabilities_dict('lmstudio_rest') def get_ollama_rest_model_capabilities_map() -> dict[str, dict[str, object]]: """Retrieve model capabilities map from Ollama REST.""" - capability_map = ollama_rest_provider.get_model_capabilities_map(is_ollama_rest_enabled()) - return {model_name: capabilities.to_dict() for model_name, capabilities in capability_map.items()} + return _get_model_capabilities_dict('ollama_rest') def get_openai_model_capabilities_map() -> dict[str, dict[str, object]]: """Retrieve model capabilities map from OpenAI-compatible provider.""" - capability_map = openai_provider.get_model_capabilities_map(is_openai_enabled()) - return {model_name: capabilities.to_dict() for model_name, capabilities in capability_map.items()} + return _get_model_capabilities_dict('openai') # ============================================================================ # GENERATION (NON-STREAMING) # ============================================================================ -def lmstudio_generate(model: str, prompt: str, keep_alive: int = 0, options=None, system_prompt: str = '') -> str: - """Legacy alias: generate text via LM Studio REST.""" - return lmstudio_rest_generate(model, prompt, keep_alive, options, system_prompt) - - -def lmstudio_generate_vision(model: str, prompt: str, keep_alive: int = 0, images=None, options=None, system_prompt: str = '') -> str: - """Legacy alias: generate vision output via LM Studio REST.""" - return lmstudio_rest_generate_vision(model, prompt, keep_alive, images, options, system_prompt) - - -def lmstudio_generate_vision_refine( - model: str, - prompt: str, - images=None, - options=None, - refine_model: str = '', - refine_prompt: str = '', - refine_options=None, -) -> tuple[str, str]: - """Legacy alias: REST-only path does not expose refine helper.""" - first_pass = lmstudio_rest_generate_vision(model, prompt, images=images, options=options) - if not refine_model or not refine_prompt: - return first_pass, '' - refined = lmstudio_rest_generate(refine_model, refine_prompt, options=refine_options) - return first_pass, refined - - +# Canonical provider operations (LM Studio REST) def lmstudio_rest_generate(model: str, prompt: str, keep_alive: int = 0, options=None, system_prompt: str = '') -> str: """Generate text via LM Studio REST.""" ensure_lmstudio_rest_initialized() - return lmstudio_rest_provider.generate_with_stream( - is_lmstudio_rest_enabled(), + return _generate_lmstudio_non_streaming( model, prompt, options=options, @@ -234,61 +411,21 @@ def lmstudio_rest_generate(model: str, prompt: str, keep_alive: int = 0, options def lmstudio_rest_generate_vision(model: str, prompt: str, keep_alive: int = 0, images=None, options=None, system_prompt: str = '') -> str: """Generate vision output via LM Studio REST.""" ensure_lmstudio_rest_initialized() - return lmstudio_rest_provider.generate_vision_with_stream( - is_lmstudio_rest_enabled(), - model, - prompt, - images, - options=options, - system_prompt=system_prompt, - ) - - -def ollama_generate(model: str, prompt: str, keep_alive: float = 0.0, options=None, system_prompt: str = '') -> str: - """Legacy alias: generate text via Ollama REST.""" - return ollama_rest_generate( - model, - prompt, - options=options, - system_prompt=system_prompt, - keep_alive=_normalize_ollama_keep_alive(keep_alive), - ) - - -def ollama_generate_vision(model: str, prompt: str, keep_alive: float = 0.0, images=None, options=None, system_prompt: str = '') -> str: - """Legacy alias: generate vision output via Ollama REST.""" - return ollama_rest_generate_vision( + return _generate_lmstudio_vision_non_streaming( model, prompt, images=images, options=options, system_prompt=system_prompt, - keep_alive=_normalize_ollama_keep_alive(keep_alive), ) -def ollama_generate_vision_refine( - model: str, - prompt: str, - images=None, - options=None, - refine_model: str = '', - refine_prompt: str = '', - refine_options=None, -) -> tuple[str, str]: - """Legacy alias: REST-only path does not expose refine helper.""" - first_pass = ollama_rest_generate_vision(model, prompt, images=images, options=options) - if not refine_model or not refine_prompt: - return first_pass, '' - refined = ollama_rest_generate(refine_model, refine_prompt, options=refine_options) - return first_pass, refined - - +# Canonical provider operations (Ollama REST) def ollama_rest_generate(model: str, prompt: str, options=None, system_prompt: str = '', keep_alive: str = '5m') -> str: """Generate text via Ollama REST.""" ensure_ollama_rest_initialized() - return ollama_rest_provider.generate( - is_ollama_rest_enabled(), + return _generate_non_streaming( + 'ollama_rest', model, prompt, options=options, @@ -300,8 +437,8 @@ def ollama_rest_generate(model: str, prompt: str, options=None, system_prompt: s def ollama_rest_generate_vision(model: str, prompt: str, images=None, options=None, system_prompt: str = '', keep_alive: str = '5m') -> str: """Generate vision output via Ollama REST.""" ensure_ollama_rest_initialized() - return ollama_rest_provider.generate_vision( - is_ollama_rest_enabled(), + return _generate_vision_non_streaming( + 'ollama_rest', model, prompt, images=images, @@ -311,11 +448,12 @@ def ollama_rest_generate_vision(model: str, prompt: str, images=None, options=No ) +# Canonical provider operations (OpenAI) def openai_generate(model: str, prompt: str, options=None, system_prompt: str = '') -> str: """Generate text via OpenAI-compatible provider.""" ensure_openai_initialized() - return openai_provider.generate( - is_openai_enabled(), + return _generate_non_streaming( + 'openai', model, prompt, options=options, @@ -326,8 +464,8 @@ def openai_generate(model: str, prompt: str, options=None, system_prompt: str = def openai_generate_vision(model: str, prompt: str, images=None, options=None, system_prompt: str = '') -> str: """Generate vision output via OpenAI-compatible provider.""" ensure_openai_initialized() - return openai_provider.generate_vision( - is_openai_enabled(), + return _generate_vision_non_streaming( + 'openai', model, prompt, images=images, @@ -340,21 +478,11 @@ def openai_generate_vision(model: str, prompt: str, images=None, options=None, s # GENERATION (STREAMING) # ============================================================================ -def lmstudio_generate_stream(model: str, prompt: str, keep_alive: int = 0, options=None, system_prompt: str = ''): - """Legacy alias: stream text via LM Studio REST.""" - return lmstudio_rest_generate_stream(model, prompt, keep_alive, options, system_prompt) - - -def lmstudio_generate_vision_stream(model: str, prompt: str, keep_alive: int = 0, images=None, options=None, system_prompt: str = ''): - """Legacy alias: stream vision output via LM Studio REST.""" - return lmstudio_rest_generate_vision_stream(model, prompt, keep_alive, images, options, system_prompt) - - +# Canonical provider operations (LM Studio REST) def lmstudio_rest_generate_stream(model: str, prompt: str, keep_alive: int = 0, options=None, system_prompt: str = ''): """Stream text via LM Studio REST.""" ensure_lmstudio_rest_initialized() - return lmstudio_rest_provider.generate_stream( - is_lmstudio_rest_enabled(), + return _generate_lmstudio_streaming( model, prompt, options=options, @@ -365,44 +493,21 @@ def lmstudio_rest_generate_stream(model: str, prompt: str, keep_alive: int = 0, def lmstudio_rest_generate_vision_stream(model: str, prompt: str, keep_alive: int = 0, images=None, options=None, system_prompt: str = ''): """Stream vision output via LM Studio REST.""" ensure_lmstudio_rest_initialized() - return lmstudio_rest_provider.generate_vision_stream( - is_lmstudio_rest_enabled(), - model, - prompt, - images, - options=options, - system_prompt=system_prompt, - ) - - -def ollama_generate_stream(model: str, prompt: str, keep_alive: float = 0.0, options=None, system_prompt: str = ''): - """Legacy alias: stream text via Ollama REST.""" - return ollama_rest_generate_stream( - model, - prompt, - options=options, - system_prompt=system_prompt, - keep_alive=_normalize_ollama_keep_alive(keep_alive), - ) - - -def ollama_generate_vision_stream(model: str, prompt: str, keep_alive: float = 0.0, images=None, options=None, system_prompt: str = ''): - """Legacy alias: stream vision output via Ollama REST.""" - return ollama_rest_generate_vision_stream( + return _generate_lmstudio_vision_streaming( model, prompt, images=images, options=options, system_prompt=system_prompt, - keep_alive=_normalize_ollama_keep_alive(keep_alive), ) +# Canonical provider operations (Ollama REST) def ollama_rest_generate_stream(model: str, prompt: str, options=None, system_prompt: str = '', keep_alive: str = '5m'): """Stream text via Ollama REST.""" ensure_ollama_rest_initialized() - return ollama_rest_provider.generate_stream( - is_ollama_rest_enabled(), + return _generate_streaming( + 'ollama_rest', model, prompt, options=options, @@ -414,8 +519,8 @@ def ollama_rest_generate_stream(model: str, prompt: str, options=None, system_pr def ollama_rest_generate_vision_stream(model: str, prompt: str, images=None, options=None, system_prompt: str = '', keep_alive: str = '5m'): """Stream vision output via Ollama REST.""" ensure_ollama_rest_initialized() - return ollama_rest_provider.generate_vision_stream( - is_ollama_rest_enabled(), + return _generate_vision_streaming( + 'ollama_rest', model, prompt, images=images, @@ -425,11 +530,12 @@ def ollama_rest_generate_vision_stream(model: str, prompt: str, images=None, opt ) +# Canonical provider operations (OpenAI) def openai_generate_stream(model: str, prompt: str, options=None, system_prompt: str = ''): """Stream text via OpenAI-compatible provider.""" ensure_openai_initialized() - return openai_provider.generate_stream( - is_openai_enabled(), + return _generate_streaming( + 'openai', model, prompt, options=options, @@ -440,8 +546,8 @@ def openai_generate_stream(model: str, prompt: str, options=None, system_prompt: def openai_generate_vision_stream(model: str, prompt: str, images=None, options=None, system_prompt: str = ''): """Stream vision output via OpenAI-compatible provider.""" ensure_openai_initialized() - return openai_provider.generate_vision_stream( - is_openai_enabled(), + return _generate_vision_streaming( + 'openai', model, prompt, images=images, @@ -453,182 +559,87 @@ def openai_generate_vision_stream(model: str, prompt: str, images=None, options= # ============================================================================ # MODEL LOAD / UNLOAD HELPERS # ============================================================================ - -def ollama_preload_model(model: str, keep_alive: float = 60.0) -> bool: - """Legacy alias: preloading is implicit in Ollama REST generate calls.""" - ensure_ollama_rest_initialized() - return ollama_rest_provider.load_model( - _is_ollama_service_enabled(), - model, - int(keep_alive), - ) - - -def ollama_generate_preloaded(model: str, prompt: str, keep_alive: float = 0.0, options=None, system_prompt: str = '') -> str: - """Legacy alias: generate using Ollama REST with keep_alive hint.""" - return ollama_generate(model, prompt, keep_alive=keep_alive, options=options, system_prompt=system_prompt) - - -def lmstudio_load_model(model: str, keep_alive: int = 0): - """Legacy alias: load model in LM Studio REST by name.""" - ensure_lmstudio_rest_initialized() - return lmstudio_rest_provider.load_model(is_lmstudio_rest_enabled(), model, keep_alive) - - -def lmstudio_generate_with_model(lms_model, prompt: str, options=None) -> str: - """Legacy alias: model-handle path not available in REST; use direct generate.""" - return lmstudio_rest_generate(str(lms_model), prompt, options=options) - - -def lmstudio_generate_vision_with_model(lms_model, prompt: str, images=None, options=None) -> str: - """Legacy alias: model-handle path not available in REST; use direct vision generate.""" - return lmstudio_rest_generate_vision(str(lms_model), prompt, images=images, options=options) - - -def lmstudio_unload_model(lms_model) -> None: - """Legacy alias: unload LM Studio REST model by name.""" - lmstudio_rest_provider.unload_model(_is_lmstudio_service_enabled(), str(lms_model)) - - +# Canonical provider operations (LM Studio REST) def lmstudio_rest_load_model(model: str, keep_alive: int = 0) -> bool: """Load LM Studio REST model by name.""" ensure_lmstudio_rest_initialized() - return lmstudio_rest_provider.load_model(is_lmstudio_rest_enabled(), model, keep_alive) + return _load_provider_model('lmstudio_rest', model, keep_alive=keep_alive) def lmstudio_rest_unload_model(model: str) -> bool: """Unload LM Studio REST model by name.""" - return lmstudio_rest_provider.unload_model(is_lmstudio_rest_enabled(), model) + return _unload_provider_model('lmstudio_rest', model) # ============================================================================ # INITIALIZATION # ============================================================================ - -def init_ollama() -> bool: - """Legacy alias: initialize Ollama REST state.""" - return init_ollama_rest() - - -def init_lmstudio() -> bool: - """Legacy alias: initialize LM Studio REST state.""" - return init_lmstudio_rest() - - +# Canonical provider operations (initialization) def init_lmstudio_rest() -> bool: """Initialize LM Studio REST provider state.""" - global _lmstudio_rest_initialized, LMSTUDIO_REST_AVAILABLE, _lmstudio_rest_last_checked - from ..settings import get_setting - - enabled = bool(get_setting('enable_lmstudio_rest', False)) - _lmstudio_rest_initialized = _init_lmstudio_rest(enabled) - LMSTUDIO_REST_AVAILABLE = _lmstudio_rest_initialized - _lmstudio_rest_last_checked = time.monotonic() - return _lmstudio_rest_initialized + return _init_provider('lmstudio_rest') def init_ollama_rest() -> bool: """Initialize Ollama REST provider state.""" - global _ollama_rest_initialized, OLLAMA_REST_AVAILABLE, _ollama_rest_last_checked - from ..settings import get_setting - - enabled = bool(get_setting('enable_ollama_rest', False)) - _ollama_rest_initialized = _init_ollama_rest(enabled) - OLLAMA_REST_AVAILABLE = _ollama_rest_initialized - _ollama_rest_last_checked = time.monotonic() - return _ollama_rest_initialized + return _init_provider('ollama_rest') def init_openai() -> bool: """Initialize OpenAI provider state.""" - global _openai_initialized, OPENAI_AVAILABLE, _openai_last_checked - from ..settings import get_setting + return _init_provider('openai') - enabled = bool(get_setting('enable_openai', False)) - _openai_initialized = _init_openai_provider(enabled) - OPENAI_AVAILABLE = _openai_initialized - _openai_last_checked = time.monotonic() - return _openai_initialized + +def _init_registered_providers() -> dict[str, bool]: + """Initialize providers in descriptor registration order.""" + results: dict[str, bool] = {} + for descriptor in _PROVIDER_DESCRIPTORS: + results[descriptor.key] = _init_provider(descriptor.key) + return results + + +def _ensure_registered_providers(force: bool = False) -> dict[str, bool]: + """Ensure providers in descriptor registration order.""" + ensure_map = { + 'lmstudio_rest': ensure_lmstudio_rest_initialized, + 'ollama_rest': ensure_ollama_rest_initialized, + 'openai': ensure_openai_initialized, + } + results: dict[str, bool] = {} + for descriptor in _PROVIDER_DESCRIPTORS: + ensure_func = ensure_map.get(descriptor.key) + if ensure_func is None: + continue + results[descriptor.key] = ensure_func(force=force) + return results def init_llm() -> None: """Initialize all configured REST/OpenAI providers.""" - init_lmstudio_rest() - init_ollama_rest() - init_openai() + _init_registered_providers() logger.info('LLM providers initialized.') -def ensure_ollama_initialized() -> bool: - """Legacy alias: ensure Ollama REST initialization.""" - return ensure_ollama_rest_initialized() - - -def ensure_lmstudio_initialized() -> bool: - """Legacy alias: ensure LM Studio REST initialization.""" - return ensure_lmstudio_rest_initialized() - - +# Canonical provider operations (ensure initialization) def ensure_lmstudio_rest_initialized(force: bool = False) -> bool: """Ensure LM Studio REST is initialized if enabled.""" - global _lmstudio_rest_initialized, LMSTUDIO_REST_AVAILABLE, _lmstudio_rest_last_checked - from ..settings import get_setting - - if not get_setting('enable_lmstudio_rest', False): - return False - - if _lmstudio_rest_initialized: - return True - - if not force and _lmstudio_rest_last_checked is not None and time.monotonic() - _lmstudio_rest_last_checked < _LLM_INITIALIZATION_RETRY_SECONDS: - return LMSTUDIO_REST_AVAILABLE - - logger.info('LM Studio REST is enabled but not initialized, initializing now...') - return init_lmstudio_rest() + return _ensure_provider_initialized('lmstudio_rest', force=force, init_func=init_lmstudio_rest) def ensure_ollama_rest_initialized(force: bool = False) -> bool: """Ensure Ollama REST is initialized if enabled.""" - global _ollama_rest_initialized, OLLAMA_REST_AVAILABLE, _ollama_rest_last_checked - from ..settings import get_setting - - if not get_setting('enable_ollama_rest', False): - return False - - if _ollama_rest_initialized: - return True - - if not force and _ollama_rest_last_checked is not None and time.monotonic() - _ollama_rest_last_checked < _LLM_INITIALIZATION_RETRY_SECONDS: - return OLLAMA_REST_AVAILABLE - - logger.info('Ollama REST is enabled but not initialized, initializing now...') - return init_ollama_rest() + return _ensure_provider_initialized('ollama_rest', force=force, init_func=init_ollama_rest) def ensure_openai_initialized(force: bool = False) -> bool: """Ensure OpenAI provider is initialized if enabled.""" - global _openai_initialized, OPENAI_AVAILABLE, _openai_last_checked - from ..settings import get_setting - - if not get_setting('enable_openai', False): - return False - - if _openai_initialized: - return True - - if not force and _openai_last_checked is not None and time.monotonic() - _openai_last_checked < _LLM_INITIALIZATION_RETRY_SECONDS: - return OPENAI_AVAILABLE - - logger.info('OpenAI provider is enabled but not initialized, initializing now...') - return init_openai() + return _ensure_provider_initialized('openai', force=force, init_func=init_openai) def ensure_llm_initialized(force: bool = False) -> bool: """Ensure all enabled LLM services are initialized.""" - lmstudio_rest_ok = ensure_lmstudio_rest_initialized(force=force) - ollama_rest_ok = ensure_ollama_rest_initialized(force=force) - openai_ok = ensure_openai_initialized(force=force) - return lmstudio_rest_ok or ollama_rest_ok or openai_ok + ensure_results = _ensure_registered_providers(force=force) + return any(ensure_results.values()) def reset_llm_initialization_state() -> None: @@ -646,3 +657,5 @@ def reset_llm_initialization_state() -> None: _lmstudio_rest_last_checked = None _ollama_rest_last_checked = None _openai_last_checked = None + + _provider_registry.reset() diff --git a/utils/performance_fix.py b/utils/performance_fix.py index 95c466e..10bb8b2 100644 --- a/utils/performance_fix.py +++ b/utils/performance_fix.py @@ -2,7 +2,7 @@ """ Performance fix for INPUT_TYPES methods that are causing 32+ second delays during node registration. -The issue: Multiple LLM nodes call get_ollama_models() and get_lmstudio_models() in their INPUT_TYPES methods. +The issue: Multiple LLM nodes call provider discovery methods during INPUT_TYPES evaluation. These functions make network calls that can timeout, causing massive startup delays. The fix: Create lightweight cached versions that don't block during node registration. @@ -195,20 +195,12 @@ def populate_llm_cache_async(): logger.info("Starting background LLM cache population...") # These calls will populate the cache asynchronously - logger.info("Fetching LLM models to populate cache (Ollama)...") - llm.get_ollama_models() - logger.info("Fetching LLM models to populate cache (LM Studio)...") - llm.get_lmstudio_models() logger.info("Fetching LLM models to populate cache (LM Studio REST)...") llm.get_lmstudio_rest_models() logger.info("Fetching LLM models to populate cache (Ollama REST)...") llm.get_ollama_rest_models() logger.info("Fetching LLM models to populate cache (OpenAI)...") llm.get_openai_models() - logger.info("Fetching LLM vision models to populate cache (Ollama)...") - llm.get_ollama_vision_models() - logger.info("Fetching LLM vision models to populate cache (LM Studio)...") - llm.get_lmstudio_vision_models() logger.info("Fetching LLM vision models to populate cache (LM Studio REST)...") llm.get_lmstudio_rest_vision_models() logger.info("Fetching LLM vision models to populate cache (Ollama REST)...")