Removed original ollama and lm studio backend in favor of rest. Got OpenAI working.

This commit is contained in:
Shanoah Alkire
2026-05-10 08:23:10 -07:00
parent bf9e23b584
commit 2aa82db665
28 changed files with 955 additions and 2491 deletions
+118 -17
View File
@@ -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;
}
+3 -3
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+2 -2
View File
@@ -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';
}
}
+1 -10
View File
@@ -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;
+2 -4
View File
@@ -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)' }
],
+2 -4
View File
@@ -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)' }
],
-9
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 -7
View File
@@ -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',
-613
View File
@@ -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),
}
+6 -2
View File
@@ -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)
-532
View File
@@ -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),
}
+10 -6
View File
@@ -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
+7
View File
@@ -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]] = {
+4 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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}")
+93
View File
@@ -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 ''