From 2aa82db665d23ba18d1cfc14564fffbbfa8d3776 Mon Sep 17 00:00:00 2001 From: Shanoah Alkire Date: Sun, 10 May 2026 08:23:10 -0700 Subject: [PATCH] Removed original ollama and lm studio backend in favor of rest. Got OpenAI working. --- js/dialogs/settingsDialog.js | 135 +++- js/llm/llmPresets.js | 6 +- js/llm/llmProviders.js | 72 +-- js/settingsUI.js | 74 ++- js/sidebar/llmTab.js | 26 +- js/sidebar/llmTab/llmEventHandlers.js | 4 +- js/sidebar/llmTab/llmGenerationHandler.js | 11 +- js/sidebar/llmTab/llmModelSelection.js | 6 +- js/sidebar/llmTab/llmPresetDialogs.js | 6 +- js/sidebar/llmTab/llmStyles.js | 9 - nodes/llm_v3.py | 221 +------ nodes/settings.py | 15 +- requirements.txt | 7 +- routes/llm_routes.py | 611 +++--------------- server_routes.py | 7 +- utils/llm/init.py | 112 ++-- utils/llm/providers/__init__.py | 8 +- utils/llm/providers/lmstudio_client.py | 613 ------------------ utils/llm/providers/lmstudio_rest_client.py | 8 +- utils/llm/providers/ollama_client.py | 532 ---------------- utils/llm/providers/ollama_rest_client.py | 16 +- utils/llm/providers/openai_client.py | 7 + utils/llm/providers/settings.py | 8 +- utils/llm/rest.py | 83 ++- utils/llm/routes_helpers.py | 48 +- utils/llm/service.py | 649 +++++++++----------- utils/settings.py | 59 +- utils/settings_crypto.py | 93 +++ 28 files changed, 955 insertions(+), 2491 deletions(-) delete mode 100644 utils/llm/providers/lmstudio_client.py delete mode 100644 utils/llm/providers/ollama_client.py create mode 100644 utils/settings_crypto.py diff --git a/js/dialogs/settingsDialog.js b/js/dialogs/settingsDialog.js index ad19dbe..0bc6e2c 100644 --- a/js/dialogs/settingsDialog.js +++ b/js/dialogs/settingsDialog.js @@ -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; } diff --git a/js/llm/llmPresets.js b/js/llm/llmPresets.js index de5c340..0d5bc1f 100644 --- a/js/llm/llmPresets.js +++ b/js/llm/llmPresets.js @@ -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', diff --git a/js/llm/llmProviders.js b/js/llm/llmProviders.js index 0d375a2..1dff641 100644 --- a/js/llm/llmProviders.js +++ b/js/llm/llmProviders.js @@ -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; diff --git a/js/settingsUI.js b/js/settingsUI.js index d927226..4dc6df9 100644 --- a/js/settingsUI.js +++ b/js/settingsUI.js @@ -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); + } } ] }); diff --git a/js/sidebar/llmTab.js b/js/sidebar/llmTab.js index aa756e2..a34e8c5 100644 --- a/js/sidebar/llmTab.js +++ b/js/sidebar/llmTab.js @@ -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} Default provider ('ollama', 'lmstudio', 'lmstudio_rest', or 'native') + * @returns {Promise} 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 diff --git a/js/sidebar/llmTab/llmEventHandlers.js b/js/sidebar/llmTab/llmEventHandlers.js index f9f5f8c..2a349a5 100644 --- a/js/sidebar/llmTab/llmEventHandlers.js +++ b/js/sidebar/llmTab/llmEventHandlers.js @@ -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'; } } diff --git a/js/sidebar/llmTab/llmGenerationHandler.js b/js/sidebar/llmTab/llmGenerationHandler.js index ca0ed9e..6d4cc99 100644 --- a/js/sidebar/llmTab/llmGenerationHandler.js +++ b/js/sidebar/llmTab/llmGenerationHandler.js @@ -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; diff --git a/js/sidebar/llmTab/llmModelSelection.js b/js/sidebar/llmTab/llmModelSelection.js index fabef70..bc010ed 100644 --- a/js/sidebar/llmTab/llmModelSelection.js +++ b/js/sidebar/llmTab/llmModelSelection.js @@ -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)' } ], diff --git a/js/sidebar/llmTab/llmPresetDialogs.js b/js/sidebar/llmTab/llmPresetDialogs.js index b105abc..6a4b790 100644 --- a/js/sidebar/llmTab/llmPresetDialogs.js +++ b/js/sidebar/llmTab/llmPresetDialogs.js @@ -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)' } ], diff --git a/js/sidebar/llmTab/llmStyles.js b/js/sidebar/llmTab/llmStyles.js index e3d9a85..8531ae2 100644 --- a/js/sidebar/llmTab/llmStyles.js +++ b/js/sidebar/llmTab/llmStyles.js @@ -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; diff --git a/nodes/llm_v3.py b/nodes/llm_v3.py index db9ada8..0a9e07e 100644 --- a/nodes/llm_v3.py +++ b/nodes/llm_v3.py @@ -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 \ No newline at end of file +] \ No newline at end of file diff --git a/nodes/settings.py b/nodes/settings.py index 1e8ad90..1030395 100644 --- a/nodes/settings.py +++ b/nodes/settings.py @@ -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(): diff --git a/requirements.txt b/requirements.txt index 51b8667..f93eb7c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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 diff --git a/routes/llm_routes.py b/routes/llm_routes.py index b299b9c..65040de 100644 --- a/routes/llm_routes.py +++ b/routes/llm_routes.py @@ -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', diff --git a/server_routes.py b/server_routes.py index a6c493f..b1ccef0 100644 --- a/server_routes.py +++ b/server_routes.py @@ -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: diff --git a/utils/llm/init.py b/utils/llm/init.py index 7cb3460..3019f42 100644 --- a/utils/llm/init.py +++ b/utils/llm/init.py @@ -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) diff --git a/utils/llm/providers/__init__.py b/utils/llm/providers/__init__.py index 5ba6996..e4b1b1d 100644 --- a/utils/llm/providers/__init__.py +++ b/utils/llm/providers/__init__.py @@ -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', diff --git a/utils/llm/providers/lmstudio_client.py b/utils/llm/providers/lmstudio_client.py deleted file mode 100644 index 4c71f07..0000000 --- a/utils/llm/providers/lmstudio_client.py +++ /dev/null @@ -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), - } diff --git a/utils/llm/providers/lmstudio_rest_client.py b/utils/llm/providers/lmstudio_rest_client.py index 1a75886..4740486 100644 --- a/utils/llm/providers/lmstudio_rest_client.py +++ b/utils/llm/providers/lmstudio_rest_client.py @@ -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) diff --git a/utils/llm/providers/ollama_client.py b/utils/llm/providers/ollama_client.py deleted file mode 100644 index ce6ce57..0000000 --- a/utils/llm/providers/ollama_client.py +++ /dev/null @@ -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 ('' 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), - } diff --git a/utils/llm/providers/ollama_rest_client.py b/utils/llm/providers/ollama_rest_client.py index ff6101c..6e29b75 100644 --- a/utils/llm/providers/ollama_rest_client.py +++ b/utils/llm/providers/ollama_rest_client.py @@ -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 diff --git a/utils/llm/providers/openai_client.py b/utils/llm/providers/openai_client.py index f9723b7..e207a33 100644 --- a/utils/llm/providers/openai_client.py +++ b/utils/llm/providers/openai_client.py @@ -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]] = { diff --git a/utils/llm/providers/settings.py b/utils/llm/providers/settings.py index b4be03b..a712679 100644 --- a/utils/llm/providers/settings.py +++ b/utils/llm/providers/settings.py @@ -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: diff --git a/utils/llm/rest.py b/utils/llm/rest.py index afdfdfb..89075a4 100644 --- a/utils/llm/rest.py +++ b/utils/llm/rest.py @@ -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]: diff --git a/utils/llm/routes_helpers.py b/utils/llm/routes_helpers.py index a7b9aa8..7361360 100644 --- a/utils/llm/routes_helpers.py +++ b/utils/llm/routes_helpers.py @@ -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() diff --git a/utils/llm/service.py b/utils/llm/service.py index ad7c83a..a4ed58e 100644 --- a/utils/llm/service.py +++ b/utils/llm/service.py @@ -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 diff --git a/utils/settings.py b/utils/settings.py index a7354a3..997a8e7 100644 --- a/utils/settings.py +++ b/utils/settings.py @@ -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}") diff --git a/utils/settings_crypto.py b/utils/settings_crypto.py new file mode 100644 index 0000000..e5852b9 --- /dev/null +++ b/utils/settings_crypto.py @@ -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 ''