Removed original ollama and lm studio backend in favor of rest. Got OpenAI working.
This commit is contained in:
+118
-17
@@ -1,7 +1,7 @@
|
||||
/**
|
||||
* Settings Dialog
|
||||
* Provides a UI for configuring SageUtils settings
|
||||
* Currently focuses on LLM provider configuration (Ollama, LM Studio, and Native)
|
||||
* Currently focuses on LLM provider configuration (REST providers, OpenAI, and Native)
|
||||
*/
|
||||
|
||||
import { api } from '../../../../scripts/api.js';
|
||||
@@ -21,18 +21,18 @@ import {
|
||||
* Mapping between backend setting keys and ComfyUI setting IDs
|
||||
*/
|
||||
const SETTING_KEY_TO_ID_MAP = {
|
||||
'enable_lmstudio': 'SageUtils.LLM Providers.enable_lmstudio',
|
||||
'enable_lmstudio_rest': 'SageUtils.LLM Providers.enable_lmstudio_rest',
|
||||
'enable_ollama_rest': 'SageUtils.LLM Providers.enable_ollama_rest',
|
||||
'enable_openai': 'SageUtils.LLM Providers.enable_openai',
|
||||
'openai_api_key': 'SageUtils.OpenAI.openai_api_key',
|
||||
'openai_use_custom_url': 'SageUtils.OpenAI.openai_use_custom_url',
|
||||
'openai_base_url': 'SageUtils.OpenAI.openai_base_url',
|
||||
'enable_ollama': 'SageUtils.LLM Providers.enable_ollama',
|
||||
'ollama_custom_url': 'SageUtils.Local Custom Ollama URL.ollama_custom_url',
|
||||
'ollama_use_custom_url': 'SageUtils.Local Custom Ollama URL.ollama_use_custom_url',
|
||||
'ollama_api_key': 'SageUtils.Ollama.ollama_api_key',
|
||||
'lmstudio_custom_url': 'SageUtils.Local Custom LM Studio URL.lmstudio_custom_url',
|
||||
'lmstudio_use_custom_url': 'SageUtils.Local Custom LM Studio URL.lmstudio_use_custom_url'
|
||||
'lmstudio_use_custom_url': 'SageUtils.Local Custom LM Studio URL.lmstudio_use_custom_url',
|
||||
'lmstudio_api_token': 'SageUtils.LM Studio.lmstudio_api_token'
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -186,7 +186,7 @@ function createLLMSection(settings) {
|
||||
section.appendChild(topSeparator);
|
||||
|
||||
// Ollama settings
|
||||
const ollamaGroup = createProviderGroup('Ollama', settings, 'ollama');
|
||||
const ollamaGroup = createProviderGroup('Ollama', settings, 'ollama_rest', 'ollama');
|
||||
section.appendChild(ollamaGroup);
|
||||
|
||||
// Separator
|
||||
@@ -198,8 +198,21 @@ function createLLMSection(settings) {
|
||||
`;
|
||||
section.appendChild(separator);
|
||||
|
||||
// OpenAI settings
|
||||
const openaiGroup = createOpenAIProviderGroup(settings);
|
||||
section.appendChild(openaiGroup);
|
||||
|
||||
// Separator
|
||||
const openaiSeparator = document.createElement('div');
|
||||
openaiSeparator.style.cssText = `
|
||||
height: 1px;
|
||||
background: #444;
|
||||
margin: 15px 0;
|
||||
`;
|
||||
section.appendChild(openaiSeparator);
|
||||
|
||||
// LM Studio settings
|
||||
const lmstudioGroup = createProviderGroup('LM Studio', settings, 'lmstudio');
|
||||
const lmstudioGroup = createProviderGroup('LM Studio', settings, 'lmstudio_rest', 'lmstudio');
|
||||
section.appendChild(lmstudioGroup);
|
||||
|
||||
return section;
|
||||
@@ -237,14 +250,12 @@ function createDefaultProviderSelector(settings) {
|
||||
|
||||
// Get current setting
|
||||
const defaultProviderSetting = settings['default_llm_provider'];
|
||||
const currentValue = defaultProviderSetting ? defaultProviderSetting.current_value : 'ollama';
|
||||
const currentValue = defaultProviderSetting ? defaultProviderSetting.current_value : 'lmstudio_rest';
|
||||
|
||||
// Create radio button group
|
||||
const radioItems = [
|
||||
{ value: 'ollama', label: 'Ollama' },
|
||||
{ value: 'lmstudio', label: 'LM Studio' },
|
||||
{ value: 'lmstudio_rest', label: 'LM Studio (REST)' },
|
||||
{ value: 'ollama_rest', label: 'Ollama (REST)' },
|
||||
{ value: 'lmstudio_rest', label: 'LM Studio' },
|
||||
{ value: 'ollama_rest', label: 'Ollama' },
|
||||
{ value: 'openai', label: 'OpenAI' },
|
||||
{ value: 'native', label: 'Native (CLIP)' }
|
||||
];
|
||||
@@ -334,13 +345,14 @@ function createTabVisibilitySection(settings) {
|
||||
}
|
||||
|
||||
/**
|
||||
* Create settings group for a provider (Ollama or LM Studio)
|
||||
* Create settings group for a provider
|
||||
* @param {string} providerName - Display name of the provider
|
||||
* @param {Object} settings - Current settings
|
||||
* @param {string} providerKey - Key prefix (ollama or lmstudio)
|
||||
* @param {string} providerKey - Enable flag key suffix (e.g., ollama_rest)
|
||||
* @param {string} urlKeyPrefix - URL key prefix (ollama or lmstudio)
|
||||
* @returns {HTMLElement} Provider group element
|
||||
*/
|
||||
function createProviderGroup(providerName, settings, providerKey) {
|
||||
function createProviderGroup(providerName, settings, providerKey, urlKeyPrefix) {
|
||||
const group = document.createElement('div');
|
||||
group.style.cssText = 'margin-bottom: 15px;';
|
||||
|
||||
@@ -358,7 +370,7 @@ function createProviderGroup(providerName, settings, providerKey) {
|
||||
group.appendChild(enableContainer);
|
||||
|
||||
// Custom URL checkbox
|
||||
const useCustomUrlKey = `${providerKey}_use_custom_url`;
|
||||
const useCustomUrlKey = `${urlKeyPrefix}_use_custom_url`;
|
||||
const useCustomUrlSetting = settings[useCustomUrlKey];
|
||||
const { container: customUrlCheckboxContainer, checkbox: customUrlCheckbox } = createCheckbox(
|
||||
`Use custom URL for ${providerName}`,
|
||||
@@ -372,13 +384,13 @@ function createProviderGroup(providerName, settings, providerKey) {
|
||||
group.appendChild(customUrlCheckboxContainer);
|
||||
|
||||
// Custom URL input
|
||||
const customUrlKey = `${providerKey}_custom_url`;
|
||||
const customUrlKey = `${urlKeyPrefix}_custom_url`;
|
||||
const customUrlSetting = settings[customUrlKey];
|
||||
const { container: customUrlInputContainer, input: customUrlInput } = createTextInput(
|
||||
`${providerName} URL`,
|
||||
customUrlSetting ? customUrlSetting.current_value : '',
|
||||
customUrlKey,
|
||||
`e.g., http://localhost:${providerKey === 'ollama' ? '11434' : '1234'}`
|
||||
`e.g., http://localhost:${urlKeyPrefix === 'ollama' ? '11434' : '1234'}`
|
||||
);
|
||||
customUrlInputContainer.style.marginLeft = '40px';
|
||||
customUrlInputContainer.style.marginTop = '10px';
|
||||
@@ -393,6 +405,95 @@ function createProviderGroup(providerName, settings, providerKey) {
|
||||
|
||||
group.appendChild(customUrlInputContainer);
|
||||
|
||||
if (urlKeyPrefix === 'lmstudio' || urlKeyPrefix === 'ollama') {
|
||||
const tokenKey = urlKeyPrefix === 'ollama' ? 'ollama_api_key' : 'lmstudio_api_token';
|
||||
const tokenLabel = urlKeyPrefix === 'ollama' ? 'Ollama API Key' : 'LM Studio API Token';
|
||||
const tokenPlaceholder = urlKeyPrefix === 'ollama' ? 'Optional bearer key' : 'Optional bearer token';
|
||||
const tokenSetting = settings[tokenKey];
|
||||
const { container: tokenInputContainer, input: tokenInput } = createTextInput(
|
||||
tokenLabel,
|
||||
tokenSetting ? tokenSetting.current_value : '',
|
||||
tokenKey,
|
||||
tokenPlaceholder
|
||||
);
|
||||
tokenInput.type = 'password';
|
||||
tokenInput.autocomplete = 'off';
|
||||
tokenInputContainer.style.marginLeft = '20px';
|
||||
tokenInputContainer.style.marginTop = '10px';
|
||||
group.appendChild(tokenInputContainer);
|
||||
}
|
||||
|
||||
return group;
|
||||
}
|
||||
|
||||
/**
|
||||
* Create OpenAI settings group with API key and custom URL fields
|
||||
* @param {Object} settings - Current settings
|
||||
* @returns {HTMLElement} OpenAI group element
|
||||
*/
|
||||
function createOpenAIProviderGroup(settings) {
|
||||
const group = document.createElement('div');
|
||||
group.style.cssText = 'margin-bottom: 15px;';
|
||||
|
||||
const enableKey = 'enable_openai';
|
||||
const enableSetting = settings[enableKey];
|
||||
const { container: enableContainer, checkbox: enableCheckbox } = createCheckbox(
|
||||
'Enable OpenAI',
|
||||
{
|
||||
checked: enableSetting ? enableSetting.current_value : false,
|
||||
id: `setting-${enableKey}`
|
||||
}
|
||||
);
|
||||
enableCheckbox.dataset.settingKey = enableKey;
|
||||
group.appendChild(enableContainer);
|
||||
|
||||
const useCustomUrlKey = 'openai_use_custom_url';
|
||||
const useCustomUrlSetting = settings[useCustomUrlKey];
|
||||
const { container: customUrlCheckboxContainer, checkbox: customUrlCheckbox } = createCheckbox(
|
||||
'Use custom URL for OpenAI',
|
||||
{
|
||||
checked: useCustomUrlSetting ? useCustomUrlSetting.current_value : false,
|
||||
id: `setting-${useCustomUrlKey}`
|
||||
}
|
||||
);
|
||||
customUrlCheckbox.dataset.settingKey = useCustomUrlKey;
|
||||
customUrlCheckboxContainer.style.marginLeft = '20px';
|
||||
customUrlCheckboxContainer.style.marginTop = '10px';
|
||||
group.appendChild(customUrlCheckboxContainer);
|
||||
|
||||
const baseUrlKey = 'openai_base_url';
|
||||
const baseUrlSetting = settings[baseUrlKey];
|
||||
const { container: baseUrlInputContainer } = createTextInput(
|
||||
'OpenAI Base URL',
|
||||
baseUrlSetting ? baseUrlSetting.current_value : '',
|
||||
baseUrlKey,
|
||||
'e.g., https://api.openai.com'
|
||||
);
|
||||
baseUrlInputContainer.style.marginLeft = '40px';
|
||||
baseUrlInputContainer.style.marginTop = '10px';
|
||||
|
||||
const updateUrlInputVisibility = () => {
|
||||
baseUrlInputContainer.style.display = customUrlCheckbox.checked ? 'block' : 'none';
|
||||
};
|
||||
updateUrlInputVisibility();
|
||||
customUrlCheckbox.addEventListener('change', updateUrlInputVisibility);
|
||||
|
||||
group.appendChild(baseUrlInputContainer);
|
||||
|
||||
const apiKeyKey = 'openai_api_key';
|
||||
const apiKeySetting = settings[apiKeyKey];
|
||||
const { container: apiKeyInputContainer, input: apiKeyInput } = createTextInput(
|
||||
'OpenAI API Key',
|
||||
apiKeySetting ? apiKeySetting.current_value : '',
|
||||
apiKeyKey,
|
||||
'sk-...'
|
||||
);
|
||||
apiKeyInput.type = 'password';
|
||||
apiKeyInput.autocomplete = 'off';
|
||||
apiKeyInputContainer.style.marginLeft = '20px';
|
||||
apiKeyInputContainer.style.marginTop = '10px';
|
||||
group.appendChild(apiKeyInputContainer);
|
||||
|
||||
return group;
|
||||
}
|
||||
|
||||
|
||||
@@ -33,7 +33,7 @@ const BUILTIN_PRESETS = {
|
||||
'descriptive_prompt': {
|
||||
name: 'Descriptive Prompt',
|
||||
description: 'Generate detailed image descriptions',
|
||||
provider: 'ollama',
|
||||
provider: 'lmstudio_rest',
|
||||
model: 'gemma3:12b', // Default vision model for image description
|
||||
promptTemplate: 'description/Descriptive Prompt',
|
||||
systemPrompt: 'e621_prompt_generator',
|
||||
@@ -50,7 +50,7 @@ const BUILTIN_PRESETS = {
|
||||
'e621_description': {
|
||||
name: 'E621 Image Description',
|
||||
description: 'Generate E621-style detailed image descriptions',
|
||||
provider: 'ollama',
|
||||
provider: 'lmstudio_rest',
|
||||
model: 'gemma3:12b', // Default vision model for image description
|
||||
promptTemplate: 'description/Descriptive Prompt',
|
||||
systemPrompt: 'e621_prompt_generator',
|
||||
@@ -67,7 +67,7 @@ const BUILTIN_PRESETS = {
|
||||
'casual_chat': {
|
||||
name: 'Casual Chat',
|
||||
description: 'Friendly conversational assistant',
|
||||
provider: 'ollama',
|
||||
provider: 'lmstudio_rest',
|
||||
model: null,
|
||||
promptTemplate: '',
|
||||
systemPrompt: 'default',
|
||||
|
||||
+25
-47
@@ -1,15 +1,22 @@
|
||||
/**
|
||||
* LLM Provider-Specific Logic
|
||||
* Handles differences between Ollama, LM Studio, and Native providers
|
||||
* Handles differences between REST providers, OpenAI, and Native providers
|
||||
*/
|
||||
|
||||
function normalizeProviderKey(provider) {
|
||||
if (provider === 'lmstudio') return 'lmstudio_rest';
|
||||
if (provider === 'ollama') return 'ollama_rest';
|
||||
return provider;
|
||||
}
|
||||
|
||||
/**
|
||||
* Build generation options for a specific provider
|
||||
* @param {string} provider - 'ollama', 'lmstudio', or 'native'
|
||||
* @param {string} provider - 'lmstudio_rest', 'ollama_rest', 'openai', or 'native'
|
||||
* @param {Object} settings - Settings object
|
||||
* @returns {Object} - Provider-specific options object
|
||||
*/
|
||||
export function buildProviderOptions(provider, settings) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
// Common options for both providers
|
||||
const options = {
|
||||
temperature: settings.temperature,
|
||||
@@ -19,16 +26,7 @@ export function buildProviderOptions(provider, settings) {
|
||||
};
|
||||
|
||||
// Add provider-specific options
|
||||
if (provider === 'ollama') {
|
||||
options.num_keep = settings.numKeep;
|
||||
options.num_predict = settings.numPredict;
|
||||
options.top_k = settings.topK;
|
||||
options.top_p = settings.topP;
|
||||
options.repeat_last_n = settings.repeatLastN;
|
||||
options.repeat_penalty = settings.repeatPenalty;
|
||||
options.presence_penalty = settings.presencePenalty;
|
||||
options.frequency_penalty = settings.frequencyPenalty;
|
||||
} else if (provider === 'lmstudio' || provider === 'lmstudio_rest') {
|
||||
if (provider === 'lmstudio_rest') {
|
||||
options.topKSampling = settings.lmsTopK;
|
||||
options.topPSampling = settings.lmsTopP;
|
||||
options.repeatPenalty = settings.lmsRepeatPenalty;
|
||||
@@ -59,10 +57,11 @@ export function buildProviderOptions(provider, settings) {
|
||||
|
||||
/**
|
||||
* Get provider-specific default settings
|
||||
* @param {string} provider - 'ollama', 'lmstudio', or 'native'
|
||||
* @param {string} provider - 'lmstudio_rest', 'ollama_rest', 'openai', or 'native'
|
||||
* @returns {Object} - Provider-specific defaults
|
||||
*/
|
||||
export function getProviderDefaults(provider) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
const commonDefaults = {
|
||||
temperature: 0.7,
|
||||
seed: 42,
|
||||
@@ -70,19 +69,7 @@ export function getProviderDefaults(provider) {
|
||||
keepAlive: 300
|
||||
};
|
||||
|
||||
if (provider === 'ollama') {
|
||||
return {
|
||||
...commonDefaults,
|
||||
numKeep: 0,
|
||||
numPredict: -1,
|
||||
topK: 40,
|
||||
topP: 0.9,
|
||||
repeatLastN: 64,
|
||||
repeatPenalty: 1.1,
|
||||
presencePenalty: 0.0,
|
||||
frequencyPenalty: 0.0
|
||||
};
|
||||
} else if (provider === 'lmstudio' || provider === 'lmstudio_rest') {
|
||||
if (provider === 'lmstudio_rest') {
|
||||
return {
|
||||
...commonDefaults,
|
||||
lmsTopK: 40,
|
||||
@@ -124,12 +111,11 @@ export function getProviderDefaults(provider) {
|
||||
* @returns {string} - Human-readable provider name
|
||||
*/
|
||||
export function getProviderDisplayName(provider) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
const names = {
|
||||
'ollama': 'Ollama',
|
||||
'lmstudio': 'LM Studio',
|
||||
'lmstudio_rest': 'LM Studio (REST)',
|
||||
'lmstudio_rest': 'LM Studio',
|
||||
'native': 'Native (CLIP)',
|
||||
'ollama_rest': 'Ollama (REST)',
|
||||
'ollama_rest': 'Ollama',
|
||||
'openai': 'OpenAI'
|
||||
};
|
||||
return names[provider] || provider;
|
||||
@@ -142,6 +128,7 @@ export function getProviderDisplayName(provider) {
|
||||
* @returns {string} - Status class name
|
||||
*/
|
||||
export function getProviderStatusClass(status, provider) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
if (!status || !status[provider]) {
|
||||
return 'status-offline';
|
||||
}
|
||||
@@ -164,6 +151,7 @@ export function getProviderStatusClass(status, provider) {
|
||||
* @returns {string} - Status text
|
||||
*/
|
||||
export function getProviderStatusText(status, provider) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
const displayName = getProviderDisplayName(provider);
|
||||
|
||||
if (!status || !status[provider]) {
|
||||
@@ -189,11 +177,12 @@ export function getProviderStatusText(status, provider) {
|
||||
* @returns {Object} - Validation result { valid: boolean, error: string }
|
||||
*/
|
||||
export function validateProviderConfig(provider, model, models) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
if (!provider) {
|
||||
return { valid: false, error: 'No provider selected' };
|
||||
}
|
||||
|
||||
if (!['ollama', 'lmstudio', 'lmstudio_rest', 'ollama_rest', 'openai', 'native'].includes(provider)) {
|
||||
if (!['lmstudio_rest', 'ollama_rest', 'openai', 'native'].includes(provider)) {
|
||||
return { valid: false, error: `Invalid provider: ${provider}` };
|
||||
}
|
||||
|
||||
@@ -219,6 +208,7 @@ export function validateProviderConfig(provider, model, models) {
|
||||
* @returns {Object} - Parameter descriptions keyed by parameter name
|
||||
*/
|
||||
export function getProviderParameterDescriptions(provider) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
const common = {
|
||||
temperature: 'Controls randomness in output. Lower = more focused, higher = more creative.',
|
||||
seed: 'Random seed for reproducible outputs. Use -1 for random.',
|
||||
@@ -226,19 +216,7 @@ export function getProviderParameterDescriptions(provider) {
|
||||
keepAlive: 'How long to keep model loaded in memory (seconds).'
|
||||
};
|
||||
|
||||
if (provider === 'ollama') {
|
||||
return {
|
||||
...common,
|
||||
numKeep: 'Number of tokens to keep from initial prompt.',
|
||||
numPredict: 'Maximum tokens to predict. -1 for infinite.',
|
||||
topK: 'Limits next token selection to top K tokens.',
|
||||
topP: 'Nucleus sampling: cumulative probability threshold.',
|
||||
repeatLastN: 'How far back to look for repetitions.',
|
||||
repeatPenalty: 'Penalty for repeating tokens.',
|
||||
presencePenalty: 'Penalty for tokens that have appeared.',
|
||||
frequencyPenalty: 'Penalty based on token frequency.'
|
||||
};
|
||||
} else if (provider === 'lmstudio' || provider === 'lmstudio_rest') {
|
||||
if (provider === 'lmstudio_rest') {
|
||||
return {
|
||||
...common,
|
||||
lmsTopK: 'Limits next token selection to top K tokens.',
|
||||
@@ -280,11 +258,10 @@ export function getProviderParameterDescriptions(provider) {
|
||||
* @returns {string[]} - Array of setting keys specific to this provider
|
||||
*/
|
||||
export function getProviderSettingKeys(provider) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
const common = ['temperature', 'seed', 'maxTokens', 'keepAlive', 'systemPrompt', 'promptTemplate', 'includeHistory', 'maxHistoryMessages'];
|
||||
|
||||
if (provider === 'ollama') {
|
||||
return [...common, 'numKeep', 'numPredict', 'topK', 'topP', 'repeatLastN', 'repeatPenalty', 'presencePenalty', 'frequencyPenalty'];
|
||||
} else if (provider === 'lmstudio' || provider === 'lmstudio_rest') {
|
||||
if (provider === 'lmstudio_rest') {
|
||||
return [...common, 'lmsTopK', 'lmsTopP', 'lmsRepeatPenalty', 'lmsMinP'];
|
||||
} else if (provider === 'ollama_rest') {
|
||||
return [...common, 'numKeep', 'numPredict', 'topK', 'topP', 'repeatLastN', 'repeatPenalty', 'presencePenalty', 'frequencyPenalty'];
|
||||
@@ -326,6 +303,7 @@ export function getModelCapabilityFlags(
|
||||
toolModelsByProvider,
|
||||
reasoningModelsByProvider
|
||||
) {
|
||||
provider = normalizeProviderKey(provider);
|
||||
const providerCapabilities = capabilitiesByProvider?.[provider] || {};
|
||||
const modelCapabilities = providerCapabilities?.[model] || null;
|
||||
|
||||
|
||||
+41
-33
@@ -33,19 +33,19 @@ async function loadSageSettings() {
|
||||
console.error('Error loading SageUtils settings:', error);
|
||||
// Return fallback settings structure
|
||||
return {
|
||||
enable_ollama: { current_value: true },
|
||||
enable_lmstudio: { current_value: true },
|
||||
enable_lmstudio_rest: { current_value: false },
|
||||
enable_ollama_rest: { current_value: false },
|
||||
enable_openai: { current_value: false },
|
||||
openai_api_key: { current_value: '' },
|
||||
openai_use_custom_url: { current_value: false },
|
||||
openai_base_url: { current_value: '' },
|
||||
default_llm_provider: { current_value: 'ollama' },
|
||||
default_llm_provider: { current_value: 'lmstudio_rest' },
|
||||
ollama_use_custom_url: { current_value: false },
|
||||
ollama_custom_url: { current_value: "" },
|
||||
ollama_api_key: { current_value: '' },
|
||||
lmstudio_use_custom_url: { current_value: false },
|
||||
lmstudio_custom_url: { current_value: "" }
|
||||
lmstudio_custom_url: { current_value: "" },
|
||||
lmstudio_api_token: { current_value: '' }
|
||||
};
|
||||
}
|
||||
}
|
||||
@@ -96,18 +96,18 @@ app.registerExtension({
|
||||
// Map backend keys to the new frontend setting IDs
|
||||
const keyToIdMap = {
|
||||
'default_llm_provider': 'SageUtils.LLM Providers.default_llm_provider',
|
||||
'enable_lmstudio': 'SageUtils.LLM Providers.enable_lmstudio',
|
||||
'enable_lmstudio_rest': 'SageUtils.LLM Providers.enable_lmstudio_rest',
|
||||
'enable_ollama_rest': 'SageUtils.LLM Providers.enable_ollama_rest',
|
||||
'enable_openai': 'SageUtils.LLM Providers.enable_openai',
|
||||
'openai_api_key': 'SageUtils.OpenAI.openai_api_key',
|
||||
'openai_use_custom_url': 'SageUtils.OpenAI.openai_use_custom_url',
|
||||
'openai_base_url': 'SageUtils.OpenAI.openai_base_url',
|
||||
'enable_ollama': 'SageUtils.LLM Providers.enable_ollama',
|
||||
'ollama_custom_url': 'SageUtils.Local Custom Ollama URL.ollama_custom_url',
|
||||
'ollama_use_custom_url': 'SageUtils.Local Custom Ollama URL.ollama_use_custom_url',
|
||||
'ollama_api_key': 'SageUtils.Ollama.ollama_api_key',
|
||||
'lmstudio_custom_url': 'SageUtils.Local Custom LM Studio URL.lmstudio_custom_url',
|
||||
'lmstudio_use_custom_url': 'SageUtils.Local Custom LM Studio URL.lmstudio_use_custom_url'
|
||||
'lmstudio_use_custom_url': 'SageUtils.Local Custom LM Studio URL.lmstudio_use_custom_url',
|
||||
'lmstudio_api_token': 'SageUtils.LM Studio.lmstudio_api_token'
|
||||
};
|
||||
|
||||
for (const [key, settingInfo] of Object.entries(serverSettings)) {
|
||||
@@ -130,8 +130,8 @@ app.registerExtension({
|
||||
id: "SageUtils.LLM Providers.default_llm_provider",
|
||||
name: "Default LLM Provider",
|
||||
type: "combo",
|
||||
defaultValue: "ollama",
|
||||
options: ["ollama", "lmstudio", "lmstudio_rest", "ollama_rest", "openai", "native"],
|
||||
defaultValue: "lmstudio_rest",
|
||||
options: ["lmstudio_rest", "ollama_rest", "openai", "native"],
|
||||
tooltip: "Default provider used by the LLM sidebar and provider-switching LLM v3 nodes",
|
||||
onChange: async (newVal, oldVal) => {
|
||||
console.log(`Default LLM provider changed from ${oldVal} to ${newVal}`);
|
||||
@@ -140,7 +140,7 @@ app.registerExtension({
|
||||
},
|
||||
{
|
||||
id: "SageUtils.LLM Providers.enable_lmstudio_rest",
|
||||
name: "Enable LM Studio (REST) Integration",
|
||||
name: "Enable LM Studio Integration",
|
||||
type: "boolean",
|
||||
defaultValue: false,
|
||||
tooltip: "Enable LM Studio REST v1 integration",
|
||||
@@ -151,7 +151,7 @@ app.registerExtension({
|
||||
},
|
||||
{
|
||||
id: "SageUtils.LLM Providers.enable_ollama_rest",
|
||||
name: "Enable Ollama (REST) Integration",
|
||||
name: "Enable Ollama Integration",
|
||||
type: "boolean",
|
||||
defaultValue: false,
|
||||
tooltip: "Enable Ollama native REST API integration (no SDK required)",
|
||||
@@ -204,28 +204,6 @@ app.registerExtension({
|
||||
await saveSageSetting('openai_base_url', newVal);
|
||||
}
|
||||
},
|
||||
{
|
||||
id: "SageUtils.LLM Providers.enable_lmstudio",
|
||||
name: "Enable LM Studio Integration",
|
||||
type: "boolean",
|
||||
defaultValue: true,
|
||||
tooltip: "Enable LM Studio LLM integration",
|
||||
onChange: async (newVal, oldVal) => {
|
||||
console.log(`LM Studio integration changed from ${oldVal} to ${newVal}`);
|
||||
await saveSageSetting('enable_lmstudio', newVal);
|
||||
}
|
||||
},
|
||||
{
|
||||
id: "SageUtils.LLM Providers.enable_ollama",
|
||||
name: "Enable Ollama Integration",
|
||||
type: "boolean",
|
||||
defaultValue: true,
|
||||
tooltip: "Enable Ollama LLM integration",
|
||||
onChange: async (newVal, oldVal) => {
|
||||
console.log(`Ollama integration changed from ${oldVal} to ${newVal}`);
|
||||
await saveSageSetting('enable_ollama', newVal);
|
||||
}
|
||||
},
|
||||
{
|
||||
id: "SageUtils.Local Custom Ollama URL.ollama_custom_url",
|
||||
name: "Address",
|
||||
@@ -248,6 +226,21 @@ app.registerExtension({
|
||||
await saveSageSetting('ollama_use_custom_url', newVal);
|
||||
}
|
||||
},
|
||||
{
|
||||
id: "SageUtils.Ollama.ollama_api_key",
|
||||
name: "API Key",
|
||||
type: "text",
|
||||
defaultValue: "",
|
||||
attrs: {
|
||||
type: 'password',
|
||||
autocomplete: 'off'
|
||||
},
|
||||
tooltip: "Ollama API key (or leave blank to use OLLAMA_API_KEY env var)",
|
||||
onChange: async (newVal) => {
|
||||
console.log('Ollama API key changed');
|
||||
await saveSageSetting('ollama_api_key', newVal);
|
||||
}
|
||||
},
|
||||
{
|
||||
id: "SageUtils.Local Custom LM Studio URL.lmstudio_custom_url",
|
||||
name: "Address",
|
||||
@@ -269,6 +262,21 @@ app.registerExtension({
|
||||
console.log(`LM Studio custom URL setting changed from ${oldVal} to ${newVal}`);
|
||||
await saveSageSetting('lmstudio_use_custom_url', newVal);
|
||||
}
|
||||
},
|
||||
{
|
||||
id: "SageUtils.LM Studio.lmstudio_api_token",
|
||||
name: "API Token",
|
||||
type: "text",
|
||||
defaultValue: "",
|
||||
attrs: {
|
||||
type: 'password',
|
||||
autocomplete: 'off'
|
||||
},
|
||||
tooltip: "LM Studio API token (or leave blank to use LMSTUDIO_API_TOKEN env var)",
|
||||
onChange: async (newVal) => {
|
||||
console.log('LM Studio API token changed');
|
||||
await saveSageSetting('lmstudio_api_token', newVal);
|
||||
}
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
+14
-12
@@ -49,7 +49,7 @@ function savePromptText(value) {
|
||||
function loadLastSelectedProvider() {
|
||||
try {
|
||||
const provider = localStorage.getItem(LLM_LAST_PROVIDER_KEY);
|
||||
if (provider === 'ollama' || provider === 'lmstudio' || provider === 'lmstudio_rest' ||
|
||||
if (provider === 'lmstudio_rest' ||
|
||||
provider === 'ollama_rest' || provider === 'openai' || provider === 'native') {
|
||||
return provider;
|
||||
}
|
||||
@@ -81,7 +81,7 @@ function logLlmDebug(...args) {
|
||||
|
||||
/**
|
||||
* Load default LLM provider from settings
|
||||
* @returns {Promise<string>} Default provider ('ollama', 'lmstudio', 'lmstudio_rest', or 'native')
|
||||
* @returns {Promise<string>} Default provider ('lmstudio_rest', 'ollama_rest', 'openai', or 'native')
|
||||
*/
|
||||
async function loadDefaultProvider() {
|
||||
try {
|
||||
@@ -95,15 +95,17 @@ async function loadDefaultProvider() {
|
||||
logLlmDebug('[LLM Tab] default_llm_provider setting:', setting);
|
||||
|
||||
// Try current_value first, then fall back to default
|
||||
const provider = setting.current_value || setting.default;
|
||||
let provider = setting.current_value || setting.default;
|
||||
if (provider === 'lmstudio') provider = 'lmstudio_rest';
|
||||
if (provider === 'ollama') provider = 'ollama_rest';
|
||||
logLlmDebug('[LLM Tab] Resolved provider value:', provider);
|
||||
|
||||
if (provider === 'ollama' || provider === 'lmstudio' || provider === 'lmstudio_rest' ||
|
||||
if (provider === 'lmstudio_rest' ||
|
||||
provider === 'ollama_rest' || provider === 'openai' || provider === 'native') {
|
||||
logLlmDebug(`[LLM Tab] Loading default LLM provider: ${provider}`);
|
||||
return provider;
|
||||
} else {
|
||||
console.warn(`[LLM Tab] Invalid provider value: ${provider}, using ollama`);
|
||||
console.warn(`[LLM Tab] Invalid provider value: ${provider}, using lmstudio_rest`);
|
||||
}
|
||||
} else {
|
||||
console.warn('[LLM Tab] default_llm_provider setting not found in response');
|
||||
@@ -112,9 +114,9 @@ async function loadDefaultProvider() {
|
||||
console.warn('[LLM Tab] Settings API request failed:', response.status);
|
||||
}
|
||||
} catch (error) {
|
||||
console.warn('[LLM Tab] Failed to load default LLM provider setting, using ollama:', error);
|
||||
console.warn('[LLM Tab] Failed to load default LLM provider setting, using lmstudio_rest:', error);
|
||||
}
|
||||
return 'ollama'; // Default fallback
|
||||
return 'lmstudio_rest'; // Default fallback
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -351,11 +353,11 @@ async function createLLMTabVanilla(container) {
|
||||
const state = {
|
||||
provider: initialProvider,
|
||||
model: null,
|
||||
models: { ollama: [], lmstudio: [], lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
visionModels: { ollama: [], lmstudio: [], lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
toolModels: { ollama: [], lmstudio: [], lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
reasoningModels: { ollama: [], lmstudio: [], lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
capabilities: { ollama: {}, lmstudio: {}, lmstudio_rest: {}, ollama_rest: {}, openai: {}, native: {} },
|
||||
models: { lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
visionModels: { lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
toolModels: { lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
reasoningModels: { lmstudio_rest: [], ollama_rest: [], openai: [], native: [] },
|
||||
capabilities: { lmstudio_rest: {}, ollama_rest: {}, openai: {}, native: {} },
|
||||
generating: false,
|
||||
streamController: null,
|
||||
// Vision support
|
||||
|
||||
@@ -709,10 +709,10 @@ function showProviderOptions(advancedOptions, provider) {
|
||||
const lmstudioSection = advancedOptions.querySelector('.llm-lmstudio-options');
|
||||
|
||||
if (ollamaSection) {
|
||||
ollamaSection.style.display = (provider === 'ollama' || provider === 'ollama_rest') ? 'block' : 'none';
|
||||
ollamaSection.style.display = provider === 'ollama_rest' ? 'block' : 'none';
|
||||
}
|
||||
if (lmstudioSection) {
|
||||
lmstudioSection.style.display = (provider === 'lmstudio' || provider === 'lmstudio_rest') ? 'block' : 'none';
|
||||
lmstudioSection.style.display = provider === 'lmstudio_rest' ? 'block' : 'none';
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -363,16 +363,7 @@ function buildGenerationOptions(state) {
|
||||
};
|
||||
|
||||
// Add provider-specific options
|
||||
if (state.provider === 'ollama') {
|
||||
options.num_keep = state.settings.numKeep;
|
||||
options.num_predict = state.settings.numPredict;
|
||||
options.top_k = state.settings.topK;
|
||||
options.top_p = state.settings.topP;
|
||||
options.repeat_last_n = state.settings.repeatLastN;
|
||||
options.repeat_penalty = state.settings.repeatPenalty;
|
||||
options.presence_penalty = state.settings.presencePenalty;
|
||||
options.frequency_penalty = state.settings.frequencyPenalty;
|
||||
} else if (state.provider === 'lmstudio' || state.provider === 'lmstudio_rest') {
|
||||
if (state.provider === 'lmstudio_rest') {
|
||||
options.topKSampling = state.settings.lmsTopK;
|
||||
options.topPSampling = state.settings.lmsTopP;
|
||||
options.repeatPenalty = state.settings.lmsRepeatPenalty;
|
||||
|
||||
@@ -81,10 +81,8 @@ export function createModelSelection() {
|
||||
|
||||
const providerSelect = createSelect({
|
||||
items: [
|
||||
{ value: 'ollama', text: 'Ollama' },
|
||||
{ value: 'lmstudio', text: 'LM Studio' },
|
||||
{ value: 'lmstudio_rest', text: 'LM Studio (REST)' },
|
||||
{ value: 'ollama_rest', text: 'Ollama (REST)' },
|
||||
{ value: 'lmstudio_rest', text: 'LM Studio' },
|
||||
{ value: 'ollama_rest', text: 'Ollama' },
|
||||
{ value: 'openai', text: 'OpenAI' },
|
||||
{ value: 'native', text: 'Native (CLIP)' }
|
||||
],
|
||||
|
||||
@@ -437,10 +437,8 @@ async function showPresetEditor(state, modelSelection, advancedOptions, inputSec
|
||||
(() => {
|
||||
const select = createSelect({
|
||||
items: [
|
||||
{ value: 'ollama', text: 'Ollama' },
|
||||
{ value: 'lmstudio', text: 'LM Studio' },
|
||||
{ value: 'lmstudio_rest', text: 'LM Studio (REST)' },
|
||||
{ value: 'ollama_rest', text: 'Ollama (REST)' },
|
||||
{ value: 'lmstudio_rest', text: 'LM Studio' },
|
||||
{ value: 'ollama_rest', text: 'Ollama' },
|
||||
{ value: 'openai', text: 'OpenAI' },
|
||||
{ value: 'native', text: 'Native (CLIP)' }
|
||||
],
|
||||
|
||||
@@ -880,15 +880,6 @@ export function addLLMStyles() {
|
||||
color: #f44336;
|
||||
}
|
||||
|
||||
/* Provider-specific colors */
|
||||
.llm-provider-select option[value="ollama"] {
|
||||
color: #6bcf7f;
|
||||
}
|
||||
|
||||
.llm-provider-select option[value="lmstudio"] {
|
||||
color: #4a9eff;
|
||||
}
|
||||
|
||||
/* History section */
|
||||
.llm-history-section {
|
||||
padding: 12px;
|
||||
|
||||
+34
-187
@@ -22,12 +22,6 @@ from ..utils.settings import get_setting
|
||||
import logging
|
||||
|
||||
from ..utils.performance_fix import (
|
||||
get_cached_ollama_models_for_input_types,
|
||||
get_cached_ollama_vision_models_for_input_types,
|
||||
)
|
||||
from ..utils.performance_fix import (
|
||||
get_cached_lmstudio_models_for_input_types,
|
||||
get_cached_lmstudio_vision_models_for_input_types,
|
||||
get_cached_lmstudio_rest_models_for_input_types,
|
||||
get_cached_lmstudio_rest_vision_models_for_input_types,
|
||||
get_cached_ollama_rest_models_for_input_types,
|
||||
@@ -47,27 +41,12 @@ from ..utils.llm.tensor import tensor_to_base64_safe
|
||||
from ..utils.constants import SAGE_UTILS_CAT
|
||||
|
||||
from .custom_io_v3 import *
|
||||
from .ollama_v3 import (
|
||||
Sage_OllamaAdvancedOptions,
|
||||
Sage_OllamaLLMPromptText,
|
||||
Sage_OllamaLLMPromptVision,
|
||||
Sage_OllamaLLMPromptVisionRefine,
|
||||
OLLAMA_NODES,
|
||||
)
|
||||
from .lmstudio_v3 import (
|
||||
Sage_LMStudioLLMPromptText,
|
||||
Sage_LMStudioLLMPromptVision,
|
||||
Sage_LMStudioLLMPromptVisionRefine,
|
||||
LMSTUDIO_NODES,
|
||||
)
|
||||
|
||||
logger = logging.getLogger('sageutils.nodes.llm_v3')
|
||||
|
||||
_PROVIDER_LABEL_BY_KEY = {
|
||||
"ollama": "Ollama",
|
||||
"lmstudio": "LM Studio",
|
||||
"lmstudio_rest": "LM Studio (REST)",
|
||||
"ollama_rest": "Ollama (REST)",
|
||||
"lmstudio_rest": "LM Studio",
|
||||
"ollama_rest": "Ollama",
|
||||
"openai": "OpenAI",
|
||||
"native": "Native",
|
||||
}
|
||||
@@ -215,14 +194,6 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
"""Unified text generation node that switches between Ollama, LM Studio, and Native providers."""
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
ollama_models = get_cached_ollama_models_for_input_types()
|
||||
if not ollama_models:
|
||||
ollama_models = ["(Ollama not available)"]
|
||||
|
||||
lm_models = get_cached_lmstudio_models_for_input_types()
|
||||
if not lm_models:
|
||||
lm_models = ["(LM Studio not available)"]
|
||||
|
||||
lmstudio_rest_models = get_cached_lmstudio_rest_models_for_input_types()
|
||||
if not lmstudio_rest_models:
|
||||
lmstudio_rest_models = ["(LM Studio REST not available)"]
|
||||
@@ -254,24 +225,8 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
]
|
||||
|
||||
provider_options = _ordered_provider_options([
|
||||
io.DynamicCombo.Option(
|
||||
"Ollama",
|
||||
[
|
||||
io.Combo.Input("ollama_model", display_name="model", options=sorted(ollama_models)),
|
||||
io.Float.Input("ollama_keep_alive", display_name="keep_alive", default=0.0, min=-1.0, max=60.0 * 60.0, step=1, advanced=True, tooltip="How long to keep the model loaded after generation (-1 keeps it resident)."),
|
||||
OllamaOptions.Input("ollama_options", display_name="options", optional=True, advanced=True, tooltip="Optional low-level Ollama generation parameters."),
|
||||
io.String.Input("ollama_system_prompt", display_name="system_prompt", default="", multiline=True, optional=True, advanced=True, tooltip="Optional system instruction prepended as model context."),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"LM Studio",
|
||||
[
|
||||
io.Combo.Input("lm_model", display_name="model", options=sorted(lm_models)),
|
||||
io.Int.Input("lm_load_for_seconds", display_name="load_for_seconds", default=0, min=-1, max=60 * 60, step=1, advanced=True, tooltip="How long to keep model loaded in LM Studio (seconds)."),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"LM Studio (REST)",
|
||||
[
|
||||
io.Combo.Input("lmstudio_rest_model", display_name="model", options=sorted(lmstudio_rest_models)),
|
||||
io.Int.Input("lmstudio_rest_load_for_seconds", display_name="load_for_seconds", default=0, min=-1, max=60 * 60, step=1, advanced=True, tooltip="Compatibility field for LM Studio REST model load duration."),
|
||||
@@ -279,7 +234,7 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"Ollama (REST)",
|
||||
"Ollama",
|
||||
[
|
||||
io.Combo.Input("ollama_rest_model", display_name="model", options=sorted(ollama_rest_models)),
|
||||
io.String.Input("ollama_rest_keep_alive", display_name="keep_alive", default="5m", advanced=True, tooltip="How long to keep the model loaded after generation (Ollama duration string)."),
|
||||
@@ -293,7 +248,7 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
io.Combo.Input("openai_model", display_name="model", options=sorted(openai_models)),
|
||||
io.String.Input("openai_system_prompt", display_name="system_prompt", default="", multiline=True, optional=True, advanced=True, tooltip="Optional system instruction prepended as model context."),
|
||||
io.Float.Input("openai_temperature", display_name="temperature", default=0.7, min=0.0, max=2.0, step=0.01, advanced=True),
|
||||
io.Int.Input("openai_max_tokens", display_name="max_tokens", default=1024, min=1, max=16384, step=1, advanced=True),
|
||||
io.Int.Input("openai_max_tokens", display_name="max_tokens", default=4096, min=1, max=16384, step=1, advanced=True),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option("Native", native_inputs),
|
||||
@@ -302,7 +257,7 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
return io.Schema(
|
||||
node_id="Sage_LLMPromptText",
|
||||
display_name="LLM Prompt (Text)",
|
||||
description="Unified provider-switching text generation node for Ollama, LM Studio, and Native CLIP.",
|
||||
description="Unified provider-switching text generation node for REST/OpenAI and Native CLIP.",
|
||||
category=f"{SAGE_UTILS_CAT}/LLM",
|
||||
inputs=[
|
||||
io.String.Input("prompt", display_name="prompt", default=DEFAULT_TEXT_PROMPT, multiline=True),
|
||||
@@ -324,25 +279,7 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
native_sampling = _get_native_sampling_config(provider_data)
|
||||
|
||||
try:
|
||||
if provider == "Ollama":
|
||||
return Sage_OllamaLLMPromptText.execute(
|
||||
prompt=prompt,
|
||||
model=provider_data.get("ollama_model"),
|
||||
seed=seed,
|
||||
keep_alive=provider_data.get("ollama_keep_alive", 0.0),
|
||||
options=provider_data.get("ollama_options") or {},
|
||||
system_prompt=provider_data.get("ollama_system_prompt", ""),
|
||||
)
|
||||
|
||||
if provider == "LM Studio":
|
||||
return Sage_LMStudioLLMPromptText.execute(
|
||||
prompt=prompt,
|
||||
model=provider_data.get("lm_model"),
|
||||
seed=seed,
|
||||
load_for_seconds=provider_data.get("lm_load_for_seconds", 0),
|
||||
)
|
||||
|
||||
if provider == "LM Studio (REST)":
|
||||
model = str(provider_data.get("lmstudio_rest_model") or "").strip()
|
||||
if not model or model.startswith("("):
|
||||
return io.NodeOutput("")
|
||||
@@ -357,7 +294,7 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
)
|
||||
)
|
||||
|
||||
if provider == "Ollama (REST)":
|
||||
if provider == "Ollama":
|
||||
model = str(provider_data.get("ollama_rest_model") or "").strip()
|
||||
if not model or model.startswith("("):
|
||||
return io.NodeOutput("")
|
||||
@@ -380,7 +317,7 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
options = {
|
||||
"seed": seed,
|
||||
"temperature": provider_data.get("openai_temperature", 0.7),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 1024),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 4096),
|
||||
}
|
||||
return io.NodeOutput(
|
||||
openai_generate(
|
||||
@@ -414,17 +351,9 @@ class Sage_LLMPromptText(io.ComfyNode):
|
||||
return io.NodeOutput("")
|
||||
|
||||
class Sage_LLMPromptVision(io.ComfyNode):
|
||||
"""Unified vision generation node that switches between Ollama, LM Studio, and Native providers."""
|
||||
"""Unified vision generation node that switches between REST/OpenAI and Native providers."""
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
ollama_models = get_cached_ollama_vision_models_for_input_types()
|
||||
if not ollama_models:
|
||||
ollama_models = ["(No Ollama vision models available)"]
|
||||
|
||||
lm_models = get_cached_lmstudio_vision_models_for_input_types()
|
||||
if not lm_models:
|
||||
lm_models = ["(No LM Studio vision models available)"]
|
||||
|
||||
lmstudio_rest_models = get_cached_lmstudio_rest_vision_models_for_input_types()
|
||||
if not lmstudio_rest_models:
|
||||
lmstudio_rest_models = ["(No LM Studio REST vision models available)"]
|
||||
@@ -456,24 +385,8 @@ class Sage_LLMPromptVision(io.ComfyNode):
|
||||
]
|
||||
|
||||
provider_options = _ordered_provider_options([
|
||||
io.DynamicCombo.Option(
|
||||
"Ollama",
|
||||
[
|
||||
io.Combo.Input("ollama_model", display_name="model", options=sorted(ollama_models)),
|
||||
io.Float.Input("ollama_keep_alive", display_name="keep_alive", default=0.0, min=-1.0, max=60.0 * 60.0, step=0.1, advanced=True, tooltip="How long to keep the model loaded after generation (-1 keeps it resident)."),
|
||||
OllamaOptions.Input("ollama_options", display_name="options", optional=True, advanced=True, tooltip="Optional low-level Ollama generation parameters."),
|
||||
io.String.Input("ollama_system_prompt", display_name="system_prompt", default="", multiline=True, optional=True, advanced=True, tooltip="Optional system instruction prepended as model context."),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"LM Studio",
|
||||
[
|
||||
io.Combo.Input("lm_model", display_name="model", options=sorted(lm_models)),
|
||||
io.Int.Input("lm_load_for_seconds", display_name="load_for_seconds", default=0, min=-1, max=60 * 60, step=1, advanced=True, tooltip="How long to keep model loaded in LM Studio (seconds)."),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"LM Studio (REST)",
|
||||
[
|
||||
io.Combo.Input("lmstudio_rest_model", display_name="model", options=sorted(lmstudio_rest_models)),
|
||||
io.Int.Input("lmstudio_rest_load_for_seconds", display_name="load_for_seconds", default=0, min=-1, max=60 * 60, step=1, advanced=True, tooltip="Compatibility field for LM Studio REST model load duration."),
|
||||
@@ -481,7 +394,7 @@ class Sage_LLMPromptVision(io.ComfyNode):
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"Ollama (REST)",
|
||||
"Ollama",
|
||||
[
|
||||
io.Combo.Input("ollama_rest_model", display_name="model", options=sorted(ollama_rest_models)),
|
||||
io.String.Input("ollama_rest_keep_alive", display_name="keep_alive", default="5m", advanced=True, tooltip="How long to keep the model loaded after generation (Ollama duration string)."),
|
||||
@@ -504,7 +417,7 @@ class Sage_LLMPromptVision(io.ComfyNode):
|
||||
return io.Schema(
|
||||
node_id="Sage_LLMPromptVision",
|
||||
display_name="LLM Prompt (Vision)",
|
||||
description="Unified provider-switching vision generation node for Ollama, LM Studio, and Native CLIP.",
|
||||
description="Unified provider-switching vision generation node for REST/OpenAI and Native CLIP.",
|
||||
category=f"{SAGE_UTILS_CAT}/LLM",
|
||||
inputs=[
|
||||
io.String.Input("prompt", display_name="prompt", default=DEFAULT_VISION_PROMPT, multiline=True),
|
||||
@@ -527,27 +440,7 @@ class Sage_LLMPromptVision(io.ComfyNode):
|
||||
native_sampling = _get_native_sampling_config(provider_data)
|
||||
|
||||
try:
|
||||
if provider == "Ollama":
|
||||
return Sage_OllamaLLMPromptVision.execute(
|
||||
prompt=prompt,
|
||||
model=provider_data.get("ollama_model"),
|
||||
image=image,
|
||||
seed=seed,
|
||||
keep_alive=provider_data.get("ollama_keep_alive", 0.0),
|
||||
options=provider_data.get("ollama_options") or {},
|
||||
system_prompt=provider_data.get("ollama_system_prompt", ""),
|
||||
)
|
||||
|
||||
if provider == "LM Studio":
|
||||
return Sage_LMStudioLLMPromptVision.execute(
|
||||
prompt=prompt,
|
||||
model=provider_data.get("lm_model"),
|
||||
image=image,
|
||||
seed=seed,
|
||||
load_for_seconds=provider_data.get("lm_load_for_seconds", 0),
|
||||
)
|
||||
|
||||
if provider == "LM Studio (REST)":
|
||||
model = str(provider_data.get("lmstudio_rest_model") or "").strip()
|
||||
if not model or model.startswith("("):
|
||||
return io.NodeOutput("")
|
||||
@@ -565,7 +458,7 @@ class Sage_LLMPromptVision(io.ComfyNode):
|
||||
)
|
||||
)
|
||||
|
||||
if provider == "Ollama (REST)":
|
||||
if provider == "Ollama":
|
||||
model = str(provider_data.get("ollama_rest_model") or "").strip()
|
||||
if not model or model.startswith("("):
|
||||
return io.NodeOutput("")
|
||||
@@ -591,13 +484,15 @@ class Sage_LLMPromptVision(io.ComfyNode):
|
||||
options = {
|
||||
"seed": seed,
|
||||
"temperature": provider_data.get("openai_temperature", 0.7),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 1024),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 4096),
|
||||
}
|
||||
# Convert image tensor to base64 for OpenAI-compatible vision payloads.
|
||||
image_base64 = tensor_to_base64_safe(image) if image is not None else None
|
||||
return io.NodeOutput(
|
||||
openai_generate_vision(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
images=image,
|
||||
images=image_base64,
|
||||
options=options,
|
||||
system_prompt=provider_data.get("openai_system_prompt", ""),
|
||||
)
|
||||
@@ -630,22 +525,6 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
"""Unified vision-refine node that performs initial generation plus a refinement pass per selected provider."""
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
ollama_models = get_cached_ollama_vision_models_for_input_types()
|
||||
if not ollama_models:
|
||||
ollama_models = ["(No Ollama vision models available)"]
|
||||
|
||||
ollama_refine_models = get_cached_ollama_models_for_input_types()
|
||||
if not ollama_refine_models:
|
||||
ollama_refine_models = ["(Ollama not available)"]
|
||||
|
||||
lm_models = get_cached_lmstudio_vision_models_for_input_types()
|
||||
if not lm_models:
|
||||
lm_models = ["(No LM Studio vision models available)"]
|
||||
|
||||
lm_refine_models = get_cached_lmstudio_models_for_input_types()
|
||||
if not lm_refine_models:
|
||||
lm_refine_models = ["(LM Studio not available)"]
|
||||
|
||||
lmstudio_rest_models = get_cached_lmstudio_rest_vision_models_for_input_types()
|
||||
if not lmstudio_rest_models:
|
||||
lmstudio_rest_models = ["(No LM Studio REST vision models available)"]
|
||||
@@ -690,22 +569,8 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
]
|
||||
|
||||
provider_options = _ordered_provider_options([
|
||||
io.DynamicCombo.Option(
|
||||
"Ollama",
|
||||
[
|
||||
io.Combo.Input("ollama_model", display_name="model", options=sorted(ollama_models)),
|
||||
io.Combo.Input("ollama_refine_model", display_name="refine_model", options=sorted(ollama_refine_models)),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"LM Studio",
|
||||
[
|
||||
io.Combo.Input("lm_model", display_name="model", options=sorted(lm_models)),
|
||||
io.Combo.Input("lm_refine_model", display_name="refine_model", options=sorted(lm_refine_models)),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"LM Studio (REST)",
|
||||
[
|
||||
io.Combo.Input("lmstudio_rest_model", display_name="model", options=sorted(lmstudio_rest_models)),
|
||||
io.Combo.Input("lmstudio_rest_refine_model", display_name="refine_model", options=sorted(lmstudio_rest_refine_models)),
|
||||
@@ -713,7 +578,7 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option(
|
||||
"Ollama (REST)",
|
||||
"Ollama",
|
||||
[
|
||||
io.Combo.Input("ollama_rest_model", display_name="model", options=sorted(ollama_rest_models)),
|
||||
io.Combo.Input("ollama_rest_refine_model", display_name="refine_model", options=sorted(ollama_rest_refine_models)),
|
||||
@@ -729,7 +594,7 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
io.Combo.Input("openai_refine_model", display_name="refine_model", options=sorted(openai_refine_models)),
|
||||
io.String.Input("openai_system_prompt", display_name="system_prompt", default="", multiline=True, optional=True, advanced=True),
|
||||
io.Float.Input("openai_temperature", display_name="temperature", default=0.7, min=0.0, max=2.0, step=0.01, advanced=True),
|
||||
io.Int.Input("openai_max_tokens", display_name="max_tokens", default=1024, min=1, max=16384, step=1, advanced=True),
|
||||
io.Int.Input("openai_max_tokens", display_name="max_tokens", default=4096, min=1, max=16384, step=1, advanced=True),
|
||||
],
|
||||
),
|
||||
io.DynamicCombo.Option("Native", native_inputs),
|
||||
@@ -766,29 +631,7 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
native_sampling = _get_native_sampling_config(provider_data)
|
||||
|
||||
try:
|
||||
if provider == "Ollama":
|
||||
return Sage_OllamaLLMPromptVisionRefine.execute(
|
||||
prompt=prompt,
|
||||
model=provider_data.get("ollama_model"),
|
||||
image=image,
|
||||
seed=seed,
|
||||
refine_prompt=refine_prompt,
|
||||
refine_model=provider_data.get("ollama_refine_model"),
|
||||
refine_seed=refine_seed,
|
||||
)
|
||||
|
||||
if provider == "LM Studio":
|
||||
return Sage_LMStudioLLMPromptVisionRefine.execute(
|
||||
prompt=prompt,
|
||||
model=provider_data.get("lm_model"),
|
||||
image=image,
|
||||
seed=seed,
|
||||
refine_prompt=refine_prompt,
|
||||
refine_model=provider_data.get("lm_refine_model"),
|
||||
refine_seed=refine_seed,
|
||||
)
|
||||
|
||||
if provider == "LM Studio (REST)":
|
||||
model = str(provider_data.get("lmstudio_rest_model") or "").strip()
|
||||
if not model or model.startswith("("):
|
||||
return io.NodeOutput("", "")
|
||||
@@ -817,7 +660,7 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
pbar.update(1)
|
||||
return io.NodeOutput(initial, refined)
|
||||
|
||||
if provider == "Ollama (REST)":
|
||||
if provider == "Ollama":
|
||||
model = str(provider_data.get("ollama_rest_model") or "").strip()
|
||||
if not model or model.startswith("("):
|
||||
return io.NodeOutput("", "")
|
||||
@@ -857,12 +700,14 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
options = {
|
||||
"seed": seed,
|
||||
"temperature": provider_data.get("openai_temperature", 0.7),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 1024),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 4096),
|
||||
}
|
||||
# Convert image tensor to base64 for OpenAI-compatible vision payloads.
|
||||
image_base64 = tensor_to_base64_safe(image) if image is not None else None
|
||||
initial = openai_generate_vision(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
images=image,
|
||||
images=image_base64,
|
||||
options=options,
|
||||
system_prompt=provider_data.get("openai_system_prompt", ""),
|
||||
)
|
||||
@@ -877,7 +722,7 @@ class Sage_LLMPromptVisionRefine(io.ComfyNode):
|
||||
options={
|
||||
"seed": refine_seed,
|
||||
"temperature": provider_data.get("openai_temperature", 0.7),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 1024),
|
||||
"max_tokens": provider_data.get("openai_max_tokens", 4096),
|
||||
},
|
||||
system_prompt=provider_data.get("openai_system_prompt", ""),
|
||||
)
|
||||
@@ -934,9 +779,13 @@ def _should_reraise_llm_node_errors() -> bool:
|
||||
|
||||
def _get_default_provider_key() -> str:
|
||||
"""Return the configured default provider key for both nodes and sidebar interoperability."""
|
||||
value = str(get_setting('default_llm_provider', 'ollama') or 'ollama').strip().lower()
|
||||
value = str(get_setting('default_llm_provider', 'lmstudio_rest') or 'lmstudio_rest').strip().lower()
|
||||
if value == 'lmstudio':
|
||||
value = 'lmstudio_rest'
|
||||
elif value == 'ollama':
|
||||
value = 'ollama_rest'
|
||||
if value not in _PROVIDER_LABEL_BY_KEY:
|
||||
return 'ollama'
|
||||
return 'lmstudio_rest'
|
||||
return value
|
||||
|
||||
|
||||
@@ -962,21 +811,19 @@ def _ordered_provider_options(options: list[io.DynamicCombo.Option]) -> list[io.
|
||||
return options
|
||||
|
||||
key_by_label = {
|
||||
"ollama": "Ollama",
|
||||
"lmstudio": "LM Studio",
|
||||
"lmstudio_rest": "LM Studio (REST)",
|
||||
"ollama_rest": "Ollama (REST)",
|
||||
"lmstudio_rest": "LM Studio",
|
||||
"ollama_rest": "Ollama",
|
||||
"openai": "OpenAI",
|
||||
"native": "Native",
|
||||
}
|
||||
default_label = key_by_label.get(_get_default_provider_key(), "Ollama")
|
||||
default_label = key_by_label.get(_get_default_provider_key(), "LM Studio")
|
||||
|
||||
options_by_label = {opt.key: opt for opt in options}
|
||||
ordered = []
|
||||
if default_label in options_by_label:
|
||||
ordered.append(options_by_label[default_label])
|
||||
|
||||
for label in ("Ollama", "LM Studio", "LM Studio (REST)", "Ollama (REST)", "OpenAI", "Native"):
|
||||
for label in ("LM Studio", "Ollama", "OpenAI", "Native"):
|
||||
option = options_by_label.get(label)
|
||||
if option and option not in ordered:
|
||||
ordered.append(option)
|
||||
@@ -1054,4 +901,4 @@ LLM_NODES = [
|
||||
Sage_LLMPromptText,
|
||||
Sage_LLMPromptVision,
|
||||
Sage_LLMPromptVisionRefine,
|
||||
] + OLLAMA_NODES + LMSTUDIO_NODES
|
||||
]
|
||||
+13
-2
@@ -79,8 +79,19 @@ class Sage_SettingsManager:
|
||||
|
||||
# Group settings by category
|
||||
categories = {
|
||||
"LLM Integration": ["default_llm_provider", "enable_ollama", "enable_lmstudio", "ollama_use_custom_url",
|
||||
"ollama_custom_url", "lmstudio_use_custom_url", "lmstudio_custom_url"]
|
||||
"LLM Integration": [
|
||||
"default_llm_provider",
|
||||
"enable_lmstudio_rest",
|
||||
"enable_ollama_rest",
|
||||
"enable_openai",
|
||||
"ollama_use_custom_url",
|
||||
"ollama_custom_url",
|
||||
"lmstudio_use_custom_url",
|
||||
"lmstudio_custom_url",
|
||||
"lmstudio_api_token",
|
||||
"openai_use_custom_url",
|
||||
"openai_base_url",
|
||||
]
|
||||
}
|
||||
|
||||
for category, setting_keys in categories.items():
|
||||
|
||||
+1
-6
@@ -1,8 +1,3 @@
|
||||
# Both of these dependencies are for llm nodes. Sage Utils will check if they were imported, and
|
||||
# will not show nodes that require them if they are not installed.
|
||||
# If you want to use LLM nodes, you must install one of these, depending on your preference.
|
||||
|
||||
# LLM support is a work in progress, and subject to revision.
|
||||
ollama>=0.5.1
|
||||
lmstudio
|
||||
dynamicprompts
|
||||
cryptography
|
||||
|
||||
+82
-529
@@ -196,17 +196,22 @@ def register_routes(routes_instance):
|
||||
@route_error_handler
|
||||
async def get_llm_status(request):
|
||||
"""
|
||||
Check if Ollama, LM Studio, and Native providers are available and enabled.
|
||||
Check if REST/OpenAI/Native providers are available and enabled.
|
||||
|
||||
Returns:
|
||||
{
|
||||
"success": true,
|
||||
"ollama": {
|
||||
"lmstudio_rest": {
|
||||
"available": bool,
|
||||
"enabled": bool,
|
||||
"url": str (if custom)
|
||||
},
|
||||
"lmstudio": {
|
||||
"ollama_rest": {
|
||||
"available": bool,
|
||||
"enabled": bool,
|
||||
"url": str (if custom)
|
||||
},
|
||||
"openai": {
|
||||
"available": bool,
|
||||
"enabled": bool,
|
||||
"url": str (if custom)
|
||||
@@ -228,15 +233,11 @@ def register_routes(routes_instance):
|
||||
logger.warning(f"LLM initialization warning: {init_error}")
|
||||
|
||||
# Check settings
|
||||
ollama_enabled = get_setting("enable_ollama", True)
|
||||
lmstudio_enabled = get_setting("enable_lmstudio", True)
|
||||
lmstudio_rest_enabled = get_setting("enable_lmstudio_rest", False)
|
||||
ollama_rest_enabled = get_setting("enable_ollama_rest", False)
|
||||
openai_enabled = get_setting("enable_openai", False)
|
||||
|
||||
# Check availability
|
||||
ollama_available = llm.OLLAMA_AVAILABLE and ollama_enabled
|
||||
lmstudio_available = llm.LMSTUDIO_AVAILABLE and lmstudio_enabled
|
||||
lmstudio_rest_available = llm.LMSTUDIO_REST_AVAILABLE and lmstudio_rest_enabled
|
||||
ollama_rest_available = llm.OLLAMA_REST_AVAILABLE and ollama_rest_enabled
|
||||
openai_available = llm.OPENAI_AVAILABLE and openai_enabled
|
||||
@@ -244,14 +245,6 @@ def register_routes(routes_instance):
|
||||
native_available = len(native_models) > 0
|
||||
|
||||
# Get custom URLs if applicable
|
||||
ollama_info = {
|
||||
"available": ollama_available,
|
||||
"enabled": ollama_enabled
|
||||
}
|
||||
lmstudio_info = {
|
||||
"available": lmstudio_available,
|
||||
"enabled": lmstudio_enabled
|
||||
}
|
||||
lmstudio_rest_info = {
|
||||
"available": lmstudio_rest_available,
|
||||
"enabled": lmstudio_rest_enabled,
|
||||
@@ -265,11 +258,7 @@ def register_routes(routes_instance):
|
||||
"enabled": openai_enabled,
|
||||
}
|
||||
|
||||
if get_setting("ollama_use_custom_url", False):
|
||||
ollama_info["url"] = get_setting("ollama_custom_url", "")
|
||||
|
||||
if get_setting("lmstudio_use_custom_url", False):
|
||||
lmstudio_info["url"] = get_setting("lmstudio_custom_url", "")
|
||||
lmstudio_rest_info["url"] = get_setting("lmstudio_custom_url", "")
|
||||
|
||||
if get_setting("openai_use_custom_url", False):
|
||||
@@ -281,8 +270,6 @@ def register_routes(routes_instance):
|
||||
}
|
||||
|
||||
return success_response(data={
|
||||
"ollama": ollama_info,
|
||||
"lmstudio": lmstudio_info,
|
||||
"lmstudio_rest": lmstudio_rest_info,
|
||||
"ollama_rest": ollama_rest_info,
|
||||
"openai": openai_info,
|
||||
@@ -296,14 +283,14 @@ def register_routes(routes_instance):
|
||||
_route_list.append({
|
||||
"method": "GET",
|
||||
"path": "/sage_llm/status",
|
||||
"description": "Check Ollama and LM Studio availability"
|
||||
"description": "Check REST/OpenAI/Native availability"
|
||||
})
|
||||
|
||||
@routes_instance.get('/sage_llm/models')
|
||||
@route_error_handler
|
||||
async def get_llm_models(request):
|
||||
"""
|
||||
Get available text models from Ollama, LM Studio, and Native CLIP.
|
||||
Get available text models from REST providers, OpenAI, and Native CLIP.
|
||||
|
||||
Query params:
|
||||
force: bool - Force re-initialization of LLM providers (default: False)
|
||||
@@ -312,13 +299,15 @@ def register_routes(routes_instance):
|
||||
{
|
||||
"success": true,
|
||||
"models": {
|
||||
"ollama": ["model1", "model2", ...],
|
||||
"lmstudio": ["model1", "model2", ...],
|
||||
"lmstudio_rest": ["model1", "model2", ...],
|
||||
"ollama_rest": ["model1", "model2", ...],
|
||||
"openai": ["model1", "model2", ...],
|
||||
"native": ["clip1.safetensors", ...]
|
||||
},
|
||||
"status": {
|
||||
"ollama_available": bool,
|
||||
"lmstudio_available": bool,
|
||||
"lmstudio_rest_available": bool,
|
||||
"ollama_rest_available": bool,
|
||||
"openai_available": bool,
|
||||
"native_available": bool
|
||||
}
|
||||
}
|
||||
@@ -338,52 +327,22 @@ def register_routes(routes_instance):
|
||||
# Initialize LLM services if needed
|
||||
llm.ensure_llm_initialized()
|
||||
|
||||
# Get models from both providers
|
||||
ollama_models = []
|
||||
lmstudio_models = []
|
||||
# Get models from configured providers
|
||||
lmstudio_rest_models = []
|
||||
native_models = _get_native_clip_models()
|
||||
ollama_tool_models = []
|
||||
lmstudio_tool_models = []
|
||||
lmstudio_rest_tool_models = []
|
||||
ollama_rest_tool_models = []
|
||||
openai_tool_models = []
|
||||
ollama_reasoning_models = []
|
||||
lmstudio_reasoning_models = []
|
||||
lmstudio_rest_reasoning_models = []
|
||||
ollama_rest_reasoning_models = []
|
||||
openai_reasoning_models = []
|
||||
ollama_capabilities = {}
|
||||
lmstudio_capabilities = {}
|
||||
lmstudio_rest_capabilities = {}
|
||||
ollama_rest_capabilities = {}
|
||||
openai_capabilities = {}
|
||||
|
||||
ollama_enabled = get_setting("enable_ollama", True)
|
||||
lmstudio_enabled = get_setting("enable_lmstudio", True)
|
||||
lmstudio_rest_enabled = get_setting("enable_lmstudio_rest", False)
|
||||
ollama_rest_enabled = get_setting("enable_ollama_rest", False)
|
||||
openai_enabled = get_setting("enable_openai", False)
|
||||
|
||||
if ollama_enabled and llm.OLLAMA_AVAILABLE:
|
||||
try:
|
||||
models = llm.get_ollama_models()
|
||||
ollama_models = get_compatible_models('ollama', models)
|
||||
ollama_tool_models = get_compatible_models('ollama', llm.get_ollama_tool_models())
|
||||
ollama_reasoning_models = get_compatible_models('ollama', llm.get_ollama_reasoning_models())
|
||||
ollama_capabilities = llm.get_ollama_model_capabilities_map()
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get Ollama models: {e}")
|
||||
|
||||
if lmstudio_enabled and llm.LMSTUDIO_AVAILABLE:
|
||||
try:
|
||||
models = llm.get_lmstudio_models()
|
||||
lmstudio_models = get_compatible_models('lmstudio', models)
|
||||
lmstudio_tool_models = get_compatible_models('lmstudio', llm.get_lmstudio_tool_models())
|
||||
lmstudio_reasoning_models = get_compatible_models('lmstudio', llm.get_lmstudio_reasoning_models())
|
||||
lmstudio_capabilities = llm.get_lmstudio_model_capabilities_map()
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get LM Studio models: {e}")
|
||||
|
||||
if lmstudio_rest_enabled and llm.LMSTUDIO_REST_AVAILABLE:
|
||||
try:
|
||||
@@ -434,40 +393,30 @@ def register_routes(routes_instance):
|
||||
|
||||
return success_response(data={
|
||||
"models": {
|
||||
"ollama": ollama_models,
|
||||
"lmstudio": lmstudio_models,
|
||||
"lmstudio_rest": lmstudio_rest_models,
|
||||
"ollama_rest": ollama_rest_models,
|
||||
"openai": openai_models,
|
||||
"native": native_models,
|
||||
},
|
||||
"capabilities": {
|
||||
"ollama": ollama_capabilities,
|
||||
"lmstudio": lmstudio_capabilities,
|
||||
"lmstudio_rest": lmstudio_rest_capabilities,
|
||||
"ollama_rest": ollama_rest_capabilities,
|
||||
"openai": openai_capabilities,
|
||||
"native": native_capabilities,
|
||||
},
|
||||
"tool_models": {
|
||||
"ollama": ollama_tool_models,
|
||||
"lmstudio": lmstudio_tool_models,
|
||||
"lmstudio_rest": lmstudio_rest_tool_models,
|
||||
"ollama_rest": ollama_rest_tool_models,
|
||||
"openai": openai_tool_models,
|
||||
"native": [],
|
||||
},
|
||||
"reasoning_models": {
|
||||
"ollama": ollama_reasoning_models,
|
||||
"lmstudio": lmstudio_reasoning_models,
|
||||
"lmstudio_rest": lmstudio_rest_reasoning_models,
|
||||
"ollama_rest": ollama_rest_reasoning_models,
|
||||
"openai": openai_reasoning_models,
|
||||
"native": [],
|
||||
},
|
||||
"status": {
|
||||
"ollama_available": len(ollama_models) > 0,
|
||||
"lmstudio_available": len(lmstudio_models) > 0,
|
||||
"lmstudio_rest_available": len(lmstudio_rest_models) > 0,
|
||||
"ollama_rest_available": len(ollama_rest_models) > 0,
|
||||
"openai_available": len(openai_models) > 0,
|
||||
@@ -489,7 +438,7 @@ def register_routes(routes_instance):
|
||||
@route_error_handler
|
||||
async def get_llm_vision_models(request):
|
||||
"""
|
||||
Get available vision models from Ollama and LM Studio.
|
||||
Get available vision models from REST providers and OpenAI.
|
||||
|
||||
Query params:
|
||||
force: bool - Force re-initialization of LLM providers (default: False)
|
||||
@@ -498,13 +447,15 @@ def register_routes(routes_instance):
|
||||
{
|
||||
"success": true,
|
||||
"models": {
|
||||
"ollama": ["llava", "bakllava", ...],
|
||||
"lmstudio": ["llava-v1.6", ...],
|
||||
"lmstudio_rest": ["llava-v1.6", ...],
|
||||
"ollama_rest": ["llava", "bakllava", ...],
|
||||
"openai": ["gpt-4o", ...],
|
||||
"native": []
|
||||
},
|
||||
"status": {
|
||||
"ollama_available": bool,
|
||||
"lmstudio_available": bool,
|
||||
"lmstudio_rest_available": bool,
|
||||
"ollama_rest_available": bool,
|
||||
"openai_available": bool,
|
||||
"native_available": bool
|
||||
}
|
||||
}
|
||||
@@ -524,52 +475,22 @@ def register_routes(routes_instance):
|
||||
# Initialize LLM services if needed
|
||||
llm.ensure_llm_initialized()
|
||||
|
||||
# Get vision models from both providers
|
||||
ollama_models = []
|
||||
lmstudio_models = []
|
||||
# Get vision models from configured providers
|
||||
lmstudio_rest_models = []
|
||||
native_models = []
|
||||
ollama_tool_models = []
|
||||
lmstudio_tool_models = []
|
||||
lmstudio_rest_tool_models = []
|
||||
ollama_rest_tool_models = []
|
||||
openai_tool_models = []
|
||||
ollama_reasoning_models = []
|
||||
lmstudio_reasoning_models = []
|
||||
lmstudio_rest_reasoning_models = []
|
||||
ollama_rest_reasoning_models = []
|
||||
openai_reasoning_models = []
|
||||
ollama_capabilities = {}
|
||||
lmstudio_capabilities = {}
|
||||
lmstudio_rest_capabilities = {}
|
||||
ollama_rest_capabilities = {}
|
||||
openai_capabilities = {}
|
||||
|
||||
ollama_enabled = get_setting("enable_ollama", True)
|
||||
lmstudio_enabled = get_setting("enable_lmstudio", True)
|
||||
lmstudio_rest_enabled = get_setting("enable_lmstudio_rest", False)
|
||||
ollama_rest_enabled = get_setting("enable_ollama_rest", False)
|
||||
openai_enabled = get_setting("enable_openai", False)
|
||||
|
||||
if ollama_enabled and llm.OLLAMA_AVAILABLE:
|
||||
try:
|
||||
models = llm.get_ollama_vision_models()
|
||||
ollama_models = get_compatible_models('ollama', models)
|
||||
ollama_tool_models = get_compatible_models('ollama', llm.get_ollama_tool_models())
|
||||
ollama_reasoning_models = get_compatible_models('ollama', llm.get_ollama_reasoning_models())
|
||||
ollama_capabilities = llm.get_ollama_model_capabilities_map()
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get Ollama vision models: {e}")
|
||||
|
||||
if lmstudio_enabled and llm.LMSTUDIO_AVAILABLE:
|
||||
try:
|
||||
models = llm.get_lmstudio_vision_models()
|
||||
lmstudio_models = get_compatible_models('lmstudio', models)
|
||||
lmstudio_tool_models = get_compatible_models('lmstudio', llm.get_lmstudio_tool_models())
|
||||
lmstudio_reasoning_models = get_compatible_models('lmstudio', llm.get_lmstudio_reasoning_models())
|
||||
lmstudio_capabilities = llm.get_lmstudio_model_capabilities_map()
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get LM Studio vision models: {e}")
|
||||
|
||||
if lmstudio_rest_enabled and llm.LMSTUDIO_REST_AVAILABLE:
|
||||
try:
|
||||
@@ -608,40 +529,30 @@ def register_routes(routes_instance):
|
||||
|
||||
return success_response(data={
|
||||
"models": {
|
||||
"ollama": ollama_models,
|
||||
"lmstudio": lmstudio_models,
|
||||
"lmstudio_rest": lmstudio_rest_models,
|
||||
"ollama_rest": ollama_rest_models,
|
||||
"openai": openai_vision_models,
|
||||
"native": native_models,
|
||||
},
|
||||
"capabilities": {
|
||||
"ollama": ollama_capabilities,
|
||||
"lmstudio": lmstudio_capabilities,
|
||||
"lmstudio_rest": lmstudio_rest_capabilities,
|
||||
"ollama_rest": ollama_rest_capabilities,
|
||||
"openai": openai_capabilities,
|
||||
"native": native_capabilities,
|
||||
},
|
||||
"tool_models": {
|
||||
"ollama": ollama_tool_models,
|
||||
"lmstudio": lmstudio_tool_models,
|
||||
"lmstudio_rest": lmstudio_rest_tool_models,
|
||||
"ollama_rest": ollama_rest_tool_models,
|
||||
"openai": openai_tool_models,
|
||||
"native": [],
|
||||
},
|
||||
"reasoning_models": {
|
||||
"ollama": ollama_reasoning_models,
|
||||
"lmstudio": lmstudio_reasoning_models,
|
||||
"lmstudio_rest": lmstudio_rest_reasoning_models,
|
||||
"ollama_rest": ollama_rest_reasoning_models,
|
||||
"openai": openai_reasoning_models,
|
||||
"native": [],
|
||||
},
|
||||
"status": {
|
||||
"ollama_available": len(ollama_models) > 0,
|
||||
"lmstudio_available": len(lmstudio_models) > 0,
|
||||
"lmstudio_rest_available": len(lmstudio_rest_models) > 0,
|
||||
"ollama_rest_available": len(ollama_rest_models) > 0,
|
||||
"openai_available": len(openai_vision_models) > 0,
|
||||
@@ -699,7 +610,7 @@ def register_routes(routes_instance):
|
||||
|
||||
Request body:
|
||||
{
|
||||
"provider": "ollama" | "lmstudio" | "native",
|
||||
"provider": "lmstudio_rest" | "ollama_rest" | "openai" | "native",
|
||||
"model": str,
|
||||
"prompt": str,
|
||||
"system_prompt": str (optional),
|
||||
@@ -727,7 +638,7 @@ def register_routes(routes_instance):
|
||||
if not is_valid:
|
||||
return _error_response_with_metadata(error_msg, status=400, error_code='LLM_VALIDATION_ERROR')
|
||||
|
||||
provider = data["provider"].lower()
|
||||
provider = routes_helpers.normalize_provider(data["provider"])
|
||||
model = data["model"]
|
||||
prompt = data["prompt"]
|
||||
system_prompt = data.get("system_prompt", "")
|
||||
@@ -741,42 +652,7 @@ def register_routes(routes_instance):
|
||||
# Generate response based on provider
|
||||
response_text = ""
|
||||
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'Ollama is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='ollama',
|
||||
operation='generate',
|
||||
)
|
||||
|
||||
response_text = llm.ollama_generate(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
system_prompt=system_prompt,
|
||||
keep_alive=0.0,
|
||||
options=options
|
||||
)
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='lmstudio',
|
||||
operation='generate',
|
||||
)
|
||||
|
||||
response_text = llm.lmstudio_generate(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
keep_alive=0,
|
||||
options=options
|
||||
)
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio REST is not available',
|
||||
@@ -873,7 +749,7 @@ def register_routes(routes_instance):
|
||||
|
||||
Request body:
|
||||
{
|
||||
"provider": "ollama" | "lmstudio" | "native",
|
||||
"provider": "lmstudio_rest" | "ollama_rest" | "openai" | "native",
|
||||
"model": str,
|
||||
"prompt": str,
|
||||
"system_prompt": str (optional),
|
||||
@@ -903,16 +779,16 @@ def register_routes(routes_instance):
|
||||
error_code='LLM_VALIDATION_ERROR',
|
||||
)
|
||||
|
||||
provider = data["provider"].lower()
|
||||
provider = routes_helpers.normalize_provider(data["provider"])
|
||||
model = data["model"]
|
||||
prompt = data["prompt"]
|
||||
system_prompt = data.get("system_prompt", "")
|
||||
options = data.get("options", {})
|
||||
|
||||
# Validate provider
|
||||
if provider not in ["ollama", "lmstudio", "lmstudio_rest", "ollama_rest", "openai", "native"]:
|
||||
if provider not in ["lmstudio_rest", "ollama_rest", "openai", "native"]:
|
||||
return _error_response_with_metadata(
|
||||
f"Invalid provider: {provider}. Must be 'ollama', 'lmstudio', 'lmstudio_rest', 'ollama_rest', 'openai', or 'native'",
|
||||
f"Invalid provider: {provider}. Must be 'lmstudio', 'ollama', 'openai', or 'native'",
|
||||
status=400,
|
||||
error_code='LLM_VALIDATION_ERROR',
|
||||
)
|
||||
@@ -931,58 +807,7 @@ def register_routes(routes_instance):
|
||||
|
||||
try:
|
||||
# Generate streaming response based on provider
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
error_chunk = _sse_error_chunk(
|
||||
'Ollama is not available',
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='ollama',
|
||||
operation='generate_stream',
|
||||
)
|
||||
await response.write(error_chunk.encode('utf-8'))
|
||||
await response.write_eof()
|
||||
return response
|
||||
|
||||
for chunk_data in llm.ollama_generate_stream(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
system_prompt=system_prompt,
|
||||
keep_alive=0.0,
|
||||
options=options
|
||||
):
|
||||
# Send chunk as SSE
|
||||
sse_data = routes_helpers.format_sse_chunk(chunk_data)
|
||||
await response.write(sse_data.encode('utf-8'))
|
||||
|
||||
if chunk_data.get("done", False):
|
||||
break
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
error_chunk = _sse_error_chunk(
|
||||
'LM Studio is not available',
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='lmstudio',
|
||||
operation='generate_stream',
|
||||
)
|
||||
await response.write(error_chunk.encode('utf-8'))
|
||||
await response.write_eof()
|
||||
return response
|
||||
|
||||
for chunk_data in llm.lmstudio_generate_stream(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
keep_alive=0,
|
||||
options=options
|
||||
):
|
||||
# Send chunk as SSE
|
||||
sse_data = routes_helpers.format_sse_chunk(chunk_data)
|
||||
await response.write(sse_data.encode('utf-8'))
|
||||
|
||||
if chunk_data.get("done", False):
|
||||
break
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
error_chunk = _sse_error_chunk(
|
||||
'LM Studio REST is not available',
|
||||
@@ -1127,7 +952,7 @@ def register_routes(routes_instance):
|
||||
|
||||
Request body:
|
||||
{
|
||||
"provider": "ollama" | "lmstudio",
|
||||
"provider": "lmstudio_rest" | "ollama_rest" | "openai",
|
||||
"model": str,
|
||||
"prompt": str,
|
||||
"images": [base64_encoded_image, ...],
|
||||
@@ -1152,7 +977,7 @@ def register_routes(routes_instance):
|
||||
if not is_valid:
|
||||
return _error_response_with_metadata(error_msg, status=400, error_code='LLM_VALIDATION_ERROR')
|
||||
|
||||
provider = data["provider"].lower()
|
||||
provider = routes_helpers.normalize_provider(data["provider"])
|
||||
model = data["model"]
|
||||
|
||||
# Check vision capability before attempting dispatch
|
||||
@@ -1177,92 +1002,7 @@ def register_routes(routes_instance):
|
||||
|
||||
response_text = ""
|
||||
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'Ollama is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='ollama',
|
||||
operation='vision_generate',
|
||||
)
|
||||
|
||||
# Build Ollama vision generation parameters
|
||||
response_parameters = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"stream": False,
|
||||
"images": images_data, # Already base64
|
||||
"keep_alive": 0.0
|
||||
}
|
||||
|
||||
if system_prompt:
|
||||
response_parameters["system"] = system_prompt
|
||||
|
||||
if options:
|
||||
response_parameters["options"] = options
|
||||
|
||||
response = llm.ollama_client.generate(**response_parameters)
|
||||
|
||||
if not response or 'response' not in response:
|
||||
return _error_response_with_metadata(
|
||||
'No valid response received from model',
|
||||
status=500,
|
||||
error_code='LLM_EMPTY_RESPONSE',
|
||||
provider='ollama',
|
||||
operation='vision_generate',
|
||||
)
|
||||
|
||||
response_text = clean_response(response['response'])
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='lmstudio',
|
||||
operation='vision_generate',
|
||||
)
|
||||
|
||||
# LM Studio requires temp file conversion
|
||||
temp_files = []
|
||||
try:
|
||||
temp_files = routes_helpers.decode_base64_images_to_temp(images_data)
|
||||
|
||||
# Use LM Studio's vision generation
|
||||
import lmstudio as lms
|
||||
keep_alive = options.get('keep_alive', 0)
|
||||
|
||||
lms_model = lms.llm(model, ttl=keep_alive) if keep_alive >= 1 else lms.llm(model)
|
||||
|
||||
# Create chat with system prompt (if provided)
|
||||
chat = lms.Chat(system_prompt) if system_prompt else lms.Chat()
|
||||
|
||||
# Prepare image handles
|
||||
image_handles = [lms.prepare_image(img_path) for img_path in temp_files]
|
||||
chat.add_user_message(prompt, images=image_handles)
|
||||
|
||||
response = lms_model.respond(chat)
|
||||
|
||||
if keep_alive < 1:
|
||||
lms_model.unload()
|
||||
|
||||
if not response:
|
||||
return _error_response_with_metadata(
|
||||
'No valid response received from model',
|
||||
status=500,
|
||||
error_code='LLM_EMPTY_RESPONSE',
|
||||
provider='lmstudio',
|
||||
operation='vision_generate',
|
||||
)
|
||||
|
||||
response_text = clean_response(response.content)
|
||||
|
||||
finally:
|
||||
routes_helpers.cleanup_temp_files(temp_files)
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio REST is not available',
|
||||
@@ -1344,7 +1084,7 @@ def register_routes(routes_instance):
|
||||
|
||||
Request body:
|
||||
{
|
||||
"provider": "ollama" | "lmstudio",
|
||||
"provider": "lmstudio_rest" | "ollama_rest" | "openai",
|
||||
"model": str,
|
||||
"prompt": str,
|
||||
"images": [base64_encoded_image, ...],
|
||||
@@ -1363,7 +1103,7 @@ def register_routes(routes_instance):
|
||||
if not is_valid:
|
||||
return _error_response_with_metadata(error_msg, status=400, error_code='LLM_VALIDATION_ERROR')
|
||||
|
||||
provider = data["provider"].lower()
|
||||
provider = routes_helpers.normalize_provider(data["provider"])
|
||||
model = data["model"]
|
||||
|
||||
# Check vision capability before attempting dispatch
|
||||
@@ -1394,85 +1134,7 @@ def register_routes(routes_instance):
|
||||
await response.prepare(request)
|
||||
|
||||
try:
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
error_chunk = _sse_error_chunk(
|
||||
'Ollama is not available',
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='ollama',
|
||||
operation='vision_generate_stream',
|
||||
)
|
||||
await response.write(error_chunk.encode('utf-8'))
|
||||
await response.write_eof()
|
||||
return response
|
||||
|
||||
# Ollama vision streaming with base64 images
|
||||
response_parameters = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
"stream": True,
|
||||
"images": images_data,
|
||||
"keep_alive": 0.0
|
||||
}
|
||||
|
||||
if system_prompt:
|
||||
response_parameters["system"] = system_prompt
|
||||
|
||||
if options:
|
||||
response_parameters["options"] = options
|
||||
|
||||
full_response = ""
|
||||
|
||||
for chunk in llm.ollama_client.generate(**response_parameters):
|
||||
if 'response' in chunk:
|
||||
chunk_text = chunk['response']
|
||||
full_response += chunk_text
|
||||
|
||||
chunk_data = {
|
||||
"chunk": chunk_text,
|
||||
"done": chunk.get('done', False)
|
||||
}
|
||||
sse_data = routes_helpers.format_sse_chunk(chunk_data)
|
||||
await response.write(sse_data.encode('utf-8'))
|
||||
|
||||
if chunk.get('done', False):
|
||||
# Send final message
|
||||
final_data = {
|
||||
"chunk": "",
|
||||
"done": True,
|
||||
"full_response": clean_response(full_response)
|
||||
}
|
||||
sse_data = routes_helpers.format_sse_chunk(final_data)
|
||||
await response.write(sse_data.encode('utf-8'))
|
||||
break
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
error_chunk = _sse_error_chunk(
|
||||
'LM Studio is not available',
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='lmstudio',
|
||||
operation='vision_generate_stream',
|
||||
)
|
||||
await response.write(error_chunk.encode('utf-8'))
|
||||
await response.write_eof()
|
||||
return response
|
||||
|
||||
keep_alive = options.get('keep_alive', 0)
|
||||
for chunk_data in llm.lmstudio_generate_vision_stream(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
keep_alive=keep_alive,
|
||||
images=images_data,
|
||||
options=options,
|
||||
):
|
||||
sse_data = routes_helpers.format_sse_chunk(chunk_data)
|
||||
await response.write(sse_data.encode('utf-8'))
|
||||
|
||||
if chunk_data.get('done', False):
|
||||
break
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
error_chunk = _sse_error_chunk(
|
||||
'LM Studio REST is not available',
|
||||
@@ -2039,7 +1701,7 @@ def register_routes(routes_instance):
|
||||
# Get all presets
|
||||
builtin_presets = {
|
||||
'descriptive_prompt': {
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': 'gemma3:12b',
|
||||
'promptTemplate': 'description/Descriptive Prompt',
|
||||
'systemPrompt': 'e621_prompt_generator',
|
||||
@@ -2051,7 +1713,7 @@ def register_routes(routes_instance):
|
||||
}
|
||||
},
|
||||
'e621_description': {
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': 'gemma3:12b',
|
||||
'promptTemplate': 'description/Descriptive Prompt',
|
||||
'systemPrompt': 'e621_prompt_generator',
|
||||
@@ -2063,7 +1725,7 @@ def register_routes(routes_instance):
|
||||
}
|
||||
},
|
||||
'casual_chat': {
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': None,
|
||||
'promptTemplate': '',
|
||||
'systemPrompt': 'default',
|
||||
@@ -2089,7 +1751,7 @@ def register_routes(routes_instance):
|
||||
return error_response(f"Preset '{preset_id}' not found", status=404)
|
||||
|
||||
# Get provider and model
|
||||
provider = preset.get('provider', 'ollama')
|
||||
provider = routes_helpers.normalize_provider(preset.get('provider', 'lmstudio_rest'))
|
||||
model = preset.get('model')
|
||||
|
||||
if not model:
|
||||
@@ -2153,7 +1815,7 @@ def register_routes(routes_instance):
|
||||
}
|
||||
|
||||
# Add provider-specific options
|
||||
if provider == 'ollama':
|
||||
if provider == 'ollama_rest':
|
||||
if 'top_k' in settings:
|
||||
options['top_k'] = settings['top_k']
|
||||
if 'top_p' in settings:
|
||||
@@ -2167,79 +1829,7 @@ def register_routes(routes_instance):
|
||||
# Generate response
|
||||
response_text = ""
|
||||
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
return error_response("Ollama is not available", status=503)
|
||||
|
||||
# Build parameters for Ollama vision generation
|
||||
response_parameters = {
|
||||
"model": model,
|
||||
"prompt": prompt_text,
|
||||
"stream": False,
|
||||
"images": images_data,
|
||||
"keep_alive": settings.get('keepAlive', 300) / 60.0 # Convert seconds to minutes
|
||||
}
|
||||
|
||||
if system_prompt_text:
|
||||
response_parameters["system"] = system_prompt_text
|
||||
|
||||
if options:
|
||||
response_parameters["options"] = options
|
||||
|
||||
response = llm.ollama_client.generate(**response_parameters)
|
||||
|
||||
if not response or 'response' not in response:
|
||||
return error_response("No valid response received from model", status=500)
|
||||
|
||||
response_text = clean_response(response['response'])
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
return error_response("LM Studio is not available", status=503)
|
||||
|
||||
# Convert base64 to temp files for LM Studio
|
||||
import base64
|
||||
import tempfile
|
||||
import os
|
||||
import lmstudio as lms
|
||||
|
||||
temp_files = []
|
||||
try:
|
||||
for img_b64 in images_data:
|
||||
img_bytes = base64.b64decode(img_b64)
|
||||
temp_fd, temp_path = tempfile.mkstemp(suffix='.png')
|
||||
os.close(temp_fd)
|
||||
with open(temp_path, 'wb') as f:
|
||||
f.write(img_bytes)
|
||||
temp_files.append(temp_path)
|
||||
|
||||
keep_alive = settings.get('keepAlive', 0)
|
||||
lms_model = lms.llm(model, ttl=keep_alive) if keep_alive >= 1 else lms.llm(model)
|
||||
|
||||
# Create chat with system prompt (if provided)
|
||||
chat = lms.Chat(system_prompt_text) if system_prompt_text else lms.Chat()
|
||||
|
||||
image_handles = [lms.prepare_image(img_path) for img_path in temp_files]
|
||||
chat.add_user_message(prompt_text, images=image_handles)
|
||||
|
||||
lms_response = lms_model.respond(chat)
|
||||
|
||||
if keep_alive < 1:
|
||||
lms_model.unload()
|
||||
|
||||
if not lms_response:
|
||||
return error_response("No valid response received from model", status=500)
|
||||
|
||||
response_text = clean_response(lms_response.content)
|
||||
|
||||
finally:
|
||||
for temp_path in temp_files:
|
||||
try:
|
||||
os.unlink(temp_path)
|
||||
except:
|
||||
pass
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
return error_response("LM Studio REST is not available", status=503)
|
||||
|
||||
@@ -2251,6 +1841,34 @@ def register_routes(routes_instance):
|
||||
options=options,
|
||||
system_prompt=system_prompt_text,
|
||||
)
|
||||
|
||||
elif provider == "ollama_rest":
|
||||
if not llm.OLLAMA_REST_AVAILABLE:
|
||||
return error_response("Ollama REST is not available", status=503)
|
||||
|
||||
response_text = llm.ollama_rest_generate_vision(
|
||||
model=model,
|
||||
prompt=prompt_text,
|
||||
images=images_data,
|
||||
options=options,
|
||||
system_prompt=system_prompt_text,
|
||||
keep_alive=str(settings.get('keepAlive', '5m') or '5m'),
|
||||
)
|
||||
|
||||
elif provider == "openai":
|
||||
if not llm.OPENAI_AVAILABLE:
|
||||
return error_response("OpenAI provider is not available", status=503)
|
||||
|
||||
response_text = llm.openai_generate_vision(
|
||||
model=model,
|
||||
prompt=prompt_text,
|
||||
images=images_data,
|
||||
options=options,
|
||||
system_prompt=system_prompt_text,
|
||||
)
|
||||
|
||||
else:
|
||||
return error_response(f"Unsupported preset provider: {provider}", status=400)
|
||||
|
||||
return success_response({
|
||||
"response": response_text,
|
||||
@@ -2290,7 +1908,7 @@ def register_routes(routes_instance):
|
||||
|
||||
Request body:
|
||||
{
|
||||
"provider": "ollama" | "lmstudio" | "lmstudio_rest" | "ollama_rest" | "openai" | "native",
|
||||
"provider": "lmstudio_rest" | "ollama_rest" | "openai" | "native",
|
||||
"model": str,
|
||||
"keep_alive": float | int (optional, seconds; default 60)
|
||||
}
|
||||
@@ -2307,13 +1925,13 @@ def register_routes(routes_instance):
|
||||
_clear_last_llm_error()
|
||||
data = await request.json()
|
||||
|
||||
provider = str(data.get("provider", "")).lower()
|
||||
provider = routes_helpers.normalize_provider(str(data.get("provider", "")))
|
||||
model = str(data.get("model", ""))
|
||||
keep_alive = data.get("keep_alive", 60)
|
||||
|
||||
if not provider or provider not in ("ollama", "lmstudio", "lmstudio_rest", "ollama_rest", "openai", "native"):
|
||||
if not provider or provider not in ("lmstudio_rest", "ollama_rest", "openai", "native"):
|
||||
return _error_response_with_metadata(
|
||||
f"Invalid or missing provider: {provider!r}. Must be 'ollama', 'lmstudio', 'lmstudio_rest', 'ollama_rest', 'openai', or 'native'",
|
||||
f"Invalid or missing provider: {provider!r}. Must be 'lmstudio', 'ollama', 'openai', or 'native'",
|
||||
status=400,
|
||||
error_code='LLM_VALIDATION_ERROR',
|
||||
)
|
||||
@@ -2328,39 +1946,7 @@ def register_routes(routes_instance):
|
||||
|
||||
llm.ensure_llm_initialized()
|
||||
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'Ollama is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='ollama',
|
||||
operation='load_model',
|
||||
)
|
||||
loaded = llm.ollama_preload_model(model=model, keep_alive=float(keep_alive))
|
||||
if not loaded:
|
||||
return _error_response_with_metadata(
|
||||
f"Failed to preload model '{model}' in Ollama",
|
||||
status=500,
|
||||
error_code='LLM_LOAD_ERROR',
|
||||
provider='ollama',
|
||||
operation='load_model',
|
||||
)
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='lmstudio',
|
||||
operation='load_model',
|
||||
)
|
||||
lms_model = llm.lmstudio_load_model(model=model, keep_alive=int(keep_alive))
|
||||
if keep_alive < 1:
|
||||
llm.lmstudio_unload_model(lms_model)
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio REST is not available',
|
||||
@@ -2431,7 +2017,7 @@ def register_routes(routes_instance):
|
||||
_route_list.append({
|
||||
"method": "POST",
|
||||
"path": "/sage_llm/load_model",
|
||||
"description": "Load a model into memory (Ollama pre-warm / LM Studio load)"
|
||||
"description": "Validate/preload selected model for REST/OpenAI/Native providers"
|
||||
})
|
||||
|
||||
@routes_instance.post('/sage_llm/generate_only')
|
||||
@@ -2447,7 +2033,7 @@ def register_routes(routes_instance):
|
||||
|
||||
Request body:
|
||||
{
|
||||
"provider": "ollama" | "lmstudio" | "native",
|
||||
"provider": "lmstudio_rest" | "ollama_rest" | "openai" | "native",
|
||||
"model": str,
|
||||
"prompt": str,
|
||||
"system_prompt": str (optional, Ollama only),
|
||||
@@ -2475,7 +2061,7 @@ def register_routes(routes_instance):
|
||||
if not is_valid:
|
||||
return _error_response_with_metadata(error_msg, status=400, error_code='LLM_VALIDATION_ERROR')
|
||||
|
||||
provider = data["provider"].lower()
|
||||
provider = routes_helpers.normalize_provider(data["provider"])
|
||||
model = data["model"]
|
||||
prompt = data["prompt"]
|
||||
system_prompt = data.get("system_prompt", "")
|
||||
@@ -2488,40 +2074,7 @@ def register_routes(routes_instance):
|
||||
|
||||
response_text = ""
|
||||
|
||||
if provider == "ollama":
|
||||
if not llm.OLLAMA_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'Ollama is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='ollama',
|
||||
operation='generate_only',
|
||||
)
|
||||
response_text = llm.ollama_generate_preloaded(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
keep_alive=float(keep_alive),
|
||||
options=options,
|
||||
system_prompt=system_prompt,
|
||||
)
|
||||
|
||||
elif provider == "lmstudio":
|
||||
if not llm.LMSTUDIO_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio is not available',
|
||||
status=503,
|
||||
error_code='LLM_PROVIDER_UNAVAILABLE',
|
||||
provider='lmstudio',
|
||||
operation='generate_only',
|
||||
)
|
||||
lms_model = llm.lmstudio_load_model(model=model, keep_alive=int(keep_alive))
|
||||
try:
|
||||
response_text = llm.lmstudio_generate_with_model(lms_model, prompt=prompt, options=options)
|
||||
finally:
|
||||
if keep_alive < 1:
|
||||
llm.lmstudio_unload_model(lms_model)
|
||||
|
||||
elif provider == "lmstudio_rest":
|
||||
if provider == "lmstudio_rest":
|
||||
if not llm.LMSTUDIO_REST_AVAILABLE:
|
||||
return _error_response_with_metadata(
|
||||
'LM Studio REST is not available',
|
||||
|
||||
+3
-4
@@ -103,13 +103,12 @@ try:
|
||||
if settings.save():
|
||||
# Check if LLM-related settings were updated and trigger lazy initialization
|
||||
llm_settings = {
|
||||
'enable_ollama',
|
||||
'enable_lmstudio',
|
||||
'enable_lmstudio_rest',
|
||||
'enable_ollama_rest',
|
||||
'enable_openai',
|
||||
'ollama_custom_url',
|
||||
'lmstudio_custom_url',
|
||||
'custom_ollama_url',
|
||||
'custom_lmstudio_url',
|
||||
'openai_base_url',
|
||||
}
|
||||
if any(setting in llm_settings for setting in updated_settings):
|
||||
try:
|
||||
|
||||
+74
-38
@@ -3,38 +3,83 @@
|
||||
from typing import Any
|
||||
|
||||
from ..logger import get_logger
|
||||
from .providers import lmstudio_rest_client as lmstudio_rest_provider
|
||||
from .providers import ollama_rest_client as ollama_rest_provider
|
||||
from .providers import openai_client as openai_provider
|
||||
|
||||
logger = get_logger('llm.init')
|
||||
|
||||
|
||||
def init_ollama_rest(enabled: bool) -> bool:
|
||||
"""Initialize Ollama REST provider state."""
|
||||
if not enabled:
|
||||
logger.info('Ollama is disabled in settings.')
|
||||
return False
|
||||
|
||||
try:
|
||||
initialized = ollama_rest_provider.is_running(True)
|
||||
if initialized:
|
||||
logger.info('Ollama REST provider initialized.')
|
||||
else:
|
||||
logger.info('Ollama REST provider is enabled but server is not reachable yet.')
|
||||
return initialized
|
||||
except Exception as e:
|
||||
logger.error(f'Failed to initialize Ollama REST provider: {e}')
|
||||
return False
|
||||
|
||||
|
||||
def init_lmstudio_rest(enabled: bool) -> bool:
|
||||
"""Initialize LM Studio REST provider state."""
|
||||
if not enabled:
|
||||
logger.info('LM Studio is disabled in settings.')
|
||||
return False
|
||||
|
||||
try:
|
||||
initialized = lmstudio_rest_provider.is_running(True)
|
||||
if initialized:
|
||||
logger.info('LM Studio REST provider initialized.')
|
||||
else:
|
||||
logger.info('LM Studio REST provider is enabled but server is not reachable yet.')
|
||||
return initialized
|
||||
except Exception as e:
|
||||
logger.error(f'Failed to initialize LM Studio REST provider: {e}')
|
||||
return False
|
||||
|
||||
|
||||
def init_openai_provider(enabled: bool) -> bool:
|
||||
"""Initialize OpenAI-compatible provider state."""
|
||||
if not enabled:
|
||||
logger.info('OpenAI provider is disabled in settings.')
|
||||
return False
|
||||
|
||||
try:
|
||||
initialized = openai_provider.is_running(True)
|
||||
if initialized:
|
||||
logger.info('OpenAI provider initialized.')
|
||||
else:
|
||||
logger.info('OpenAI provider is enabled but endpoint is not reachable yet.')
|
||||
return initialized
|
||||
except Exception as e:
|
||||
logger.error(f'Failed to initialize OpenAI provider: {e}')
|
||||
return False
|
||||
|
||||
|
||||
def init_ollama_client(
|
||||
ollama_available: bool,
|
||||
ollama_module: Any,
|
||||
enabled: bool,
|
||||
custom_url: str,
|
||||
) -> tuple[Any, bool]:
|
||||
"""Create and return an Ollama client and initialized state."""
|
||||
if not ollama_available:
|
||||
logger.warning('Ollama library is not available.')
|
||||
return None, False
|
||||
if ollama_module is None:
|
||||
logger.warning('Ollama module is not loaded.')
|
||||
return None, False
|
||||
if not enabled:
|
||||
logger.info('Ollama is disabled in settings.')
|
||||
return None, False
|
||||
"""Compatibility alias for legacy callers.
|
||||
|
||||
try:
|
||||
if custom_url and custom_url.strip():
|
||||
client = ollama_module.Client(host=custom_url)
|
||||
logger.info('Ollama client initialized with custom URL.')
|
||||
else:
|
||||
client = ollama_module.Client()
|
||||
logger.info('Ollama client initialized with default URL.')
|
||||
return client, True
|
||||
except Exception as e:
|
||||
logger.error(f'Failed to initialize Ollama client: {e}')
|
||||
return None, False
|
||||
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(
|
||||
@@ -43,21 +88,12 @@ def init_lmstudio_client(
|
||||
enabled: bool,
|
||||
custom_url: str,
|
||||
) -> bool:
|
||||
"""Initialize LM Studio configuration and return initialized state."""
|
||||
if not lmstudio_available or lms_module is None:
|
||||
logger.info('LM Studio is not available.')
|
||||
return False
|
||||
if not enabled:
|
||||
logger.info('LM Studio is disabled in settings.')
|
||||
return False
|
||||
"""Compatibility alias for legacy callers.
|
||||
|
||||
try:
|
||||
if custom_url and custom_url.strip():
|
||||
lms_module.get_default_client(custom_url)
|
||||
logger.info('LM Studio client configured with custom URL.')
|
||||
else:
|
||||
logger.info('LM Studio using default configuration.')
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f'Failed to configure LM Studio: {e}')
|
||||
return False
|
||||
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)
|
||||
|
||||
@@ -1,22 +1,16 @@
|
||||
"""Provider-specific LLM helpers."""
|
||||
|
||||
from .settings import (
|
||||
is_ollama_enabled,
|
||||
is_lmstudio_enabled,
|
||||
is_lmstudio_rest_enabled,
|
||||
is_ollama_rest_enabled,
|
||||
is_openai_enabled,
|
||||
)
|
||||
from . import ollama_client, lmstudio_client, lmstudio_rest_client, ollama_rest_client, openai_client
|
||||
from . import lmstudio_rest_client, ollama_rest_client, openai_client
|
||||
|
||||
__all__ = [
|
||||
'is_ollama_enabled',
|
||||
'is_lmstudio_enabled',
|
||||
'is_lmstudio_rest_enabled',
|
||||
'is_ollama_rest_enabled',
|
||||
'is_openai_enabled',
|
||||
'ollama_client',
|
||||
'lmstudio_client',
|
||||
'lmstudio_rest_client',
|
||||
'ollama_rest_client',
|
||||
'openai_client',
|
||||
|
||||
@@ -1,613 +0,0 @@
|
||||
"""LM Studio provider operations extracted from llm/service.py."""
|
||||
|
||||
import base64
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Any, cast
|
||||
|
||||
from ..cache import get_llm_cache
|
||||
from ...logger import get_logger
|
||||
from ...helpers_image import tensor_to_temp_image
|
||||
from ..common import clean_response, build_lmstudio_config
|
||||
from ..errors import raise_llm_error, report_llm_error, stringify_llm_error
|
||||
from ..capabilities import ModelCapabilities, get_capability_cache
|
||||
|
||||
logger = get_logger('llm.providers.lmstudio')
|
||||
|
||||
_PROVIDER_NAME = 'lmstudio'
|
||||
|
||||
|
||||
def _prepare_image_paths(images) -> tuple[list[str], bool]:
|
||||
"""Normalize tensor or base64 image inputs into temporary file paths."""
|
||||
input_images = tensor_to_temp_image(images) if images is not None else []
|
||||
if input_images:
|
||||
return input_images, False
|
||||
|
||||
image_entries = images if isinstance(images, list) else [images] if images is not None else []
|
||||
temp_paths: list[str] = []
|
||||
|
||||
for image in image_entries:
|
||||
if not isinstance(image, str) or not image.strip():
|
||||
continue
|
||||
|
||||
image_data = image.strip()
|
||||
if image_data.startswith('data:') and ',' in image_data:
|
||||
image_data = image_data.split(',', 1)[1]
|
||||
|
||||
image_bytes = base64.b64decode(image_data)
|
||||
temp_fd, temp_path = tempfile.mkstemp(suffix='.png')
|
||||
os.close(temp_fd)
|
||||
with open(temp_path, 'wb') as file_handle:
|
||||
file_handle.write(image_bytes)
|
||||
temp_paths.append(temp_path)
|
||||
|
||||
return temp_paths, True
|
||||
|
||||
|
||||
def _cleanup_temp_paths(paths: list[str]) -> None:
|
||||
for path in paths:
|
||||
try:
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _safe_get(mapping_or_obj: Any, key: str, default=None):
|
||||
if isinstance(mapping_or_obj, dict):
|
||||
return mapping_or_obj.get(key, default)
|
||||
return getattr(mapping_or_obj, key, default)
|
||||
|
||||
|
||||
def _detect_capabilities_from_model_object(model_obj: Any, model_name: str) -> ModelCapabilities:
|
||||
print(f"Detecting capabilities for model '{model_name}' from model object.")
|
||||
info = _safe_get(model_obj, 'info', None)
|
||||
capabilities_obj = _safe_get(model_obj, 'capabilities', None)
|
||||
|
||||
vision = bool(_safe_get(info, 'vision', False))
|
||||
tool_use = bool(_safe_get(info, 'trained_for_tool_use', False) or _safe_get(info, 'tool_use', False))
|
||||
reasoning = False
|
||||
thinking = False
|
||||
print(f"Model info for '{model_name}': {info}")
|
||||
print(f"Model capabilities object for '{model_name}': {capabilities_obj}")
|
||||
|
||||
if isinstance(capabilities_obj, dict):
|
||||
print(f"Capabilities object for model '{model_name}': {capabilities_obj}")
|
||||
vision = vision or bool(capabilities_obj.get('vision'))
|
||||
tool_use = tool_use or bool(capabilities_obj.get('trained_for_tool_use') or capabilities_obj.get('tool_use'))
|
||||
# Check for reasoning in API metadata
|
||||
reasoning_obj = capabilities_obj.get('reasoning')
|
||||
if isinstance(reasoning_obj, dict):
|
||||
reasoning = True
|
||||
# reasoning implies thinking capability
|
||||
thinking = True
|
||||
|
||||
context_window = _safe_get(info, 'max_context_length', None)
|
||||
if context_window is None:
|
||||
context_window = _safe_get(model_obj, 'max_context_length', None)
|
||||
if context_window is not None:
|
||||
try:
|
||||
context_window = int(context_window)
|
||||
except (TypeError, ValueError):
|
||||
context_window = None
|
||||
|
||||
confidence = 'api' if (info is not None or capabilities_obj is not None) else 'heuristic'
|
||||
|
||||
return ModelCapabilities(
|
||||
name=model_name,
|
||||
provider=_PROVIDER_NAME,
|
||||
vision=vision,
|
||||
tool_use=tool_use,
|
||||
reasoning=reasoning,
|
||||
thinking=thinking,
|
||||
supported_modalities=['text'] + (['image'] if vision else []),
|
||||
context_window=context_window,
|
||||
metadata={
|
||||
'info': info if isinstance(info, dict) else {},
|
||||
'capabilities': capabilities_obj if isinstance(capabilities_obj, dict) else {},
|
||||
},
|
||||
confidence=confidence,
|
||||
)
|
||||
|
||||
|
||||
def get_model_capabilities(lmstudio_available: bool, lms: Any, enabled: bool, model_obj: Any) -> ModelCapabilities:
|
||||
print('Getting model capabilities for model object:', model_obj)
|
||||
print("Info:", model_obj.get_info() if hasattr(model_obj, 'get_info') else 'N/A')
|
||||
model_name = str(_safe_get(model_obj, 'model_key', '') or _safe_get(model_obj, 'id', '') or 'unknown')
|
||||
if not lmstudio_available or lms is None or not enabled:
|
||||
print('LM Studio not available or not enabled, returning default capabilities with low confidence.')
|
||||
return ModelCapabilities(name=model_name, provider=_PROVIDER_NAME, confidence='guess')
|
||||
|
||||
cap_cache = get_capability_cache()
|
||||
cached = cap_cache.get(_PROVIDER_NAME, model_name)
|
||||
if cached is not None:
|
||||
print(f'Capabilities for model "{model_name}" found in cache: {cached}')
|
||||
return cached
|
||||
|
||||
capabilities = _detect_capabilities_from_model_object(model_obj, model_name)
|
||||
cap_cache.set(capabilities)
|
||||
print(f'Capabilities for model "{model_name}" detected and cached: {capabilities}')
|
||||
return capabilities
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# PRIMITIVE FUNCTIONS FOR TWO-PHASE LOADING / GENERATION
|
||||
# ===========================================================================
|
||||
|
||||
def load_model(lmstudio_available: bool, lms: Any, enabled: bool, model: str, keep_alive: int) -> Any:
|
||||
"""Load and return an lms_model handle. Caller is responsible for unloading."""
|
||||
if not lmstudio_available or lms is None:
|
||||
raise_llm_error(ImportError, 'LM Studio is not available.', provider='lmstudio', operation='load_model')
|
||||
if not enabled:
|
||||
raise_llm_error(RuntimeError, 'LM Studio is not enabled.', provider='lmstudio', operation='load_model')
|
||||
logger.info(f"Loading model '{model}'...")
|
||||
lms_model = None
|
||||
try:
|
||||
if keep_alive >= 1:
|
||||
lms_model = lms.llm(model, ttl=keep_alive)
|
||||
else:
|
||||
lms_model = lms.llm(model)
|
||||
except Exception as e:
|
||||
normalized_cause = RuntimeError(stringify_llm_error(e))
|
||||
raise_llm_error(
|
||||
RuntimeError,
|
||||
f"Failed to load model '{model}'",
|
||||
provider='lmstudio',
|
||||
operation='load_model',
|
||||
cause=normalized_cause,
|
||||
)
|
||||
if lms_model is None:
|
||||
raise_llm_error(RuntimeError, f"Failed to load model '{model}'", provider='lmstudio', operation='load_model')
|
||||
logger.info(f"Model '{model}' loaded successfully.")
|
||||
return lms_model
|
||||
|
||||
|
||||
def generate_with_model(lms_model: Any, lms: Any, prompt: str, options) -> str:
|
||||
"""Run text inference on an already-loaded lms_model."""
|
||||
logger.debug(f'Generating with prompt: {prompt[:200]!r}')
|
||||
chat = lms.Chat()
|
||||
chat.add_user_message(prompt)
|
||||
config = cast(Any, build_lmstudio_config(options or {}))
|
||||
if config:
|
||||
response = lms_model.respond(chat, config=config)
|
||||
else:
|
||||
response = lms_model.respond(chat)
|
||||
if not response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the model.', provider='lmstudio', operation='generate_with_model')
|
||||
result = clean_response(response.content)
|
||||
logger.debug(f'Generation complete. Response length: {len(result)} characters.')
|
||||
return result
|
||||
|
||||
|
||||
def generate_stream_with_model(lms_model: Any, lms: Any, prompt: str, options):
|
||||
"""Run streaming text inference on an already-loaded lms_model."""
|
||||
chat = lms.Chat()
|
||||
chat.add_user_message(prompt)
|
||||
config = cast(Any, build_lmstudio_config(options or {}))
|
||||
stream = lms_model.respond_stream(chat, config=config) if config else lms_model.respond_stream(chat)
|
||||
|
||||
full_response = ''
|
||||
for fragment in stream:
|
||||
chunk = getattr(fragment, 'content', '')
|
||||
if not isinstance(chunk, str) or not chunk:
|
||||
continue
|
||||
full_response += chunk
|
||||
yield {
|
||||
'chunk': chunk,
|
||||
'done': False,
|
||||
}
|
||||
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'full_response': clean_response(full_response),
|
||||
}
|
||||
|
||||
|
||||
def generate_vision_with_model(lms_model: Any, lms: Any, prompt: str, images, options) -> str:
|
||||
"""Run vision inference on an already-loaded lms_model."""
|
||||
logger.debug(f'Generating vision with prompt: {prompt[:200]!r}')
|
||||
input_images, should_cleanup = _prepare_image_paths(images)
|
||||
try:
|
||||
chat = lms.Chat()
|
||||
if not input_images:
|
||||
chat.add_user_message(prompt)
|
||||
else:
|
||||
image_handles = [lms.prepare_image(image) for image in input_images]
|
||||
chat.add_user_message(prompt, images=image_handles)
|
||||
config = cast(Any, build_lmstudio_config(options or {}))
|
||||
if config:
|
||||
response = lms_model.respond(chat, config=config)
|
||||
else:
|
||||
response = lms_model.respond(chat)
|
||||
if not response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the model.', provider='lmstudio', operation='generate_vision_with_model')
|
||||
result = clean_response(response.content)
|
||||
logger.debug(f'Generation complete. Response length: {len(result)} characters.')
|
||||
return result
|
||||
finally:
|
||||
if should_cleanup:
|
||||
_cleanup_temp_paths(input_images)
|
||||
|
||||
|
||||
def generate_vision_stream_with_model(lms_model: Any, lms: Any, prompt: str, images, options):
|
||||
"""Run streaming vision inference on an already-loaded lms_model."""
|
||||
input_images, should_cleanup = _prepare_image_paths(images)
|
||||
try:
|
||||
chat = lms.Chat()
|
||||
if not input_images:
|
||||
raise_llm_error(ValueError, 'No images provided for vision model.', provider='lmstudio', operation='generate_vision_stream_with_model')
|
||||
|
||||
image_handles = [lms.prepare_image(image) for image in input_images]
|
||||
chat.add_user_message(prompt, images=image_handles)
|
||||
config = cast(Any, build_lmstudio_config(options or {}))
|
||||
stream = lms_model.respond_stream(chat, config=config) if config else lms_model.respond_stream(chat)
|
||||
|
||||
full_response = ''
|
||||
for fragment in stream:
|
||||
chunk = getattr(fragment, 'content', '')
|
||||
if not isinstance(chunk, str) or not chunk:
|
||||
continue
|
||||
full_response += chunk
|
||||
yield {
|
||||
'chunk': chunk,
|
||||
'done': False,
|
||||
}
|
||||
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'full_response': clean_response(full_response),
|
||||
}
|
||||
finally:
|
||||
if should_cleanup:
|
||||
_cleanup_temp_paths(input_images)
|
||||
|
||||
|
||||
def unload_model(lms_model: Any) -> None:
|
||||
"""Unload an lms_model handle."""
|
||||
logger.debug('Unloading model.')
|
||||
lms_model.unload()
|
||||
|
||||
|
||||
def is_running(lmstudio_available: bool, lms: Any, enabled: bool) -> bool:
|
||||
"""Check if LM Studio server is running by attempting a lightweight API call."""
|
||||
if not lmstudio_available or lms is None:
|
||||
return False
|
||||
|
||||
if not enabled:
|
||||
return False
|
||||
|
||||
try:
|
||||
lms.list_downloaded_models('llm')
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def get_models(lmstudio_available: bool, lms: Any, enabled: bool) -> list[str]:
|
||||
"""Retrieve a list of available models from LM Studio."""
|
||||
if not lmstudio_available or lms is None:
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
if not enabled:
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
def _fetch_lmstudio_models():
|
||||
if lms is None or not is_running(lmstudio_available, lms, enabled):
|
||||
return []
|
||||
|
||||
try:
|
||||
logger.debug('Retrieving models from LM Studio...')
|
||||
response = lms.list_downloaded_models('llm')
|
||||
return [model.model_key for model in response if hasattr(model, 'model_key') and model.model_key is not None]
|
||||
except Exception as e:
|
||||
report_llm_error('Error retrieving models from LM Studio', provider='lmstudio', operation='get_models', cause=e)
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'models',
|
||||
_fetch_lmstudio_models,
|
||||
label='LM Studio models',
|
||||
)
|
||||
|
||||
|
||||
def get_vision_models(lmstudio_available: bool, lms: Any, enabled: bool) -> list[str]:
|
||||
"""Retrieve a list of available vision models from LM Studio."""
|
||||
if not lmstudio_available or lms is None:
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
if not enabled:
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
def _fetch_lmstudio_vision_models(cache_instance):
|
||||
if lms is None or not is_running(lmstudio_available, lms, enabled):
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
try:
|
||||
logger.debug('Retrieving vision models from LM Studio...')
|
||||
response = lms.list_downloaded_models('llm')
|
||||
models = []
|
||||
|
||||
for model in response:
|
||||
if not (hasattr(model, 'model_key') and model.model_key is not None):
|
||||
continue
|
||||
|
||||
print(model.__dict__)
|
||||
capabilities = get_model_capabilities(lmstudio_available, lms, enabled, model)
|
||||
cache_instance.set_model_capability(_PROVIDER_NAME, str(model.model_key), capabilities.vision)
|
||||
if capabilities.vision:
|
||||
models.append(model.model_key)
|
||||
|
||||
return models
|
||||
except Exception as e:
|
||||
report_llm_error('Error retrieving vision models from LM Studio', provider='lmstudio', operation='get_vision_models', cause=e)
|
||||
return []
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'vision_models',
|
||||
_fetch_lmstudio_vision_models,
|
||||
label='LM Studio vision models',
|
||||
pass_self=True,
|
||||
)
|
||||
|
||||
|
||||
def get_tool_models(lmstudio_available: bool, lms: Any, enabled: bool) -> list[str]:
|
||||
if not lmstudio_available or lms is None or not enabled:
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
def _fetch_tool_models() -> list[str]:
|
||||
if lms is None or not is_running(lmstudio_available, lms, enabled):
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
response = lms.list_downloaded_models('llm')
|
||||
models: list[str] = []
|
||||
for model in response:
|
||||
model_name = str(_safe_get(model, 'model_key', '') or '')
|
||||
if not model_name:
|
||||
continue
|
||||
if get_model_capabilities(lmstudio_available, lms, enabled, model).tool_use:
|
||||
models.append(model_name)
|
||||
return models
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'tool_models',
|
||||
_fetch_tool_models,
|
||||
label='LM Studio tool-capable models',
|
||||
)
|
||||
|
||||
|
||||
def get_reasoning_models(lmstudio_available: bool, lms: Any, enabled: bool) -> list[str]:
|
||||
if not lmstudio_available or lms is None or not enabled:
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
def _fetch_reasoning_models() -> list[str]:
|
||||
if lms is None or not is_running(lmstudio_available, lms, enabled):
|
||||
return ['(LM Studio not available)']
|
||||
|
||||
response = lms.list_downloaded_models('llm')
|
||||
models: list[str] = []
|
||||
for model in response:
|
||||
model_name = str(_safe_get(model, 'model_key', '') or '')
|
||||
if not model_name:
|
||||
continue
|
||||
if get_model_capabilities(lmstudio_available, lms, enabled, model).reasoning:
|
||||
models.append(model_name)
|
||||
return models
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'reasoning_models',
|
||||
_fetch_reasoning_models,
|
||||
label='LM Studio reasoning models',
|
||||
)
|
||||
|
||||
|
||||
def get_model_capabilities_map(lmstudio_available: bool, lms: Any, enabled: bool) -> dict[str, ModelCapabilities]:
|
||||
if not lmstudio_available or lms is None or not enabled:
|
||||
return {}
|
||||
|
||||
if not is_running(lmstudio_available, lms, enabled):
|
||||
return {}
|
||||
|
||||
response = lms.list_downloaded_models('llm')
|
||||
capabilities_map: dict[str, ModelCapabilities] = {}
|
||||
for model in response:
|
||||
model_name = str(_safe_get(model, 'model_key', '') or '')
|
||||
if not model_name:
|
||||
continue
|
||||
capabilities_map[model_name] = get_model_capabilities(lmstudio_available, lms, enabled, model)
|
||||
return capabilities_map
|
||||
|
||||
|
||||
def generate_vision(
|
||||
lmstudio_available: bool,
|
||||
lms: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: int,
|
||||
images,
|
||||
options,
|
||||
) -> str:
|
||||
"""Generate a response from an LM Studio vision model."""
|
||||
if not lmstudio_available or lms is None:
|
||||
raise_llm_error(ImportError, 'LM Studio is not available. Please install it to use this function.', provider='lmstudio', operation='generate_vision')
|
||||
|
||||
model_list = get_vision_models(lmstudio_available, lms, enabled)
|
||||
if model not in model_list:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio', operation='generate_vision')
|
||||
|
||||
lms_model = None
|
||||
try:
|
||||
lms_model = load_model(lmstudio_available, lms, enabled, model, keep_alive)
|
||||
response = generate_vision_with_model(lms_model, lms, prompt, images, options)
|
||||
if keep_alive < 1:
|
||||
unload_model(lms_model)
|
||||
return response
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from LM Studio vision model', provider='lmstudio', operation='generate_vision', cause=e)
|
||||
if lms_model is not None and keep_alive < 1:
|
||||
lms_model.unload()
|
||||
return ''
|
||||
|
||||
|
||||
def generate(
|
||||
lmstudio_available: bool,
|
||||
lms: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: int,
|
||||
options,
|
||||
) -> str:
|
||||
"""Generate a response from an LM Studio model."""
|
||||
if not lmstudio_available or lms is None:
|
||||
raise_llm_error(ImportError, 'LM Studio is not available. Please install it to use this function.', provider='lmstudio', operation='generate')
|
||||
|
||||
model_list = get_models(lmstudio_available, lms, enabled)
|
||||
if model not in model_list:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio', operation='generate')
|
||||
|
||||
lms_model = None
|
||||
try:
|
||||
lms_model = load_model(lmstudio_available, lms, enabled, model, keep_alive)
|
||||
response = generate_with_model(lms_model, lms, prompt, options)
|
||||
if keep_alive < 1:
|
||||
unload_model(lms_model)
|
||||
return response
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from LM Studio', provider='lmstudio', operation='generate', cause=e)
|
||||
if lms_model is not None and keep_alive < 1:
|
||||
lms_model.unload()
|
||||
return ''
|
||||
|
||||
|
||||
def generate_vision_refine(
|
||||
lmstudio_available: bool,
|
||||
lms: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
images,
|
||||
options,
|
||||
refine_model: str,
|
||||
refine_prompt: str,
|
||||
refine_options,
|
||||
) -> tuple[str, str]:
|
||||
"""Generate a response from an LM Studio vision model and refine it with another model."""
|
||||
if not lmstudio_available or lms is None:
|
||||
raise_llm_error(ImportError, 'LM Studio is not available. Please install it to use this function.', provider='lmstudio', operation='generate_vision_refine')
|
||||
|
||||
model_list = get_vision_models(lmstudio_available, lms, enabled)
|
||||
if model not in model_list:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio', operation='generate_vision_refine')
|
||||
|
||||
seed = (options or {}).get('seed', 0)
|
||||
lms_model = None
|
||||
try:
|
||||
lms_model = load_model(lmstudio_available, lms, enabled, model, 0)
|
||||
initial_response = generate_vision_with_model(lms_model, lms, prompt, images, options)
|
||||
|
||||
if refine_model == '':
|
||||
refine_model = model
|
||||
if refine_prompt == '':
|
||||
refine_prompt = prompt
|
||||
|
||||
if refine_model != model:
|
||||
unload_model(lms_model)
|
||||
lms_model = load_model(lmstudio_available, lms, enabled, refine_model, 0)
|
||||
|
||||
refine_prompt = f'{refine_prompt}\n{initial_response}'
|
||||
refine_options = refine_options or {}
|
||||
refine_options['seed'] = seed
|
||||
refined_response = generate_with_model(lms_model, lms, refine_prompt, refine_options)
|
||||
unload_model(lms_model)
|
||||
lms_model = None
|
||||
|
||||
return (initial_response, refined_response)
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from LM Studio model', provider='lmstudio', operation='generate_vision_refine', cause=e)
|
||||
if lms_model is not None:
|
||||
lms_model.unload()
|
||||
return ('', '')
|
||||
|
||||
|
||||
def generate_stream(
|
||||
lmstudio_available: bool,
|
||||
lms: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: int,
|
||||
options,
|
||||
):
|
||||
"""Generate a real streaming response from an LM Studio model."""
|
||||
if not lmstudio_available or lms is None:
|
||||
raise_llm_error(ImportError, 'LM Studio is not available. Please install it to use this function.', provider='lmstudio', operation='generate_stream')
|
||||
|
||||
model_list = get_models(lmstudio_available, lms, enabled)
|
||||
if model not in model_list:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio', operation='generate_stream')
|
||||
|
||||
lms_model = None
|
||||
try:
|
||||
lms_model = load_model(lmstudio_available, lms, enabled, model, keep_alive)
|
||||
for chunk_data in generate_stream_with_model(lms_model, lms, prompt, options):
|
||||
yield chunk_data
|
||||
if keep_alive < 1:
|
||||
unload_model(lms_model)
|
||||
except Exception as e:
|
||||
report_llm_error('Error streaming response from LM Studio', provider='lmstudio', operation='generate_stream', cause=e)
|
||||
if lms_model is not None and keep_alive < 1:
|
||||
lms_model.unload()
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'error': stringify_llm_error(e),
|
||||
}
|
||||
|
||||
|
||||
def generate_vision_stream(
|
||||
lmstudio_available: bool,
|
||||
lms: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: int,
|
||||
images,
|
||||
options,
|
||||
):
|
||||
"""Generate a real streaming response from an LM Studio vision model."""
|
||||
if not lmstudio_available or lms is None:
|
||||
raise_llm_error(ImportError, 'LM Studio is not available. Please install it to use this function.', provider='lmstudio', operation='generate_vision_stream')
|
||||
|
||||
model_list = get_vision_models(lmstudio_available, lms, enabled)
|
||||
if model not in model_list:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {model_list}", provider='lmstudio', operation='generate_vision_stream')
|
||||
|
||||
lms_model = None
|
||||
try:
|
||||
lms_model = load_model(lmstudio_available, lms, enabled, model, keep_alive)
|
||||
if images is None:
|
||||
raise_llm_error(ValueError, 'No images provided for vision model.', provider='lmstudio', operation='generate_vision_stream')
|
||||
|
||||
for chunk_data in generate_vision_stream_with_model(lms_model, lms, prompt, images, options):
|
||||
yield chunk_data
|
||||
if keep_alive < 1:
|
||||
unload_model(lms_model)
|
||||
except Exception as e:
|
||||
report_llm_error('Error streaming response from LM Studio vision model', provider='lmstudio', operation='generate_vision_stream', cause=e)
|
||||
if lms_model is not None and keep_alive < 1:
|
||||
lms_model.unload()
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'error': stringify_llm_error(e),
|
||||
}
|
||||
@@ -36,8 +36,12 @@ def _get_base_url() -> str:
|
||||
|
||||
|
||||
def _get_headers() -> Dict[str, str]:
|
||||
# Optional token support. LM Studio defaults to no auth.
|
||||
token = os.environ.get('LMSTUDIO_API_TOKEN', '')
|
||||
from ...settings import get_setting
|
||||
|
||||
# Prefer explicit setting, then environment variable.
|
||||
token = str(get_setting('lmstudio_api_token', '')).strip()
|
||||
if not token:
|
||||
token = os.environ.get('LMSTUDIO_API_TOKEN', '').strip()
|
||||
return with_bearer_auth({}, token)
|
||||
|
||||
|
||||
|
||||
@@ -1,532 +0,0 @@
|
||||
"""Ollama provider operations extracted from llm/service.py."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..cache import get_llm_cache
|
||||
from ...logger import get_logger
|
||||
from ..common import clean_response, build_response_parameters
|
||||
from ..errors import raise_llm_error, report_llm_error, stringify_llm_error
|
||||
from ..capabilities import ModelCapabilities, get_capability_cache
|
||||
|
||||
logger = get_logger('llm.providers.ollama')
|
||||
|
||||
_PROVIDER_NAME = 'ollama'
|
||||
|
||||
_OLLAMA_UNAVAILABLE_MESSAGE = '(Ollama not available)'
|
||||
|
||||
|
||||
def _ollama_unavailable_models() -> list[str]:
|
||||
return [_OLLAMA_UNAVAILABLE_MESSAGE]
|
||||
|
||||
|
||||
def _is_ollama_unavailable(ollama_available: bool, ollama_client: Any, enabled: bool) -> bool:
|
||||
return (not ollama_available) or ollama_client is None or (not enabled)
|
||||
|
||||
|
||||
def _safe_get(mapping_or_obj: Any, key: str, default=None):
|
||||
if isinstance(mapping_or_obj, dict):
|
||||
return mapping_or_obj.get(key, default)
|
||||
return getattr(mapping_or_obj, key, default)
|
||||
|
||||
|
||||
def _detect_capabilities_from_metadata(model_name: str, model_obj: Any, show_response: Any) -> ModelCapabilities:
|
||||
capabilities_obj = _safe_get(model_obj, 'capabilities', None)
|
||||
capabilities_list: list[str] = []
|
||||
if isinstance(capabilities_obj, list):
|
||||
capabilities_list = [str(item).lower() for item in capabilities_obj]
|
||||
|
||||
show_capabilities = _safe_get(show_response, 'capabilities', None)
|
||||
if isinstance(show_capabilities, list):
|
||||
for item in show_capabilities:
|
||||
lowered = str(item).lower()
|
||||
if lowered not in capabilities_list:
|
||||
capabilities_list.append(lowered)
|
||||
|
||||
details = _safe_get(show_response, 'details', None)
|
||||
families: list[str] = []
|
||||
if isinstance(details, dict):
|
||||
families_obj = details.get('families')
|
||||
if isinstance(families_obj, list):
|
||||
families = [str(item).lower() for item in families_obj]
|
||||
|
||||
template = str(_safe_get(show_response, 'template', '')).lower()
|
||||
lowered_name = model_name.lower()
|
||||
|
||||
vision = ('vision' in capabilities_list) or ('vision' in families) or ('clip' in families)
|
||||
tool_use = any(token in capabilities_list for token in ('tools', 'tool_use', 'function_calling'))
|
||||
thinking = ('thinking' in capabilities_list) or ('<think>' in template)
|
||||
|
||||
# Check for reasoning in API metadata (from /api/show response)
|
||||
reasoning = 'thinking' in capabilities_list or thinking
|
||||
|
||||
context_window = None
|
||||
model_info = _safe_get(show_response, 'model_info', None)
|
||||
if isinstance(model_info, dict):
|
||||
for key, value in model_info.items():
|
||||
if 'context_length' in str(key).lower():
|
||||
try:
|
||||
context_window = int(value)
|
||||
break
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
|
||||
has_api_metadata = bool(capabilities_list or families or template or context_window)
|
||||
confidence = 'api' if has_api_metadata else 'heuristic'
|
||||
|
||||
return ModelCapabilities(
|
||||
name=model_name,
|
||||
provider=_PROVIDER_NAME,
|
||||
vision=vision,
|
||||
tool_use=tool_use,
|
||||
reasoning=reasoning,
|
||||
thinking=thinking,
|
||||
supported_modalities=['text'] + (['image'] if vision else []),
|
||||
context_window=context_window,
|
||||
metadata={
|
||||
'capabilities': capabilities_list,
|
||||
'details': details if isinstance(details, dict) else {},
|
||||
'template': template,
|
||||
},
|
||||
confidence=confidence,
|
||||
)
|
||||
|
||||
|
||||
def get_model_capabilities(ollama_available: bool, ollama_client: Any, enabled: bool, model_name: str, model_obj: Any = None) -> ModelCapabilities:
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return ModelCapabilities(name=model_name, provider=_PROVIDER_NAME, confidence='guess')
|
||||
|
||||
cap_cache = get_capability_cache()
|
||||
cached = cap_cache.get(_PROVIDER_NAME, model_name)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
show_response = {}
|
||||
if ollama_client is not None:
|
||||
try:
|
||||
show_response = ollama_client.show(model_name)
|
||||
except Exception as e:
|
||||
logger.debug(f'Failed to query Ollama show metadata for {model_name}: {e}')
|
||||
|
||||
capabilities = _detect_capabilities_from_metadata(model_name, model_obj or {}, show_response or {})
|
||||
cap_cache.set(capabilities)
|
||||
return capabilities
|
||||
|
||||
|
||||
def get_vision_models(ollama_available: bool, ollama_client: Any, enabled: bool) -> list[str]:
|
||||
"""Retrieve a list of available vision models from Ollama."""
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return _ollama_unavailable_models()
|
||||
|
||||
def _fetch_ollama_vision_models(cache_instance):
|
||||
if ollama_client is None:
|
||||
return _ollama_unavailable_models()
|
||||
|
||||
try:
|
||||
logger.debug('Fetching vision models from Ollama...')
|
||||
response = ollama_client.list()
|
||||
models = []
|
||||
|
||||
for model in response.models:
|
||||
if model.model is None:
|
||||
continue
|
||||
|
||||
logger.debug(f'Checking model: {model.model}')
|
||||
|
||||
capabilities = get_model_capabilities(ollama_available, ollama_client, enabled, str(model.model), model)
|
||||
logger.debug(f'Caching vision capability for {model.model}: {capabilities.vision}')
|
||||
cache_instance.set_model_capability(_PROVIDER_NAME, str(model.model), capabilities.vision)
|
||||
if capabilities.vision:
|
||||
models.append(model.model)
|
||||
|
||||
logger.debug(f'Found {len(models)} vision models.')
|
||||
return models
|
||||
except Exception as e:
|
||||
report_llm_error('Error retrieving vision models from Ollama', provider='ollama', operation='get_vision_models', cause=e)
|
||||
return []
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'vision_models',
|
||||
_fetch_ollama_vision_models,
|
||||
label='Ollama vision models',
|
||||
pass_self=True,
|
||||
)
|
||||
|
||||
|
||||
def get_tool_models(ollama_available: bool, ollama_client: Any, enabled: bool) -> list[str]:
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return _ollama_unavailable_models()
|
||||
|
||||
def _fetch_tool_models() -> list[str]:
|
||||
models = get_models(ollama_available, ollama_client, enabled)
|
||||
return [
|
||||
model_name
|
||||
for model_name in models
|
||||
if not model_name.startswith('(')
|
||||
and get_model_capabilities(ollama_available, ollama_client, enabled, model_name).tool_use
|
||||
]
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'tool_models',
|
||||
_fetch_tool_models,
|
||||
label='Ollama tool-capable models',
|
||||
)
|
||||
|
||||
|
||||
def get_reasoning_models(ollama_available: bool, ollama_client: Any, enabled: bool) -> list[str]:
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return _ollama_unavailable_models()
|
||||
|
||||
def _fetch_reasoning_models() -> list[str]:
|
||||
models = get_models(ollama_available, ollama_client, enabled)
|
||||
return [
|
||||
model_name
|
||||
for model_name in models
|
||||
if not model_name.startswith('(')
|
||||
and get_model_capabilities(ollama_available, ollama_client, enabled, model_name).reasoning
|
||||
]
|
||||
|
||||
cache = get_llm_cache()
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'reasoning_models',
|
||||
_fetch_reasoning_models,
|
||||
label='Ollama reasoning models',
|
||||
)
|
||||
|
||||
|
||||
def get_model_capabilities_map(ollama_available: bool, ollama_client: Any, enabled: bool) -> dict[str, ModelCapabilities]:
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return {}
|
||||
|
||||
model_names = get_models(ollama_available, ollama_client, enabled)
|
||||
capabilities_map: dict[str, ModelCapabilities] = {}
|
||||
for model_name in model_names:
|
||||
if model_name.startswith('('):
|
||||
continue
|
||||
capabilities_map[model_name] = get_model_capabilities(ollama_available, ollama_client, enabled, model_name)
|
||||
return capabilities_map
|
||||
|
||||
|
||||
def get_models(ollama_available: bool, ollama_client: Any, enabled: bool) -> list[str]:
|
||||
"""Retrieve a list of available models from Ollama."""
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return _ollama_unavailable_models()
|
||||
|
||||
def _fetch_ollama_models():
|
||||
if ollama_client is None:
|
||||
return _ollama_unavailable_models()
|
||||
|
||||
try:
|
||||
logger.info('Fetching models from Ollama...')
|
||||
response = ollama_client.list()
|
||||
logger.info(f'Found {len(response.models)} models.')
|
||||
return [model.model for model in response.models if model.model is not None]
|
||||
except Exception as e:
|
||||
report_llm_error('Error retrieving models from Ollama', provider='ollama', operation='get_models', cause=e)
|
||||
return []
|
||||
|
||||
cache = get_llm_cache()
|
||||
logger.debug('Fetching Ollama models from cache...')
|
||||
return cache.get_model_list(
|
||||
_PROVIDER_NAME,
|
||||
'models',
|
||||
_fetch_ollama_models,
|
||||
label='Ollama models',
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# PRIMITIVE FUNCTIONS FOR TWO-PHASE LOADING / GENERATION
|
||||
# ===========================================================================
|
||||
|
||||
def preload_model(ollama_available: bool, ollama_client: Any, enabled: bool, model: str, keep_alive: float) -> bool:
|
||||
"""Pre-warm a model by sending a generate request with an empty prompt.
|
||||
|
||||
Ollama loads the model into GPU memory and responds with done_reason='load'.
|
||||
Subsequent generate calls will not incur the initial load delay as long as
|
||||
keep_alive has not expired.
|
||||
|
||||
Returns True if the model was loaded successfully, False otherwise.
|
||||
"""
|
||||
if _is_ollama_unavailable(ollama_available, ollama_client, enabled):
|
||||
return False
|
||||
|
||||
logger.info(f"Loading model '{model}'...")
|
||||
try:
|
||||
response = ollama_client.generate(model=model, prompt='', keep_alive=keep_alive)
|
||||
loaded = getattr(response, 'done_reason', None) == 'load'
|
||||
if loaded:
|
||||
logger.info(f"Model '{model}' loaded successfully.")
|
||||
else:
|
||||
logger.debug(f"Preload response for '{model}': done_reason={getattr(response, 'done_reason', None)!r}")
|
||||
return True
|
||||
except Exception as e:
|
||||
report_llm_error(f"Failed to preload model '{model}'", provider='ollama', operation='preload_model', cause=e)
|
||||
return False
|
||||
|
||||
|
||||
def generate_preloaded(
|
||||
ollama_available: bool,
|
||||
ollama_client: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: float,
|
||||
options,
|
||||
system_prompt: str,
|
||||
) -> str:
|
||||
"""Generate a response assuming the model is already loaded in Ollama.
|
||||
|
||||
Semantically identical to generate() but signals to the caller that no
|
||||
initial load delay is expected (the model was pre-warmed via preload_model).
|
||||
"""
|
||||
logger.debug(f'Generating with prompt: {prompt[:200]!r}')
|
||||
if not ollama_available or ollama_client is None:
|
||||
raise_llm_error(ImportError, 'Ollama is not available. Please install it to use this function.', provider='ollama', operation='generate_preloaded')
|
||||
|
||||
try:
|
||||
options = options or {}
|
||||
response_parameters = build_response_parameters(model, prompt, keep_alive, options, system_prompt, None)
|
||||
response = ollama_client.generate(**response_parameters)
|
||||
if not response or 'response' not in response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the model.', provider='ollama', operation='generate_preloaded')
|
||||
result = clean_response(response['response'])
|
||||
logger.debug(f'Generation complete. Response length: {len(result)} characters.')
|
||||
return result
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from Ollama (preloaded)', provider='ollama', operation='generate_preloaded', cause=e)
|
||||
return ''
|
||||
|
||||
|
||||
def generate_vision(
|
||||
ollama_available: bool,
|
||||
ollama_client: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: float,
|
||||
images,
|
||||
options,
|
||||
system_prompt: str,
|
||||
) -> str:
|
||||
"""Generate a response from an Ollama vision model."""
|
||||
if not ollama_available or ollama_client is None:
|
||||
raise_llm_error(ImportError, 'Ollama is not available. Please install it to use this function.', provider='ollama', operation='generate_vision')
|
||||
|
||||
vision_models = get_vision_models(ollama_available, ollama_client, enabled)
|
||||
if model not in vision_models:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {vision_models}", provider='ollama', operation='generate_vision')
|
||||
if images is None:
|
||||
raise_llm_error(ValueError, 'No images provided for vision model.', provider='ollama', operation='generate_vision')
|
||||
|
||||
try:
|
||||
options = options or {}
|
||||
options['seed'] = options.get('seed', 0)
|
||||
response_parameters = build_response_parameters(model, prompt, keep_alive, options, system_prompt, images)
|
||||
response = ollama_client.generate(**response_parameters)
|
||||
if not response or 'response' not in response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the model.', provider='ollama', operation='generate_vision')
|
||||
return clean_response(response['response'])
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from Ollama vision model', provider='ollama', operation='generate_vision', cause=e)
|
||||
return ''
|
||||
|
||||
|
||||
def generate(
|
||||
ollama_available: bool,
|
||||
ollama_client: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: float,
|
||||
options,
|
||||
system_prompt: str,
|
||||
) -> str:
|
||||
"""Generate a response from an Ollama model."""
|
||||
if not ollama_available or ollama_client is None:
|
||||
raise_llm_error(ImportError, 'Ollama is not available. Please install it to use this function.', provider='ollama', operation='generate')
|
||||
|
||||
models = get_models(ollama_available, ollama_client, enabled)
|
||||
if model not in models:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {models}", provider='ollama', operation='generate')
|
||||
|
||||
try:
|
||||
options = options or {}
|
||||
response_parameters = build_response_parameters(model, prompt, keep_alive, options, system_prompt, None)
|
||||
response = ollama_client.generate(**response_parameters)
|
||||
if not response or 'response' not in response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the model.', provider='ollama', operation='generate')
|
||||
return clean_response(response['response'])
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from Ollama', provider='ollama', operation='generate', cause=e)
|
||||
return ''
|
||||
|
||||
|
||||
def generate_vision_refine(
|
||||
ollama_available: bool,
|
||||
ollama_client: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
images,
|
||||
options,
|
||||
refine_model: str,
|
||||
refine_prompt: str,
|
||||
refine_options,
|
||||
) -> tuple[str, str]:
|
||||
"""Generate a response from an Ollama vision model and refine it with another model."""
|
||||
if not ollama_available or ollama_client is None:
|
||||
raise_llm_error(ImportError, 'Ollama is not available. Please install it to use this function.', provider='ollama', operation='generate_vision_refine')
|
||||
|
||||
vision_models = get_vision_models(ollama_available, ollama_client, enabled)
|
||||
if model not in vision_models:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {vision_models}", provider='ollama', operation='generate_vision_refine')
|
||||
if images is None:
|
||||
raise_llm_error(ValueError, 'No images provided for vision model.', provider='ollama', operation='generate_vision_refine')
|
||||
|
||||
try:
|
||||
options = options or {}
|
||||
options['seed'] = options.get('seed', 0)
|
||||
refine_options = refine_options or {}
|
||||
refine_options['seed'] = refine_options.get('seed', 0)
|
||||
|
||||
if refine_model == '':
|
||||
refine_model = model
|
||||
if refine_prompt == '':
|
||||
refine_prompt = prompt
|
||||
|
||||
response_parameters = build_response_parameters(model, prompt, 0, options, '', images)
|
||||
response = ollama_client.generate(**response_parameters)
|
||||
if not response or 'response' not in response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the vision model.', provider='ollama', operation='generate_vision_refine')
|
||||
|
||||
initial_response = clean_response(response['response'])
|
||||
refine_prompt = f'{refine_prompt}\n{initial_response}'
|
||||
refine_options['seed'] = options.get('seed', 0)
|
||||
refined_response_parameters = build_response_parameters(refine_model, refine_prompt, 0, refine_options, '', None)
|
||||
|
||||
refined_response = ollama_client.generate(**refined_response_parameters)
|
||||
if not refined_response or 'response' not in refined_response:
|
||||
raise_llm_error(ValueError, 'No valid response received from the refining model.', provider='ollama', operation='generate_vision_refine')
|
||||
|
||||
return (initial_response, clean_response(refined_response['response']))
|
||||
except Exception as e:
|
||||
report_llm_error('Error generating response from Ollama vision model', provider='ollama', operation='generate_vision_refine', cause=e)
|
||||
return ('', '')
|
||||
|
||||
|
||||
def generate_stream(
|
||||
ollama_available: bool,
|
||||
ollama_client: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: float,
|
||||
options,
|
||||
system_prompt: str,
|
||||
):
|
||||
"""Generate a streaming response from an Ollama model."""
|
||||
if not ollama_available or ollama_client is None:
|
||||
raise_llm_error(ImportError, 'Ollama is not available. Please install it to use this function.', provider='ollama', operation='generate_stream')
|
||||
|
||||
models = get_models(ollama_available, ollama_client, enabled)
|
||||
if model not in models:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {models}", provider='ollama', operation='generate_stream')
|
||||
|
||||
try:
|
||||
options = options or {}
|
||||
response_parameters = build_response_parameters(model, prompt, keep_alive, options, system_prompt, None)
|
||||
response_parameters['stream'] = True
|
||||
|
||||
full_response = ''
|
||||
for chunk in ollama_client.generate(**response_parameters):
|
||||
if 'response' not in chunk:
|
||||
continue
|
||||
|
||||
chunk_text = chunk['response']
|
||||
full_response += chunk_text
|
||||
done = chunk.get('done', False)
|
||||
|
||||
yield {
|
||||
'chunk': chunk_text,
|
||||
'done': done,
|
||||
}
|
||||
|
||||
if done:
|
||||
break
|
||||
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'full_response': clean_response(full_response),
|
||||
}
|
||||
except Exception as e:
|
||||
report_llm_error('Error streaming response from Ollama', provider='ollama', operation='generate_stream', cause=e)
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'error': stringify_llm_error(e),
|
||||
}
|
||||
|
||||
|
||||
def generate_vision_stream(
|
||||
ollama_available: bool,
|
||||
ollama_client: Any,
|
||||
enabled: bool,
|
||||
model: str,
|
||||
prompt: str,
|
||||
keep_alive: float,
|
||||
images,
|
||||
options,
|
||||
system_prompt: str,
|
||||
):
|
||||
"""Generate a streaming response from an Ollama vision model."""
|
||||
if not ollama_available or ollama_client is None:
|
||||
raise_llm_error(ImportError, 'Ollama is not available. Please install it to use this function.', provider='ollama', operation='generate_vision_stream')
|
||||
|
||||
vision_models = get_vision_models(ollama_available, ollama_client, enabled)
|
||||
if model not in vision_models:
|
||||
raise_llm_error(ValueError, f"Model '{model}' is not available. Available models: {vision_models}", provider='ollama', operation='generate_vision_stream')
|
||||
if images is None:
|
||||
raise_llm_error(ValueError, 'No images provided for vision model.', provider='ollama', operation='generate_vision_stream')
|
||||
|
||||
try:
|
||||
options = options or {}
|
||||
options['seed'] = options.get('seed', 0)
|
||||
|
||||
response_parameters = build_response_parameters(model, prompt, keep_alive, options, system_prompt, images)
|
||||
response_parameters['stream'] = True
|
||||
|
||||
full_response = ''
|
||||
for chunk in ollama_client.generate(**response_parameters):
|
||||
if 'response' not in chunk:
|
||||
continue
|
||||
|
||||
chunk_text = chunk['response']
|
||||
full_response += chunk_text
|
||||
done = chunk.get('done', False)
|
||||
|
||||
yield {
|
||||
'chunk': chunk_text,
|
||||
'done': done,
|
||||
}
|
||||
|
||||
if done:
|
||||
break
|
||||
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'full_response': clean_response(full_response),
|
||||
}
|
||||
except Exception as e:
|
||||
report_llm_error('Error streaming response from Ollama vision model', provider='ollama', operation='generate_vision_stream', cause=e)
|
||||
yield {
|
||||
'chunk': '',
|
||||
'done': True,
|
||||
'error': stringify_llm_error(e),
|
||||
}
|
||||
@@ -7,7 +7,7 @@ from ..cache import get_llm_cache
|
||||
from ...logger import get_logger
|
||||
from ..common import clean_response
|
||||
from ..errors import raise_llm_error, report_llm_error, stringify_llm_error
|
||||
from ..rest import iter_json_lines, normalize_base_url, normalize_image_data_url, request_json, request_stream, with_bearer_auth
|
||||
from ..rest import iter_json_lines, normalize_base_url, normalize_raw_image_base64, request_json, request_stream, with_bearer_auth
|
||||
from ..capabilities import ModelCapabilities, get_capability_cache
|
||||
|
||||
logger = get_logger('llm.providers.ollama_rest')
|
||||
@@ -36,7 +36,14 @@ def _get_base_url() -> str:
|
||||
|
||||
|
||||
def _get_headers() -> Dict[str, str]:
|
||||
token = os.environ.get('OLLAMA_API_TOKEN', '')
|
||||
from ...settings import get_setting
|
||||
|
||||
# Prefer explicit setting, then environment variable(s).
|
||||
token = str(get_setting('ollama_api_key', '')).strip()
|
||||
if not token:
|
||||
token = os.environ.get('OLLAMA_API_KEY', '').strip()
|
||||
if not token:
|
||||
token = os.environ.get('OLLAMA_API_TOKEN', '').strip()
|
||||
return with_bearer_auth({}, token)
|
||||
|
||||
|
||||
@@ -165,10 +172,7 @@ def _normalize_raw_images(images=None) -> list[str]:
|
||||
image_entries = images if isinstance(images, list) else [images]
|
||||
raw_images: list[str] = []
|
||||
for img in image_entries:
|
||||
img_str = str(img)
|
||||
if img_str.startswith('data:') and ',' in img_str:
|
||||
img_str = img_str.split(',', 1)[1]
|
||||
raw_images.append(img_str)
|
||||
raw_images.append(normalize_raw_image_base64(str(img)))
|
||||
return raw_images
|
||||
|
||||
|
||||
|
||||
@@ -26,6 +26,7 @@ _KNOWN_VISION_MODEL_PREFIXES = (
|
||||
'gpt-4o',
|
||||
'gpt-4-turbo',
|
||||
'gpt-4-vision',
|
||||
'glm-4.6v',
|
||||
'o1',
|
||||
'o3',
|
||||
'o4',
|
||||
@@ -37,6 +38,12 @@ _KNOWN_VISION_MODEL_PREFIXES = (
|
||||
'qwen-vl',
|
||||
'qwen2-vl',
|
||||
'qwen2.5-vl',
|
||||
'qwen3-vl',
|
||||
'qwen3.5-vl',
|
||||
'qwen3.6',
|
||||
'gemma-3',
|
||||
'gemma-4',
|
||||
'llama-joycaption'
|
||||
)
|
||||
|
||||
_OPENAI_MODEL_CAPABILITIES: dict[str, dict[str, bool]] = {
|
||||
|
||||
@@ -18,13 +18,13 @@ def _is_provider_enabled(feature_key: str) -> bool:
|
||||
|
||||
|
||||
def is_ollama_enabled() -> bool:
|
||||
"""Check if Ollama is enabled in settings."""
|
||||
return _is_provider_enabled('enable_ollama')
|
||||
"""Compatibility helper for removed Ollama SDK provider."""
|
||||
return False
|
||||
|
||||
|
||||
def is_lmstudio_enabled() -> bool:
|
||||
"""Check if LM Studio is enabled in settings."""
|
||||
return _is_provider_enabled('enable_lmstudio')
|
||||
"""Compatibility helper for removed LM Studio SDK provider."""
|
||||
return False
|
||||
|
||||
|
||||
def is_lmstudio_rest_enabled() -> bool:
|
||||
|
||||
+79
-4
@@ -1,8 +1,11 @@
|
||||
"""Shared REST helpers for LLM providers."""
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Dict, Generator, Iterator, Optional
|
||||
from urllib.parse import urlparse
|
||||
from urllib import error as urllib_error
|
||||
from urllib import request as urllib_request
|
||||
|
||||
@@ -178,11 +181,83 @@ def iter_json_lines(response: Any) -> Generator[dict[str, Any], None, None]:
|
||||
|
||||
|
||||
def normalize_image_data_url(image_data: str, default_mime: str = 'image/png') -> str:
|
||||
"""Normalize an image payload to a data URL string."""
|
||||
image_value = (image_data or '').strip()
|
||||
if image_value.startswith('data:'):
|
||||
"""Normalize an image payload to a data URL string.
|
||||
|
||||
Accepts:
|
||||
- data URLs containing base64 image content
|
||||
- raw base64 image content
|
||||
- remote http/https image URLs
|
||||
"""
|
||||
image_value = _coerce_image_value(image_data)
|
||||
if _is_http_url(image_value):
|
||||
return image_value
|
||||
return f'data:{default_mime};base64,{image_value}'
|
||||
|
||||
if image_value.startswith('data:'):
|
||||
header, payload = _split_data_url(image_value)
|
||||
if ';base64' not in header.lower():
|
||||
raise ValueError('Image data URL must include ;base64 payload.')
|
||||
normalized_payload = _validate_base64_payload(payload)
|
||||
return f'{header},{normalized_payload}'
|
||||
|
||||
normalized_payload = _validate_base64_payload(image_value)
|
||||
return f'data:{default_mime};base64,{normalized_payload}'
|
||||
|
||||
|
||||
def normalize_raw_image_base64(image_data: str) -> str:
|
||||
"""Normalize an image payload to raw base64 bytes for providers like Ollama."""
|
||||
image_value = _coerce_image_value(image_data)
|
||||
if _is_http_url(image_value):
|
||||
raise ValueError('Remote image URLs are not supported for this provider; provide base64 image data instead.')
|
||||
|
||||
if image_value.startswith('data:'):
|
||||
header, payload = _split_data_url(image_value)
|
||||
if ';base64' not in header.lower():
|
||||
raise ValueError('Image data URL must include ;base64 payload.')
|
||||
return _validate_base64_payload(payload)
|
||||
|
||||
return _validate_base64_payload(image_value)
|
||||
|
||||
|
||||
def _coerce_image_value(image_data: str) -> str:
|
||||
image_value = str(image_data or '').strip()
|
||||
if not image_value:
|
||||
raise ValueError('Image payload is empty.')
|
||||
|
||||
lowered = image_value.lower()
|
||||
if lowered.startswith('tensor(') or lowered.startswith('torch.tensor('):
|
||||
raise ValueError('Image payload appears to be a tensor object; expected base64 image bytes or a data URL.')
|
||||
|
||||
return image_value
|
||||
|
||||
|
||||
def _split_data_url(image_value: str) -> tuple[str, str]:
|
||||
if ',' not in image_value:
|
||||
raise ValueError('Invalid image data URL: missing comma separator.')
|
||||
|
||||
header, payload = image_value.split(',', 1)
|
||||
if not header.lower().startswith('data:image/'):
|
||||
raise ValueError('Invalid image data URL: expected data:image/* MIME type.')
|
||||
if not payload.strip():
|
||||
raise ValueError('Invalid image data URL: missing base64 payload.')
|
||||
return header, payload.strip()
|
||||
|
||||
|
||||
def _validate_base64_payload(payload: str) -> str:
|
||||
compact = ''.join(payload.split())
|
||||
if not compact:
|
||||
raise ValueError('Image payload is empty after whitespace normalization.')
|
||||
try:
|
||||
base64.b64decode(compact, validate=True)
|
||||
except (binascii.Error, ValueError) as e:
|
||||
raise ValueError('Image payload is not valid base64 data.') from e
|
||||
return compact
|
||||
|
||||
|
||||
def _is_http_url(value: str) -> bool:
|
||||
parsed = urlparse(value)
|
||||
if parsed.scheme not in ('http', 'https'):
|
||||
return False
|
||||
return bool(parsed.netloc)
|
||||
|
||||
|
||||
def iter_text_chunks(text: str, chunk_size: int = 5) -> Generator[str, None, None]:
|
||||
|
||||
+25
-23
@@ -14,6 +14,16 @@ from . import raise_llm_error
|
||||
logger = get_logger('llm.routes_helpers')
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def get_compatible_models(provider: str, model_list: list[str]) -> list[str]:
|
||||
"""Filter out placeholder unavailable-provider messages from model lists."""
|
||||
if not model_list:
|
||||
@@ -77,7 +87,7 @@ def load_preset(preset_id: str) -> dict[str, Any]:
|
||||
# Built-in presets
|
||||
builtin_presets = {
|
||||
'descriptive_prompt': {
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': 'gemma3:12b',
|
||||
'promptTemplate': 'description/Descriptive Prompt',
|
||||
'systemPrompt': 'e621_prompt_generator',
|
||||
@@ -89,7 +99,7 @@ def load_preset(preset_id: str) -> dict[str, Any]:
|
||||
},
|
||||
},
|
||||
'e621_description': {
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': 'gemma3:12b',
|
||||
'promptTemplate': 'description/Descriptive Prompt',
|
||||
'systemPrompt': 'e621_prompt_generator',
|
||||
@@ -101,7 +111,7 @@ def load_preset(preset_id: str) -> dict[str, Any]:
|
||||
},
|
||||
},
|
||||
'casual_chat': {
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': None,
|
||||
'promptTemplate': '',
|
||||
'systemPrompt': 'default',
|
||||
@@ -141,13 +151,14 @@ def load_preset(preset_id: str) -> dict[str, Any]:
|
||||
|
||||
def build_llm_options(provider: str, settings: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Build provider-specific options from settings."""
|
||||
provider = normalize_provider(provider)
|
||||
options = {
|
||||
'temperature': settings.get('temperature', 0.7),
|
||||
'seed': settings.get('seed', -1),
|
||||
}
|
||||
|
||||
# Add provider-specific options
|
||||
if provider == 'ollama':
|
||||
if provider == 'ollama_rest':
|
||||
if 'top_k' in settings:
|
||||
options['top_k'] = settings['top_k']
|
||||
if 'top_p' in settings:
|
||||
@@ -205,7 +216,7 @@ def get_available_presets_full(
|
||||
'descriptive_prompt': {
|
||||
'name': 'Descriptive Prompt',
|
||||
'description': 'Generate detailed image descriptions',
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': 'gemma3:12b',
|
||||
'promptTemplate': 'description/Descriptive Prompt',
|
||||
'systemPrompt': 'e621_prompt_generator',
|
||||
@@ -222,7 +233,7 @@ def get_available_presets_full(
|
||||
'e621_description': {
|
||||
'name': 'E621 Image Description',
|
||||
'description': 'Generate E621-style detailed image descriptions',
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': 'gemma3:12b',
|
||||
'promptTemplate': 'description/Descriptive Prompt',
|
||||
'systemPrompt': 'e621_prompt_generator',
|
||||
@@ -239,7 +250,7 @@ def get_available_presets_full(
|
||||
'casual_chat': {
|
||||
'name': 'Casual Chat',
|
||||
'description': 'Friendly conversational assistant',
|
||||
'provider': 'ollama',
|
||||
'provider': 'lmstudio_rest',
|
||||
'model': None,
|
||||
'promptTemplate': '',
|
||||
'systemPrompt': 'default',
|
||||
@@ -302,8 +313,9 @@ def validate_provider(provider: str) -> tuple[bool, Optional[str]]:
|
||||
Returns:
|
||||
Tuple of (is_valid, error_message)
|
||||
"""
|
||||
if provider not in ['ollama', 'lmstudio', 'lmstudio_rest', 'ollama_rest', 'openai', 'native']:
|
||||
return False, f"Invalid provider: {provider}. Must be 'ollama', 'lmstudio', 'lmstudio_rest', 'ollama_rest', 'openai', or 'native'"
|
||||
provider = normalize_provider(provider)
|
||||
if provider not in ['lmstudio_rest', 'ollama_rest', 'openai', 'native']:
|
||||
return False, f"Invalid provider: {provider}. Must be 'lmstudio', 'ollama', 'openai', or 'native'"
|
||||
return True, None
|
||||
|
||||
|
||||
@@ -329,7 +341,7 @@ def validate_vision_data(
|
||||
return False, error_msg
|
||||
|
||||
# Validate provider
|
||||
is_valid, error_msg = validate_provider(data['provider'].lower())
|
||||
is_valid, error_msg = validate_provider(normalize_provider(data['provider']))
|
||||
if not is_valid:
|
||||
return False, error_msg
|
||||
|
||||
@@ -358,7 +370,7 @@ def validate_generation_data(
|
||||
return False, error_msg
|
||||
|
||||
# Validate provider
|
||||
is_valid, error_msg = validate_provider(data['provider'].lower())
|
||||
is_valid, error_msg = validate_provider(normalize_provider(data['provider']))
|
||||
if not is_valid:
|
||||
return False, error_msg
|
||||
|
||||
@@ -378,25 +390,15 @@ def check_model_vision_capability(provider: str, model: str) -> tuple[bool, Opti
|
||||
- If supports_vision is True, error_message is None
|
||||
- If supports_vision is False, error_message describes why
|
||||
"""
|
||||
provider = provider.lower()
|
||||
provider = normalize_provider(provider)
|
||||
|
||||
try:
|
||||
from . import service as llm
|
||||
from .providers.settings import (
|
||||
is_ollama_enabled, is_lmstudio_enabled,
|
||||
is_lmstudio_rest_enabled, is_ollama_rest_enabled,
|
||||
is_openai_enabled
|
||||
)
|
||||
|
||||
# Ensure LLM services are initialized
|
||||
llm.ensure_llm_initialized()
|
||||
|
||||
# Get capability map for the provider
|
||||
if provider == 'ollama':
|
||||
cap_map = llm.get_ollama_model_capabilities_map()
|
||||
elif provider == 'lmstudio':
|
||||
cap_map = llm.get_lmstudio_model_capabilities_map()
|
||||
elif provider == 'lmstudio_rest':
|
||||
if provider == 'lmstudio_rest':
|
||||
cap_map = llm.get_lmstudio_rest_model_capabilities_map()
|
||||
elif provider == 'ollama_rest':
|
||||
cap_map = llm.get_ollama_rest_model_capabilities_map()
|
||||
|
||||
+274
-375
@@ -1,206 +1,143 @@
|
||||
from ..logger import get_logger, get_sageutils_logger
|
||||
from .init import init_ollama_client, init_lmstudio_client
|
||||
from .providers.settings import is_ollama_enabled, is_lmstudio_enabled, is_lmstudio_rest_enabled, is_ollama_rest_enabled, is_openai_enabled
|
||||
from .providers import ollama_client as ollama_provider
|
||||
from .providers import lmstudio_client as lmstudio_provider
|
||||
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,
|
||||
)
|
||||
from .init import (
|
||||
init_lmstudio_rest as _init_lmstudio_rest,
|
||||
init_ollama_rest as _init_ollama_rest,
|
||||
init_openai_provider as _init_openai_provider,
|
||||
)
|
||||
from .providers import lmstudio_rest_client as lmstudio_rest_provider
|
||||
from .providers import ollama_rest_client as ollama_rest_provider
|
||||
from .providers import openai_client as openai_provider
|
||||
|
||||
logger = get_logger('llm')
|
||||
root_logger = get_sageutils_logger()
|
||||
|
||||
# Initialization flags to track if services have been initialized
|
||||
_ollama_initialized = False
|
||||
_lmstudio_initialized = False
|
||||
_lmstudio_rest_initialized = False
|
||||
_ollama_rest_initialized = False
|
||||
_openai_initialized = False
|
||||
|
||||
# Attempt to import ollama, if available. Set a flag if it is not available.
|
||||
try:
|
||||
import ollama
|
||||
OLLAMA_AVAILABLE = True
|
||||
ollama_client = None # Will be initialized in init_ollama
|
||||
except ImportError:
|
||||
ollama = None
|
||||
OLLAMA_AVAILABLE = False
|
||||
ollama_client = None
|
||||
root_logger.warning('Ollama library not found.')
|
||||
|
||||
try:
|
||||
import lmstudio as lms
|
||||
LMSTUDIO_AVAILABLE = True
|
||||
except ImportError:
|
||||
lms = None
|
||||
LMSTUDIO_AVAILABLE = False
|
||||
root_logger.warning('LM Studio library not found.')
|
||||
|
||||
# LM Studio REST provider uses HTTP API and does not require lmstudio package.
|
||||
# REST providers use HTTP APIs and do not require local SDK packages.
|
||||
LMSTUDIO_REST_AVAILABLE = True
|
||||
|
||||
# Ollama REST provider uses HTTP API and does not require ollama package.
|
||||
OLLAMA_REST_AVAILABLE = True
|
||||
|
||||
# OpenAI provider uses HTTP API and does not require openai package.
|
||||
OPENAI_AVAILABLE = True
|
||||
|
||||
|
||||
def get_ollama_vision_models() -> list[str]:
|
||||
"""Retrieve a list of available vision models from Ollama."""
|
||||
return ollama_provider.get_vision_models(OLLAMA_AVAILABLE, ollama_client, is_ollama_enabled())
|
||||
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()
|
||||
|
||||
|
||||
def get_ollama_models() -> list[str]:
|
||||
"""Retrieve a list of available models from Ollama."""
|
||||
return ollama_provider.get_models(OLLAMA_AVAILABLE, ollama_client, is_ollama_enabled())
|
||||
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 ollama_generate_vision(model: str, prompt: str, keep_alive: float = 0.0, images=None, options=None, system_prompt: str = '') -> str:
|
||||
"""Generate a response from an Ollama vision model."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.generate_vision(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama_client,
|
||||
is_ollama_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
images,
|
||||
options,
|
||||
system_prompt,
|
||||
)
|
||||
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 ollama_generate(model: str, prompt: str, keep_alive: float = 0.0, options=None, system_prompt: str = '') -> str:
|
||||
"""Generate a response from an Ollama model."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.generate(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama_client,
|
||||
is_ollama_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
options,
|
||||
system_prompt,
|
||||
)
|
||||
|
||||
|
||||
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]:
|
||||
"""Generate a response from an Ollama vision model and refine it with another model."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.generate_vision_refine(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama_client,
|
||||
is_ollama_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
images,
|
||||
options,
|
||||
refine_model,
|
||||
refine_prompt,
|
||||
refine_options,
|
||||
)
|
||||
|
||||
|
||||
def is_lmstudio_running() -> bool:
|
||||
"""Check if LM Studio server is running by attempting a lightweight API call."""
|
||||
return lmstudio_provider.is_running(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled())
|
||||
|
||||
# ============================================================================
|
||||
# MODEL DISCOVERY
|
||||
# ============================================================================
|
||||
|
||||
def get_lmstudio_models() -> list[str]:
|
||||
"""Retrieve a list of available models from LM Studio."""
|
||||
return lmstudio_provider.get_models(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled())
|
||||
"""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]:
|
||||
"""Retrieve a list of available vision models from LM Studio."""
|
||||
return lmstudio_provider.get_vision_models(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled())
|
||||
"""Legacy alias: retrieve vision models from LM Studio REST."""
|
||||
return lmstudio_rest_provider.get_vision_models(_is_lmstudio_service_enabled())
|
||||
|
||||
|
||||
def is_lmstudio_rest_running() -> bool:
|
||||
"""Check if LM Studio REST server is running."""
|
||||
return lmstudio_rest_provider.is_running(is_lmstudio_rest_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 is_ollama_rest_running() -> bool:
|
||||
"""Check if the Ollama REST server is running."""
|
||||
return ollama_rest_provider.is_running(is_ollama_rest_enabled())
|
||||
|
||||
|
||||
def is_openai_running() -> bool:
|
||||
"""Check if the OpenAI-compatible endpoint is reachable."""
|
||||
return openai_provider.is_running(is_openai_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 a list of available models from LM Studio REST."""
|
||||
"""Retrieve text models from LM Studio REST."""
|
||||
return lmstudio_rest_provider.get_models(is_lmstudio_rest_enabled())
|
||||
|
||||
|
||||
def get_lmstudio_rest_vision_models() -> list[str]:
|
||||
"""Retrieve a list of available vision models from LM Studio REST."""
|
||||
"""Retrieve vision models from LM Studio REST."""
|
||||
return lmstudio_rest_provider.get_vision_models(is_lmstudio_rest_enabled())
|
||||
|
||||
|
||||
def get_ollama_rest_models() -> list[str]:
|
||||
"""Retrieve a list of available models from Ollama REST."""
|
||||
"""Retrieve text models from Ollama REST."""
|
||||
return ollama_rest_provider.get_models(is_ollama_rest_enabled())
|
||||
|
||||
|
||||
def get_ollama_rest_vision_models() -> list[str]:
|
||||
"""Retrieve a list of available vision models from Ollama REST."""
|
||||
"""Retrieve vision models from Ollama REST."""
|
||||
return ollama_rest_provider.get_vision_models(is_ollama_rest_enabled())
|
||||
|
||||
|
||||
def get_openai_models() -> list[str]:
|
||||
"""Retrieve a list of available models from the OpenAI-compatible endpoint."""
|
||||
"""Retrieve text models from OpenAI-compatible provider."""
|
||||
return openai_provider.get_models(is_openai_enabled())
|
||||
|
||||
|
||||
def get_openai_vision_models() -> list[str]:
|
||||
"""Retrieve a list of available vision models from the OpenAI-compatible endpoint."""
|
||||
"""Retrieve vision models from OpenAI-compatible provider."""
|
||||
return openai_provider.get_vision_models(is_openai_enabled())
|
||||
|
||||
|
||||
def get_ollama_tool_models() -> list[str]:
|
||||
"""Retrieve tool-capable models from Ollama SDK provider."""
|
||||
return ollama_provider.get_tool_models(OLLAMA_AVAILABLE, ollama_client, is_ollama_enabled())
|
||||
|
||||
|
||||
def get_ollama_reasoning_models() -> list[str]:
|
||||
"""Retrieve reasoning-capable models from Ollama SDK provider."""
|
||||
return ollama_provider.get_reasoning_models(OLLAMA_AVAILABLE, ollama_client, is_ollama_enabled())
|
||||
|
||||
|
||||
def get_lmstudio_tool_models() -> list[str]:
|
||||
"""Retrieve tool-capable models from LM Studio SDK provider."""
|
||||
return lmstudio_provider.get_tool_models(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled())
|
||||
"""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]:
|
||||
"""Retrieve reasoning-capable models from LM Studio SDK provider."""
|
||||
return lmstudio_provider.get_reasoning_models(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled())
|
||||
"""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())
|
||||
|
||||
|
||||
def get_lmstudio_rest_tool_models() -> list[str]:
|
||||
"""Retrieve tool-capable models from LM Studio REST provider."""
|
||||
"""Retrieve tool-capable models from LM Studio REST."""
|
||||
return lmstudio_rest_provider.get_tool_models(is_lmstudio_rest_enabled())
|
||||
|
||||
|
||||
def get_lmstudio_rest_reasoning_models() -> list[str]:
|
||||
"""Retrieve reasoning-capable models from LM Studio REST provider."""
|
||||
"""Retrieve reasoning-capable models from LM Studio REST."""
|
||||
return lmstudio_rest_provider.get_reasoning_models(is_lmstudio_rest_enabled())
|
||||
|
||||
|
||||
def get_ollama_rest_tool_models() -> list[str]:
|
||||
"""Retrieve tool-capable models from Ollama REST provider."""
|
||||
"""Retrieve tool-capable models from Ollama REST."""
|
||||
return ollama_rest_provider.get_tool_models(is_ollama_rest_enabled())
|
||||
|
||||
|
||||
def get_ollama_rest_reasoning_models() -> list[str]:
|
||||
"""Retrieve reasoning-capable models from Ollama REST provider."""
|
||||
"""Retrieve reasoning-capable models from Ollama REST."""
|
||||
return ollama_rest_provider.get_reasoning_models(is_ollama_rest_enabled())
|
||||
|
||||
|
||||
@@ -214,67 +151,98 @@ def get_openai_reasoning_models() -> list[str]:
|
||||
return openai_provider.get_reasoning_models(is_openai_enabled())
|
||||
|
||||
|
||||
def get_ollama_model_capabilities_map() -> dict[str, dict[str, object]]:
|
||||
"""Retrieve capabilities map for Ollama SDK models."""
|
||||
capability_map = ollama_provider.get_model_capabilities_map(OLLAMA_AVAILABLE, ollama_client, is_ollama_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_lmstudio_model_capabilities_map() -> dict[str, dict[str, object]]:
|
||||
"""Retrieve capabilities map for LM Studio SDK models."""
|
||||
capability_map = lmstudio_provider.get_model_capabilities_map(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled())
|
||||
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()}
|
||||
|
||||
|
||||
def get_lmstudio_rest_model_capabilities_map() -> dict[str, dict[str, object]]:
|
||||
"""Retrieve capabilities map for LM Studio REST models."""
|
||||
"""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()}
|
||||
|
||||
|
||||
def get_ollama_rest_model_capabilities_map() -> dict[str, dict[str, object]]:
|
||||
"""Retrieve capabilities map for Ollama REST models."""
|
||||
"""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()}
|
||||
|
||||
|
||||
def get_openai_model_capabilities_map() -> dict[str, dict[str, object]]:
|
||||
"""Retrieve capabilities map for OpenAI-compatible models."""
|
||||
"""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()}
|
||||
|
||||
|
||||
def lmstudio_generate_vision(model: str, prompt: str, keep_alive: int = 0, images=None, options=None) -> str:
|
||||
"""Generate a response from an LM Studio vision model."""
|
||||
ensure_lmstudio_initialized()
|
||||
return lmstudio_provider.generate_vision(
|
||||
LMSTUDIO_AVAILABLE,
|
||||
lms,
|
||||
is_lmstudio_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
images,
|
||||
options,
|
||||
)
|
||||
# ============================================================================
|
||||
# HEALTH
|
||||
# ============================================================================
|
||||
|
||||
def is_lmstudio_running() -> bool:
|
||||
"""Legacy alias: check LM Studio REST availability."""
|
||||
return lmstudio_rest_provider.is_running(_is_lmstudio_service_enabled())
|
||||
|
||||
|
||||
def lmstudio_generate(model: str, prompt: str, keep_alive: int = 0, options=None) -> str:
|
||||
"""Generate a response from an LM Studio model."""
|
||||
ensure_lmstudio_initialized()
|
||||
return lmstudio_provider.generate(
|
||||
LMSTUDIO_AVAILABLE,
|
||||
lms,
|
||||
is_lmstudio_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
options,
|
||||
)
|
||||
def is_ollama_running() -> bool:
|
||||
"""Legacy alias: check Ollama REST availability."""
|
||||
return ollama_rest_provider.is_running(_is_ollama_service_enabled())
|
||||
|
||||
|
||||
def is_lmstudio_rest_running() -> bool:
|
||||
"""Check LM Studio REST availability."""
|
||||
return lmstudio_rest_provider.is_running(is_lmstudio_rest_enabled())
|
||||
|
||||
|
||||
def is_ollama_rest_running() -> bool:
|
||||
"""Check Ollama REST availability."""
|
||||
return ollama_rest_provider.is_running(is_ollama_rest_enabled())
|
||||
|
||||
|
||||
def is_openai_running() -> bool:
|
||||
"""Check OpenAI-compatible provider availability."""
|
||||
return openai_provider.is_running(is_openai_enabled())
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 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
|
||||
|
||||
|
||||
def lmstudio_rest_generate(model: str, prompt: str, keep_alive: int = 0, options=None, system_prompt: str = '') -> str:
|
||||
"""Generate a response from an LM Studio REST model."""
|
||||
"""Generate text via LM Studio REST."""
|
||||
ensure_lmstudio_rest_initialized()
|
||||
return lmstudio_rest_provider.generate(
|
||||
is_lmstudio_rest_enabled(),
|
||||
@@ -285,25 +253,8 @@ def lmstudio_rest_generate(model: str, prompt: str, keep_alive: int = 0, options
|
||||
)
|
||||
|
||||
|
||||
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]:
|
||||
"""Generate a response from an LM Studio vision model and refine it with another model."""
|
||||
ensure_lmstudio_initialized()
|
||||
return lmstudio_provider.generate_vision_refine(
|
||||
LMSTUDIO_AVAILABLE,
|
||||
lms,
|
||||
is_lmstudio_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
images,
|
||||
options,
|
||||
refine_model,
|
||||
refine_prompt,
|
||||
refine_options,
|
||||
)
|
||||
|
||||
|
||||
def lmstudio_rest_generate_vision(model: str, prompt: str, keep_alive: int = 0, images=None, options=None, system_prompt: str = '') -> str:
|
||||
"""Generate a vision response from an LM Studio REST model."""
|
||||
"""Generate vision output via LM Studio REST."""
|
||||
ensure_lmstudio_rest_initialized()
|
||||
return lmstudio_rest_provider.generate_vision(
|
||||
is_lmstudio_rest_enabled(),
|
||||
@@ -315,8 +266,48 @@ def lmstudio_rest_generate_vision(model: str, prompt: str, keep_alive: int = 0,
|
||||
)
|
||||
|
||||
|
||||
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(
|
||||
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
|
||||
|
||||
|
||||
def ollama_rest_generate(model: str, prompt: str, options=None, system_prompt: str = '', keep_alive: str = '5m') -> str:
|
||||
"""Generate a response from an Ollama REST model."""
|
||||
"""Generate text via Ollama REST."""
|
||||
ensure_ollama_rest_initialized()
|
||||
return ollama_rest_provider.generate(
|
||||
is_ollama_rest_enabled(),
|
||||
@@ -329,7 +320,7 @@ 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 a vision response from an Ollama REST model."""
|
||||
"""Generate vision output via Ollama REST."""
|
||||
ensure_ollama_rest_initialized()
|
||||
return ollama_rest_provider.generate_vision(
|
||||
is_ollama_rest_enabled(),
|
||||
@@ -343,7 +334,7 @@ def ollama_rest_generate_vision(model: str, prompt: str, images=None, options=No
|
||||
|
||||
|
||||
def openai_generate(model: str, prompt: str, options=None, system_prompt: str = '') -> str:
|
||||
"""Generate a response from an OpenAI-compatible model."""
|
||||
"""Generate text via OpenAI-compatible provider."""
|
||||
ensure_openai_initialized()
|
||||
return openai_provider.generate(
|
||||
is_openai_enabled(),
|
||||
@@ -355,7 +346,7 @@ 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 a vision response from an OpenAI-compatible model."""
|
||||
"""Generate vision output via OpenAI-compatible provider."""
|
||||
ensure_openai_initialized()
|
||||
return openai_provider.generate_vision(
|
||||
is_openai_enabled(),
|
||||
@@ -368,71 +359,21 @@ def openai_generate_vision(model: str, prompt: str, images=None, options=None, s
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# STREAMING FUNCTIONS
|
||||
# GENERATION (STREAMING)
|
||||
# ============================================================================
|
||||
|
||||
def ollama_generate_stream(model: str, prompt: str, keep_alive: float = 0.0, options=None, system_prompt: str = ''):
|
||||
"""Generate a streaming response from an Ollama model."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.generate_stream(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama_client,
|
||||
is_ollama_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
options,
|
||||
system_prompt,
|
||||
)
|
||||
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 ollama_generate_vision_stream(model: str, prompt: str, keep_alive: float = 0.0, images=None, options=None, system_prompt: str = ''):
|
||||
"""Generate a streaming response from an Ollama vision model."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.generate_vision_stream(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama_client,
|
||||
is_ollama_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
images,
|
||||
options,
|
||||
system_prompt,
|
||||
)
|
||||
|
||||
|
||||
def lmstudio_generate_stream(model: str, prompt: str, keep_alive: int = 0, options=None):
|
||||
"""Generate a streaming response from an LM Studio model."""
|
||||
ensure_lmstudio_initialized()
|
||||
return lmstudio_provider.generate_stream(
|
||||
LMSTUDIO_AVAILABLE,
|
||||
lms,
|
||||
is_lmstudio_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
options,
|
||||
)
|
||||
|
||||
|
||||
def lmstudio_generate_vision_stream(model: str, prompt: str, keep_alive: int = 0, images=None, options=None):
|
||||
"""Generate a streaming response from an LM Studio vision model."""
|
||||
ensure_lmstudio_initialized()
|
||||
return lmstudio_provider.generate_vision_stream(
|
||||
LMSTUDIO_AVAILABLE,
|
||||
lms,
|
||||
is_lmstudio_enabled(),
|
||||
model,
|
||||
prompt,
|
||||
keep_alive,
|
||||
images,
|
||||
options,
|
||||
)
|
||||
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)
|
||||
|
||||
|
||||
def lmstudio_rest_generate_stream(model: str, prompt: str, keep_alive: int = 0, options=None, system_prompt: str = ''):
|
||||
"""Generate a streaming response from an LM Studio REST model."""
|
||||
"""Stream text via LM Studio REST."""
|
||||
ensure_lmstudio_rest_initialized()
|
||||
return lmstudio_rest_provider.generate_stream(
|
||||
is_lmstudio_rest_enabled(),
|
||||
@@ -444,7 +385,7 @@ 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 = ''):
|
||||
"""Generate a streaming vision response from an LM Studio REST model."""
|
||||
"""Stream vision output via LM Studio REST."""
|
||||
ensure_lmstudio_rest_initialized()
|
||||
return lmstudio_rest_provider.generate_vision_stream(
|
||||
is_lmstudio_rest_enabled(),
|
||||
@@ -456,8 +397,31 @@ def lmstudio_rest_generate_vision_stream(model: str, prompt: str, keep_alive: in
|
||||
)
|
||||
|
||||
|
||||
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(
|
||||
model,
|
||||
prompt,
|
||||
images=images,
|
||||
options=options,
|
||||
system_prompt=system_prompt,
|
||||
keep_alive=_normalize_ollama_keep_alive(keep_alive),
|
||||
)
|
||||
|
||||
|
||||
def ollama_rest_generate_stream(model: str, prompt: str, options=None, system_prompt: str = '', keep_alive: str = '5m'):
|
||||
"""Generate a streaming response from an Ollama REST model."""
|
||||
"""Stream text via Ollama REST."""
|
||||
ensure_ollama_rest_initialized()
|
||||
return ollama_rest_provider.generate_stream(
|
||||
is_ollama_rest_enabled(),
|
||||
@@ -470,7 +434,7 @@ 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'):
|
||||
"""Generate a streaming vision response from an Ollama REST model."""
|
||||
"""Stream vision output via Ollama REST."""
|
||||
ensure_ollama_rest_initialized()
|
||||
return ollama_rest_provider.generate_vision_stream(
|
||||
is_ollama_rest_enabled(),
|
||||
@@ -484,7 +448,7 @@ def ollama_rest_generate_vision_stream(model: str, prompt: str, images=None, opt
|
||||
|
||||
|
||||
def openai_generate_stream(model: str, prompt: str, options=None, system_prompt: str = ''):
|
||||
"""Generate a streaming response from an OpenAI-compatible model."""
|
||||
"""Stream text via OpenAI-compatible provider."""
|
||||
ensure_openai_initialized()
|
||||
return openai_provider.generate_stream(
|
||||
is_openai_enabled(),
|
||||
@@ -496,7 +460,7 @@ 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 = ''):
|
||||
"""Generate a streaming vision response from an OpenAI-compatible model."""
|
||||
"""Stream vision output via OpenAI-compatible provider."""
|
||||
ensure_openai_initialized()
|
||||
return openai_provider.generate_vision_stream(
|
||||
is_openai_enabled(),
|
||||
@@ -508,184 +472,124 @@ 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:
|
||||
"""Pre-warm an Ollama model so subsequent generate calls incur no load delay."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.preload_model(OLLAMA_AVAILABLE, ollama_client, is_ollama_enabled(), model, keep_alive)
|
||||
|
||||
|
||||
def ollama_generate_preloaded(model: str, prompt: str, keep_alive: float = 0.0, options=None, system_prompt: str = '') -> str:
|
||||
"""Generate from an Ollama model assumed to already be loaded."""
|
||||
ensure_ollama_initialized()
|
||||
return ollama_provider.generate_preloaded(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama_client,
|
||||
is_ollama_enabled(),
|
||||
"""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,
|
||||
prompt,
|
||||
keep_alive,
|
||||
options,
|
||||
system_prompt,
|
||||
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):
|
||||
"""Load an LM Studio model and return the model handle."""
|
||||
ensure_lmstudio_initialized()
|
||||
return lmstudio_provider.load_model(LMSTUDIO_AVAILABLE, lms, is_lmstudio_enabled(), model, keep_alive)
|
||||
"""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:
|
||||
"""Run text inference on an already-loaded LM Studio model handle."""
|
||||
return lmstudio_provider.generate_with_model(lms_model, lms, prompt, options)
|
||||
"""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:
|
||||
"""Run vision inference on an already-loaded LM Studio model handle."""
|
||||
return lmstudio_provider.generate_vision_with_model(lms_model, lms, prompt, images, options)
|
||||
"""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:
|
||||
"""Unload a previously loaded LM Studio model handle."""
|
||||
lmstudio_provider.unload_model(lms_model)
|
||||
"""Legacy alias: unload LM Studio REST model by name."""
|
||||
lmstudio_rest_provider.unload_model(_is_lmstudio_service_enabled(), str(lms_model))
|
||||
|
||||
|
||||
def lmstudio_rest_load_model(model: str, keep_alive: int = 0) -> bool:
|
||||
"""Ask LM Studio REST to load a model by name."""
|
||||
"""Load LM Studio REST model by name."""
|
||||
ensure_lmstudio_rest_initialized()
|
||||
return lmstudio_rest_provider.load_model(is_lmstudio_rest_enabled(), model, keep_alive)
|
||||
|
||||
|
||||
def lmstudio_rest_unload_model(model: str) -> bool:
|
||||
"""Ask LM Studio REST to unload a model by name."""
|
||||
"""Unload LM Studio REST model by name."""
|
||||
return lmstudio_rest_provider.unload_model(is_lmstudio_rest_enabled(), model)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# INITIALIZATION FUNCTIONS
|
||||
# INITIALIZATION
|
||||
# ============================================================================
|
||||
|
||||
def init_ollama():
|
||||
"""Initialize Ollama client."""
|
||||
global ollama_client, _ollama_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
ollama_custom_url = str(
|
||||
get_setting('ollama_custom_url', get_setting('custom_ollama_url', 'http://localhost:11434'))
|
||||
)
|
||||
|
||||
ollama_client, _ollama_initialized = init_ollama_client(
|
||||
OLLAMA_AVAILABLE,
|
||||
ollama,
|
||||
bool(get_setting('enable_ollama', False)),
|
||||
ollama_custom_url,
|
||||
)
|
||||
return _ollama_initialized
|
||||
def init_ollama() -> bool:
|
||||
"""Legacy alias: initialize Ollama REST state."""
|
||||
return init_ollama_rest()
|
||||
|
||||
|
||||
def init_lmstudio():
|
||||
"""Initialize LM Studio if available."""
|
||||
global _lmstudio_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
lmstudio_custom_url = str(
|
||||
get_setting('lmstudio_custom_url', get_setting('custom_lmstudio_url', ''))
|
||||
)
|
||||
|
||||
_lmstudio_initialized = init_lmstudio_client(
|
||||
LMSTUDIO_AVAILABLE,
|
||||
lms,
|
||||
bool(get_setting('enable_lmstudio', False)),
|
||||
lmstudio_custom_url,
|
||||
)
|
||||
return _lmstudio_initialized
|
||||
def init_lmstudio() -> bool:
|
||||
"""Legacy alias: initialize LM Studio REST state."""
|
||||
return init_lmstudio_rest()
|
||||
|
||||
|
||||
def init_lmstudio_rest():
|
||||
def init_lmstudio_rest() -> bool:
|
||||
"""Initialize LM Studio REST provider state."""
|
||||
global _lmstudio_rest_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
if not bool(get_setting('enable_lmstudio_rest', False)):
|
||||
_lmstudio_rest_initialized = False
|
||||
return False
|
||||
|
||||
_lmstudio_rest_initialized = lmstudio_rest_provider.is_running(True)
|
||||
if _lmstudio_rest_initialized:
|
||||
logger.info('LM Studio REST provider initialized.')
|
||||
else:
|
||||
logger.info('LM Studio REST provider is enabled but server is not reachable yet.')
|
||||
_lmstudio_rest_initialized = _init_lmstudio_rest(
|
||||
bool(get_setting('enable_lmstudio_rest', False))
|
||||
)
|
||||
return _lmstudio_rest_initialized
|
||||
|
||||
|
||||
def init_ollama_rest():
|
||||
def init_ollama_rest() -> bool:
|
||||
"""Initialize Ollama REST provider state."""
|
||||
global _ollama_rest_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
if not bool(get_setting('enable_ollama_rest', False)):
|
||||
_ollama_rest_initialized = False
|
||||
return False
|
||||
|
||||
_ollama_rest_initialized = ollama_rest_provider.is_running(True)
|
||||
if _ollama_rest_initialized:
|
||||
logger.info('Ollama REST provider initialized.')
|
||||
else:
|
||||
logger.info('Ollama REST provider is enabled but server is not reachable yet.')
|
||||
_ollama_rest_initialized = _init_ollama_rest(
|
||||
bool(get_setting('enable_ollama_rest', False))
|
||||
)
|
||||
return _ollama_rest_initialized
|
||||
|
||||
|
||||
def init_openai():
|
||||
def init_openai() -> bool:
|
||||
"""Initialize OpenAI provider state."""
|
||||
global _openai_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
if not bool(get_setting('enable_openai', False)):
|
||||
_openai_initialized = False
|
||||
return False
|
||||
|
||||
_openai_initialized = openai_provider.is_running(True)
|
||||
if _openai_initialized:
|
||||
logger.info('OpenAI provider initialized.')
|
||||
else:
|
||||
logger.info('OpenAI provider is enabled but endpoint is not reachable yet.')
|
||||
_openai_initialized = _init_openai_provider(
|
||||
bool(get_setting('enable_openai', False))
|
||||
)
|
||||
return _openai_initialized
|
||||
|
||||
|
||||
def init_llm():
|
||||
"""Initialize LLM clients."""
|
||||
init_ollama()
|
||||
init_lmstudio()
|
||||
def init_llm() -> None:
|
||||
"""Initialize all configured REST/OpenAI providers."""
|
||||
init_lmstudio_rest()
|
||||
init_ollama_rest()
|
||||
init_openai()
|
||||
logger.info('LLM clients initialized.')
|
||||
logger.info('LLM providers initialized.')
|
||||
|
||||
|
||||
def ensure_ollama_initialized():
|
||||
"""Ensure Ollama is initialized if enabled in settings and not already initialized."""
|
||||
global _ollama_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
if get_setting('enable_ollama', False) and not _ollama_initialized:
|
||||
logger.info('Ollama is enabled but not initialized, initializing now...')
|
||||
return init_ollama()
|
||||
return _ollama_initialized
|
||||
def ensure_ollama_initialized() -> bool:
|
||||
"""Legacy alias: ensure Ollama REST initialization."""
|
||||
return ensure_ollama_rest_initialized()
|
||||
|
||||
|
||||
def ensure_lmstudio_initialized():
|
||||
"""Ensure LM Studio is initialized if enabled in settings and not already initialized."""
|
||||
global _lmstudio_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
if get_setting('enable_lmstudio', False) and not _lmstudio_initialized:
|
||||
logger.info('LM Studio is enabled but not initialized, initializing now...')
|
||||
return init_lmstudio()
|
||||
return _lmstudio_initialized
|
||||
def ensure_lmstudio_initialized() -> bool:
|
||||
"""Legacy alias: ensure LM Studio REST initialization."""
|
||||
return ensure_lmstudio_rest_initialized()
|
||||
|
||||
|
||||
def ensure_lmstudio_rest_initialized():
|
||||
"""Ensure LM Studio REST is initialized if enabled in settings and not already initialized."""
|
||||
def ensure_lmstudio_rest_initialized() -> bool:
|
||||
"""Ensure LM Studio REST is initialized if enabled."""
|
||||
global _lmstudio_rest_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
@@ -695,8 +599,8 @@ def ensure_lmstudio_rest_initialized():
|
||||
return _lmstudio_rest_initialized
|
||||
|
||||
|
||||
def ensure_ollama_rest_initialized():
|
||||
"""Ensure Ollama REST is initialized if enabled in settings and not already initialized."""
|
||||
def ensure_ollama_rest_initialized() -> bool:
|
||||
"""Ensure Ollama REST is initialized if enabled."""
|
||||
global _ollama_rest_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
@@ -706,8 +610,8 @@ def ensure_ollama_rest_initialized():
|
||||
return _ollama_rest_initialized
|
||||
|
||||
|
||||
def ensure_openai_initialized():
|
||||
"""Ensure OpenAI provider is initialized if enabled in settings and not already initialized."""
|
||||
def ensure_openai_initialized() -> bool:
|
||||
"""Ensure OpenAI provider is initialized if enabled."""
|
||||
global _openai_initialized
|
||||
from ..settings import get_setting
|
||||
|
||||
@@ -717,22 +621,17 @@ def ensure_openai_initialized():
|
||||
return _openai_initialized
|
||||
|
||||
|
||||
def ensure_llm_initialized():
|
||||
def ensure_llm_initialized() -> bool:
|
||||
"""Ensure all enabled LLM services are initialized."""
|
||||
ollama_ok = ensure_ollama_initialized()
|
||||
lmstudio_ok = ensure_lmstudio_initialized()
|
||||
lmstudio_rest_ok = ensure_lmstudio_rest_initialized()
|
||||
ollama_rest_ok = ensure_ollama_rest_initialized()
|
||||
openai_ok = ensure_openai_initialized()
|
||||
return ollama_ok or lmstudio_ok or lmstudio_rest_ok or ollama_rest_ok or openai_ok
|
||||
return lmstudio_rest_ok or ollama_rest_ok or openai_ok
|
||||
|
||||
|
||||
def reset_llm_initialization_state() -> None:
|
||||
"""Reset initialization flags and local client references."""
|
||||
global _ollama_initialized, _lmstudio_initialized, _lmstudio_rest_initialized, _ollama_rest_initialized, _openai_initialized, ollama_client
|
||||
_ollama_initialized = False
|
||||
_lmstudio_initialized = False
|
||||
"""Reset initialization flags."""
|
||||
global _lmstudio_rest_initialized, _ollama_rest_initialized, _openai_initialized
|
||||
_lmstudio_rest_initialized = False
|
||||
_ollama_rest_initialized = False
|
||||
_openai_initialized = False
|
||||
ollama_client = None
|
||||
|
||||
+41
-18
@@ -10,6 +10,12 @@ This module provides a centralized way to manage settings with:
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from .config_manager import ConfigManager
|
||||
from .settings_crypto import (
|
||||
decrypt_sensitive_value,
|
||||
encrypt_sensitive_value,
|
||||
is_encrypted_value,
|
||||
is_sensitive_setting_key,
|
||||
)
|
||||
|
||||
from .logger import get_logger
|
||||
logger = get_logger('settings')
|
||||
@@ -17,16 +23,6 @@ logger = get_logger('settings')
|
||||
# Define the schema for all SageUtils settings
|
||||
SETTINGS_SCHEMA = {
|
||||
# LLM Integration Settings
|
||||
"enable_ollama": {
|
||||
"default": True,
|
||||
"type": bool,
|
||||
"description": "Enable Ollama LLM integration"
|
||||
},
|
||||
"enable_lmstudio": {
|
||||
"default": True,
|
||||
"type": bool,
|
||||
"description": "Enable LM Studio LLM integration"
|
||||
},
|
||||
"enable_lmstudio_rest": {
|
||||
"default": False,
|
||||
"type": bool,
|
||||
@@ -67,6 +63,11 @@ SETTINGS_SCHEMA = {
|
||||
"type": str,
|
||||
"description": "Custom URL for Ollama service (e.g., 'http://localhost:11434')"
|
||||
},
|
||||
"ollama_api_key": {
|
||||
"default": "",
|
||||
"type": str,
|
||||
"description": "API key for Ollama endpoint authentication (can also be set via OLLAMA_API_KEY env var)"
|
||||
},
|
||||
"lmstudio_use_custom_url": {
|
||||
"default": False,
|
||||
"type": bool,
|
||||
@@ -77,11 +78,16 @@ SETTINGS_SCHEMA = {
|
||||
"type": str,
|
||||
"description": "Custom URL for LM Studio service (e.g., 'http://localhost:1234')"
|
||||
},
|
||||
"default_llm_provider": {
|
||||
"default": "ollama",
|
||||
"lmstudio_api_token": {
|
||||
"default": "",
|
||||
"type": str,
|
||||
"description": "Default LLM provider to use for LLM sidebar and provider-switching LLM v3 nodes",
|
||||
"valid_values": ["ollama", "lmstudio", "lmstudio_rest", "ollama_rest", "openai", "native"]
|
||||
"description": "API token for LM Studio REST endpoint (can also be set via LMSTUDIO_API_TOKEN env var)"
|
||||
},
|
||||
"default_llm_provider": {
|
||||
"default": "lmstudio_rest",
|
||||
"type": str,
|
||||
"description": "Default LLM provider for the LLM sidebar and provider-switching LLM v3 nodes",
|
||||
"valid_values": ["lmstudio", "ollama", "lmstudio_rest", "ollama_rest", "openai", "native"]
|
||||
},
|
||||
"llm_raise_node_exceptions": {
|
||||
"default": False,
|
||||
@@ -195,11 +201,25 @@ class SageSettings:
|
||||
def load_and_validate(self) -> None:
|
||||
"""Load settings from config manager and validate against schema."""
|
||||
# Load current settings
|
||||
current_settings = self._config_manager.load() or {}
|
||||
raw_settings = self._config_manager.load() or {}
|
||||
|
||||
# Decrypt sensitive values for in-memory use.
|
||||
current_settings: Dict[str, Any] = {}
|
||||
settings_updated = False
|
||||
for key, raw_value in raw_settings.items():
|
||||
current_settings[key] = decrypt_sensitive_value(key, raw_value)
|
||||
# Migrate existing plaintext secrets to encrypted-at-rest values on next save.
|
||||
if (
|
||||
is_sensitive_setting_key(key)
|
||||
and isinstance(raw_value, str)
|
||||
and raw_value.strip()
|
||||
and not is_encrypted_value(raw_value)
|
||||
):
|
||||
settings_updated = True
|
||||
logger.info(f"Setting '{key}' will be migrated to encrypted storage.")
|
||||
|
||||
# Start with defaults and update with current values
|
||||
self._settings = {}
|
||||
settings_updated = False
|
||||
|
||||
for key, schema_entry in SETTINGS_SCHEMA.items():
|
||||
default_value = schema_entry["default"]
|
||||
@@ -221,7 +241,7 @@ class SageSettings:
|
||||
logger.info(f"Setting '{key}' added with default value: {default_value}")
|
||||
|
||||
# Remove any settings not in schema (cleanup old/deprecated settings)
|
||||
for key in current_settings:
|
||||
for key in raw_settings:
|
||||
if key not in SETTINGS_SCHEMA:
|
||||
logger.warning(f"Removing deprecated setting: '{key}'")
|
||||
settings_updated = True
|
||||
@@ -257,7 +277,10 @@ class SageSettings:
|
||||
def save(self) -> bool:
|
||||
"""Save current settings to file."""
|
||||
try:
|
||||
self._config_manager.data = self._settings.copy()
|
||||
persisted_settings: Dict[str, Any] = {}
|
||||
for key, value in self._settings.items():
|
||||
persisted_settings[key] = encrypt_sensitive_value(key, value)
|
||||
self._config_manager.data = persisted_settings
|
||||
return self._config_manager.save()
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save settings: {e}")
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Encryption helpers for sensitive SageUtils settings values."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .logger import get_logger
|
||||
from .path_manager import path_manager
|
||||
|
||||
logger = get_logger('settings.crypto')
|
||||
|
||||
try:
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
except ImportError: # pragma: no cover - handled at runtime when dependency is missing
|
||||
Fernet = None # type: ignore[assignment]
|
||||
InvalidToken = Exception # type: ignore[assignment]
|
||||
|
||||
|
||||
ENC_PREFIX = 'enc:v1:'
|
||||
KEY_FILENAME = '.settings_secret.key'
|
||||
SENSITIVE_SETTING_KEYS = {
|
||||
'openai_api_key',
|
||||
'ollama_api_key',
|
||||
'lmstudio_api_token',
|
||||
}
|
||||
|
||||
|
||||
def is_sensitive_setting_key(key: str) -> bool:
|
||||
return key in SENSITIVE_SETTING_KEYS
|
||||
|
||||
|
||||
def is_encrypted_value(value: Any) -> bool:
|
||||
return isinstance(value, str) and value.startswith(ENC_PREFIX)
|
||||
|
||||
|
||||
def _get_or_create_key() -> bytes:
|
||||
if Fernet is None:
|
||||
raise RuntimeError(
|
||||
'cryptography is required for encrypted settings. Please install dependencies from requirements.txt.'
|
||||
)
|
||||
|
||||
key_path: Path = path_manager.get_user_file_path(KEY_FILENAME)
|
||||
if key_path.is_file():
|
||||
key = key_path.read_bytes().strip()
|
||||
# Validate key format early.
|
||||
Fernet(key)
|
||||
return key
|
||||
|
||||
key = Fernet.generate_key()
|
||||
key_path.write_bytes(key + b'\n')
|
||||
try:
|
||||
os.chmod(key_path, 0o600)
|
||||
except OSError:
|
||||
# Best effort: chmod may not be available/allowed on some platforms.
|
||||
pass
|
||||
logger.info('Created local settings encryption key.')
|
||||
return key
|
||||
|
||||
|
||||
def encrypt_sensitive_value(key: str, value: Any) -> Any:
|
||||
if not is_sensitive_setting_key(key):
|
||||
return value
|
||||
if value is None:
|
||||
return ''
|
||||
|
||||
plain_value = str(value)
|
||||
if not plain_value:
|
||||
return ''
|
||||
if is_encrypted_value(plain_value):
|
||||
return plain_value
|
||||
|
||||
fernet = Fernet(_get_or_create_key())
|
||||
token = fernet.encrypt(plain_value.encode('utf-8')).decode('utf-8')
|
||||
return f'{ENC_PREFIX}{token}'
|
||||
|
||||
|
||||
def decrypt_sensitive_value(key: str, value: Any) -> Any:
|
||||
if not is_sensitive_setting_key(key):
|
||||
return value
|
||||
if not isinstance(value, str) or not value:
|
||||
return value
|
||||
if not is_encrypted_value(value):
|
||||
return value
|
||||
|
||||
token = value[len(ENC_PREFIX):]
|
||||
try:
|
||||
fernet = Fernet(_get_or_create_key())
|
||||
return fernet.decrypt(token.encode('utf-8')).decode('utf-8')
|
||||
except (InvalidToken, ValueError, RuntimeError) as e:
|
||||
logger.warning(f"Unable to decrypt setting '{key}': {e}. Using empty value.")
|
||||
return ''
|
||||
Reference in New Issue
Block a user