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:
+342
-250
@@ -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>
|
||||
@@ -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>
|
||||
@@ -0,0 +1,81 @@
|
||||
const htmlTemplateCache = new Map();
|
||||
|
||||
function escapeHtml(value) {
|
||||
return value.replace(/[&<>"]|\'/g, (char) => {
|
||||
switch (char) {
|
||||
case '&': return '&';
|
||||
case '<': return '<';
|
||||
case '>': return '>';
|
||||
case '"': return '"';
|
||||
case "'": return ''';
|
||||
default: return char;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function getTemplateValue(key, data) {
|
||||
return key.split('.').reduce((current, part) => {
|
||||
return current && current[part] !== undefined ? current[part] : undefined;
|
||||
}, data);
|
||||
}
|
||||
|
||||
export async function loadHtmlTemplate(templatePath) {
|
||||
if (typeof fetch === 'undefined') {
|
||||
return '';
|
||||
}
|
||||
|
||||
if (htmlTemplateCache.has(templatePath)) {
|
||||
return htmlTemplateCache.get(templatePath);
|
||||
}
|
||||
|
||||
const response = await fetch(templatePath, { cache: 'force-cache' });
|
||||
if (!response.ok) {
|
||||
throw new Error(`Unable to load HTML template: ${templatePath} (${response.status})`);
|
||||
}
|
||||
|
||||
const text = await response.text();
|
||||
htmlTemplateCache.set(templatePath, text);
|
||||
return text;
|
||||
}
|
||||
|
||||
export function renderHtmlTemplate(templateString, data = {}) {
|
||||
return templateString
|
||||
.replace(/\{\{\{\s*([\w.]+)\s*\}\}\}/g, (_, key) => {
|
||||
const value = getTemplateValue(key, data);
|
||||
return value == null ? '' : String(value);
|
||||
})
|
||||
.replace(/\{\{\s*([\w.]+)\s*\}\}/g, (_, key) => {
|
||||
const value = getTemplateValue(key, data);
|
||||
return value == null ? '' : escapeHtml(String(value));
|
||||
});
|
||||
}
|
||||
|
||||
export function createElementFromTemplate(templateString, data = {}) {
|
||||
if (typeof document === 'undefined') {
|
||||
return null;
|
||||
}
|
||||
|
||||
const rendered = renderHtmlTemplate(templateString, data).trim();
|
||||
const template = document.createElement('template');
|
||||
template.innerHTML = rendered;
|
||||
|
||||
if (template.content.childElementCount === 1) {
|
||||
return template.content.firstElementChild;
|
||||
}
|
||||
|
||||
return template.content;
|
||||
}
|
||||
|
||||
export async function loadAndCreateHtmlTemplate(templatePath, data = {}) {
|
||||
const templateString = await loadHtmlTemplate(templatePath);
|
||||
return createElementFromTemplate(templateString, data);
|
||||
}
|
||||
|
||||
export function clearHtmlTemplateCache(templatePath) {
|
||||
if (templatePath) {
|
||||
htmlTemplateCache.delete(templatePath);
|
||||
return;
|
||||
}
|
||||
|
||||
htmlTemplateCache.clear();
|
||||
}
|
||||
+32
-28
@@ -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
@@ -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
@@ -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
@@ -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',
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
"""Compatibility module retained for import stability during refactor migration."""
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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()
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)...")
|
||||
|
||||
Reference in New Issue
Block a user