Combination of some work on the llm code, and work on pulling out inline html. Meant to commit separately, but rather forgot I hadn't committed the llm changes yet.

This commit is contained in:
Shanoah Alkire
2026-06-05 19:36:56 -07:00
parent 3d15042f7a
commit 6fbc702585
26 changed files with 1730 additions and 1850 deletions
+342 -250
View File
@@ -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 = `
<div class="civitai-search-status-panel">
<h3 class="civitai-search-status-title">🔍 Civitai Model Search</h3>
<p>Search for models on Civitai to discover and download new content.</p>
<p>Use the search form above to find models by name, type, or creator.</p>
</div>
`;
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 = `
<label class="civitai-search-label" for="civitai-search-input">Search Query:</label>
<input type="text" id="civitai-search-input" class="civitai-search-input" placeholder="Enter model name, creator, or keywords...">
`;
// Type filter
const typeGroup = document.createElement('div');
typeGroup.className = 'civitai-search-group';
typeGroup.innerHTML = `
<label class="civitai-search-label" for="civitai-type-filter">Type:</label>
<select id="civitai-type-filter" class="civitai-search-select">
<option value="">All Types</option>
${SEARCH_CONFIG.SUPPORTED_TYPES.map(type => `<option value="${type}">${type}</option>`).join('')}
</select>
`;
// Sort filter
const sortGroup = document.createElement('div');
sortGroup.className = 'civitai-search-group';
sortGroup.innerHTML = `
<label class="civitai-search-label" for="civitai-sort-filter">Sort:</label>
<select id="civitai-sort-filter" class="civitai-search-select">
${SEARCH_CONFIG.SORT_OPTIONS.map(sort => `<option value="${sort}">${sort}</option>`).join('')}
</select>
`;
// 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 = `
<input type="checkbox" id="civitai-nsfw-toggle">
Include NSFW Results
`;
// Results per page
const limitGroup = document.createElement('div');
limitGroup.className = 'civitai-limit-group';
limitGroup.innerHTML = `
<label class="civitai-search-label" for="civitai-limit-select">Results:</label>
<select id="civitai-limit-select" class="civitai-search-select">
<option value="10">10</option>
<option value="20" selected>20</option>
<option value="50">50</option>
<option value="100">100</option>
</select>
`;
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 = `
<div class="civitai-search-status-panel">
<div class="civitai-search-status-icon">🔍</div>
<h3 class="civitai-search-status-title">Civitai Model Search</h3>
<p>Search for models on Civitai to discover and download new content.</p>
<p>Use the search form above to find models by name, type, or creator.</p>
</div>
`;
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 = `
<div class="civitai-search-status-panel civitai-search-status-error">
<h3 class="civitai-search-status-title">❌ Search Failed</h3>
<p>Failed to search Civitai: ${escapeHtml(error.message)}</p>
<p class="civitai-search-status-note">Please check your internet connection and try again.</p>
</div>
`;
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 = `
<div class="civitai-search-status-panel">
<h3 class="civitai-search-status-title">No Results Found</h3>
<p>No models found matching your search criteria.</p>
<p>Try adjusting your search terms or filters.</p>
</div>
`;
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 = `
<h3 class="civitai-results-header-title">Found ${results.length} models${metadata && metadata.totalItems ? ` (${metadata.totalItems} total)` : ''}</h3>
<p class="civitai-results-subtitle">Click on any model to view details and download options</p>
`;
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 = `
<div>Page ${metadata.currentPage}${metadata.totalPages ? ` of ${metadata.totalPages}` : ''}</div>
<div class="civitai-results-pagination-detail">${metadata.pageSize || results.length} per page</div>
`;
} else if (metadata.nextCursor) {
paginationInfo.innerHTML = `
<div>Cursor-based pagination</div>
<div class="civitai-results-pagination-detail">More results available</div>
`;
}
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 = `
<div class="civitai-model-card-inner">
<div class="civitai-model-card-image">
${imageUrl ?
`<img src="${escapeHtml(imageUrl)}" class="civitai-model-card-image-img" alt="Model preview" loading="lazy">` :
`<div class="civitai-model-card-no-image">No Image</div>`
}
</div>
<div class="civitai-model-card-content">
<h4 class="civitai-model-card-title">${escapeHtml(modelName)}</h4>
<p class="civitai-model-card-meta">by ${escapeHtml(creator)} • ${escapeHtml(type)}</p>
<p class="civitai-model-card-description">${escapeHtml(description)}</p>
const card = await loadAndCreateHtmlTemplate(
'extensions/comfyui_sageutils/sidebar/partials/civitaiModelCard.html',
{
modelName,
creator,
type,
description
}
);
<div class="civitai-model-card-stats-row">
<div class="civitai-model-card-stats">
${stats.downloadCount ? `<span>📥 ${stats.downloadCount.toLocaleString()}</span>` : ''}
${stats.favoriteCount ? `<span>❤️ ${stats.favoriteCount.toLocaleString()}</span>` : ''}
${stats.rating ? `<span>⭐ ${stats.rating.toFixed(1)}</span>` : ''}
</div>
${trainedWords.length > 0 ? `<div class="civitai-model-card-triggers">🏷️ ${trainedWords.length} triggers</div>` : ''}
</div>
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);
}
}
<div class="civitai-model-card-actions">
<div class="civitai-card-btn-group">
<button class="civitai-card-btn civitai-details-btn">📖 Details</button>
${latestVersion ? `<button class="civitai-card-btn civitai-download-btn" data-version-id="${latestVersion.id}" data-version-name="${escapeHtml(latestVersion.name || 'Latest Version')}" data-model-id="${model.id}">📥 Download Latest</button>` : ''}
</div>
${latestVersion && latestVersion.files && latestVersion.files[0] ? `<div class="civitai-card-stats-small">${formatFileSize((latestVersion.files[0].sizeKB || 0) * 1024)}</div>` : ''}
</div>
</div>
</div>
`;
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 = `
<h2 class="civitai-dialog-title">${escapeHtml(modelName)}</h2>
<div class="civitai-dialog-meta">
<div>
<span class="civitai-dialog-meta-label">by</span> <strong class="civitai-dialog-meta-creator">${escapeHtml(creator)}</strong>
<span class="civitai-dialog-meta-label">Type:</span> <strong class="civitai-dialog-meta-type">${escapeHtml(type)}</strong>
</div>
<a href="${getModelUrl(model.id)}" target="_blank" class="civitai-dialog-link">View on Civitai →</a>
</div>
<div class="civitai-dialog-stats">
${stats.downloadCount ? `<span>📥 ${stats.downloadCount.toLocaleString()} downloads</span>` : ''}
${stats.favoriteCount ? `<span>❤️ ${stats.favoriteCount.toLocaleString()} favorites</span>` : ''}
${stats.rating ? `<span>⭐ ${stats.rating.toFixed(1)} rating</span>` : ''}
${stats.commentCount ? `<span>💬 ${stats.commentCount.toLocaleString()} comments</span>` : ''}
</div>
`;
const descriptionSection = document.createElement('div');
descriptionSection.className = 'civitai-dialog-description';
descriptionSection.innerHTML = `
<h3 class="civitai-dialog-description-title">Description</h3>
<div class="civitai-dialog-description-text">${escapeHtml(stripHtml(description))}</div>
`;
const versionsSection = document.createElement('div');
versionsSection.className = 'civitai-versions-section';
if (versions.length > 0) {
versionsSection.innerHTML = `
<h3 class="civitai-versions-title">Available Versions (${versions.length})</h3>
`;
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 = `
<div class="civitai-version-header">
<div class="civitai-version-info">
<h4 class="civitai-version-title">${escapeHtml(versionName)}${isLatest ? ' <span class="civitai-version-badge">LATEST</span>' : ''}</h4>
<p class="civitai-version-subtitle">Released: ${createdAt}</p>
${version.description ? `<p class="civitai-version-description">${escapeHtml(version.description)}</p>` : ''}
${trainedWords.length > 0 ? `
<div class="civitai-trigger-words-group">
<strong class="civitai-trigger-words-title">Trigger Words:</strong>
<div class="civitai-trigger-words">${trainedWords.map(word => `<span class="civitai-trigger-word">${escapeHtml(word)}</span>`).join('')}</div>
</div>
` : ''}
${primaryFile ? `
<div class="civitai-version-file-info">
<span>File Size:</span> ${formatFileSize(primaryFile.sizeKB * 1024)}
${primaryFile.metadata ? `
${primaryFile.metadata.format ? `<span>Format:</span> ${primaryFile.metadata.format}` : ''}
${primaryFile.metadata.fp ? `<span>Precision:</span> ${primaryFile.metadata.fp}` : ''}
` : ''}
</div>
` : ''}
</div>
<div class="civitai-version-download-container">
<button class="civitai-version-download-button" data-version-id="${version.id}" data-version-name="${escapeHtml(versionName)}" data-model-id="${modelId}">📥 Download</button>
</div>
</div>
${appropriateImages.length > 0 ? `
<div class="civitai-version-images">
${appropriateImages.map(img => `<img src="${escapeHtml(img.url)}" class="civitai-version-image" data-image-url="${escapeHtml(img.url)}" alt="Example image" loading="lazy">`).join('')}
</div>
` : ''}
`;
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;
}
@@ -0,0 +1,19 @@
<div class="civitai-dialog-content">
<div class="civitai-dialog-header">
<h2 class="civitai-dialog-title">{{modelName}}</h2>
<div class="civitai-dialog-meta">
<div>
<span class="civitai-dialog-meta-label">by</span>
<strong class="civitai-dialog-meta-creator">{{creator}}</strong>
<span class="civitai-dialog-meta-label">Type:</span>
<strong class="civitai-dialog-meta-type">{{type}}</strong>
</div>
<a href="{{modelUrl}}" target="_blank" class="civitai-dialog-link">View on Civitai →</a>
</div>
<div class="civitai-dialog-stats"></div>
</div>
<div class="civitai-dialog-description">
<h3 class="civitai-dialog-description-title">Description</h3>
<div class="civitai-dialog-description-text">{{descriptionText}}</div>
</div>
</div>
+22
View File
@@ -0,0 +1,22 @@
<div class="civitai-model-card">
<div class="civitai-model-card-inner">
<div class="civitai-model-card-image"></div>
<div class="civitai-model-card-content">
<h4 class="civitai-model-card-title">{{modelName}}</h4>
<p class="civitai-model-card-meta">by {{creator}} • {{type}}</p>
<p class="civitai-model-card-description">{{description}}</p>
<div class="civitai-model-card-stats-row">
<div class="civitai-model-card-stats"></div>
<div class="civitai-model-card-triggers"></div>
</div>
<div class="civitai-model-card-actions">
<div class="civitai-card-btn-group">
<button class="civitai-card-btn civitai-details-btn">📖 Details</button>
</div>
<div class="civitai-card-stats-small"></div>
</div>
</div>
</div>
</div>
@@ -0,0 +1,4 @@
<div class="civitai-results-pagination-info">
<div>{{pageLine}}</div>
<div class="civitai-results-pagination-detail">{{detailText}}</div>
</div>
@@ -0,0 +1,6 @@
<div class="civitai-results-header">
<div class="civitai-results-header-text">
<h3 class="civitai-results-header-title">Found {{resultCount}} models{{{totalItemsText}}}</h3>
<p class="civitai-results-subtitle">Click on any model to view details and download options</p>
</div>
</div>
@@ -0,0 +1,40 @@
<div class="civitai-search-form">
<div class="civitai-search-group">
<label class="civitai-search-label" for="civitai-search-input">Search Query:</label>
<input type="text" id="civitai-search-input" class="civitai-search-input" placeholder="Enter model name, creator, or keywords...">
</div>
<div class="civitai-search-group">
<label class="civitai-search-label" for="civitai-type-filter">Type:</label>
<select id="civitai-type-filter" class="civitai-search-select">
<option value="">All Types</option>
{{{typeOptions}}}
</select>
</div>
<div class="civitai-search-group">
<label class="civitai-search-label" for="civitai-sort-filter">Sort:</label>
<select id="civitai-sort-filter" class="civitai-search-select">
{{{sortOptions}}}
</select>
</div>
<div class="civitai-search-button-row">
<div id="civitai-search-button-spot"></div>
<label class="civitai-nsfw-toggle">
<input type="checkbox" id="civitai-nsfw-toggle">
Include NSFW Results
</label>
<div class="civitai-limit-group">
<label class="civitai-search-label" for="civitai-limit-select">Results:</label>
<select id="civitai-limit-select" class="civitai-search-select">
<option value="10">10</option>
<option value="20" selected>20</option>
<option value="50">50</option>
<option value="100">100</option>
</select>
</div>
</div>
</div>
@@ -0,0 +1,7 @@
<div class="civitai-search-status-panel{{{extraClass}}}">
<div class="civitai-search-status-icon">{{icon}}</div>
<h3 class="civitai-search-status-title">{{title}}</h3>
<p>{{messageLine1}}</p>
<p>{{messageLine2}}</p>
<p class="civitai-search-status-note">{{note}}</p>
</div>
@@ -0,0 +1,13 @@
<div class="civitai-version-card{{latestClass}}">
<div class="civitai-version-header">
<div class="civitai-version-info">
<h4 class="civitai-version-title">{{versionName}}</h4>
<p class="civitai-version-subtitle">Released: {{createdAt}}</p>
<div class="civitai-version-description"></div>
<div class="civitai-trigger-words-group"></div>
<div class="civitai-version-file-info"></div>
</div>
<div class="civitai-version-download-container"></div>
</div>
<div class="civitai-version-images"></div>
</div>
@@ -0,0 +1,4 @@
<div class="civitai-versions-section">
<h3 class="civitai-versions-title">Available Versions ({{versionCount}})</h3>
<div class="civitai-version-list"></div>
</div>
+81
View File
@@ -0,0 +1,81 @@
const htmlTemplateCache = new Map();
function escapeHtml(value) {
return value.replace(/[&<>"]|\'/g, (char) => {
switch (char) {
case '&': return '&amp;';
case '<': return '&lt;';
case '>': return '&gt;';
case '"': return '&quot;';
case "'": return '&#39;';
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();
}
+32 -28
View File
@@ -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():
+44 -24
View File
@@ -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 = "", ""
+63 -865
View File
@@ -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`.
+127 -114
View File
@@ -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',
)
+4
View File
@@ -0,0 +1,4 @@
"""Compatibility module retained for import stability during refactor migration."""
-35
View File
@@ -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)
+30
View File
@@ -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)
+84
View File
@@ -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
+118 -56
View File
@@ -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)
+148 -66
View File
@@ -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."""
+114 -44
View File
@@ -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)
-10
View File
@@ -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')
+63
View File
@@ -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()
+9 -7
View File
@@ -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
+355 -342
View File
@@ -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()
+1 -9
View File
@@ -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)...")