diff --git a/orangecontrib/text/tests/test_documentembedder.py b/orangecontrib/text/tests/test_documentembedder.py
index ae2944b1a..8bcb8c5ca 100644
--- a/orangecontrib/text/tests/test_documentembedder.py
+++ b/orangecontrib/text/tests/test_documentembedder.py
@@ -1,11 +1,12 @@
import unittest
-from unittest.mock import patch, ANY
+from unittest.mock import patch, ANY, MagicMock
import asyncio
-from Orange.misc.utils.embedder_utils import EmbedderCache
from numpy.testing import assert_array_equal
-from orangecontrib.text.vectorization.document_embedder import DocumentEmbedder
+from Orange.misc.utils.embedder_utils import EmbedderCache
+
+from orangecontrib.text.vectorization.document_embedder import DocumentEmbedder, OAIDocumentEmbedder
from orangecontrib.text import Corpus
PATCH_METHOD = 'httpx.AsyncClient.post'
@@ -161,5 +162,297 @@ def test_set_language(self, m):
)
+class TestOAIDocumentEmbedder(unittest.TestCase):
+ """Test OAIDocumentEmbedder"""
+
+ def setUp(self):
+ self.corpus = Corpus.from_file('deerwester')
+ self.base_url = "https://api.openai.com/v1"
+ self.api_key = "test-api-key"
+ self.model = "text-embedding-ada-002"
+ self.embedder = OAIDocumentEmbedder(
+ base_url=self.base_url,
+ api_key=self.api_key,
+ model=self.model
+ )
+ self.embedder._cache._cache_dict.clear()
+
+ def _make_mock_embedding(self, embedding_list):
+ """Create a mock embedding response object."""
+ mock_embedding = MagicMock()
+ mock_embedding.embedding = embedding_list
+ return mock_embedding
+
+ def _make_mock_response(self, embeddings_list):
+ """Create a mock API response with a list of embeddings."""
+ mock_response = MagicMock()
+ mock_response.data = [self._make_mock_embedding(e) for e in embeddings_list]
+ return mock_response
+
+ @patch('openai.OpenAI')
+ def test_init(self, mock_openai):
+ """Test OAIDocumentEmbedder initialization."""
+ embedder = OAIDocumentEmbedder(
+ base_url="https://api.example.com",
+ api_key="key123",
+ model="gpt-4"
+ )
+ self.assertEqual(embedder.base_url, "https://api.example.com")
+ self.assertEqual(embedder.api_key, "key123")
+ self.assertEqual(embedder.model, "gpt-4")
+
+ @patch('openai.OpenAI')
+ def test_with_empty_corpus(self, mock_openai):
+ """Test transform with an empty corpus."""
+ empty_corpus = self.corpus[:0]
+ result, skipped = self.embedder.transform(empty_corpus)
+ # Empty corpus should return None for the first element of the tuple
+ self.assertEqual(len(result), 0)
+ self.assertIsNone(skipped)
+ # No API calls should be made
+ mock_openai.return_value.embeddings.create.assert_not_called()
+ # Cache should remain empty
+ self.assertEqual(len(self.embedder._cache._cache_dict), 0)
+
+ @patch('openai.OpenAI')
+ def test_success_single_document(self, mock_openai):
+ """Test successful embedding of a single document."""
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.1, 0.2, 0.3]])
+ mock_client.embeddings.create.return_value = mock_response
+
+ result, skipped = self.embedder.transform(self.corpus[[0]])
+
+ # Check result is a Corpus
+ self.assertIsNotNone(result)
+ self.assertIsNone(skipped)
+ # Check embedding values
+ assert_array_equal(result.X, [[0.1, 0.2, 0.3]])
+ # Check cache was populated
+ self.assertEqual(len(self.embedder._cache._cache_dict), 1)
+ # Verify API was called with correct parameters
+ mock_client.embeddings.create.assert_called_once()
+ call_args = mock_client.embeddings.create.call_args
+ self.assertEqual(call_args.kwargs['model'], self.model)
+ self.assertEqual(call_args.kwargs['encoding_format'], "float")
+ self.assertIn(self.corpus.documents[0], call_args.kwargs['input'])
+
+ @patch('openai.OpenAI')
+ def test_success_multiple_documents(self, mock_openai):
+ """Test successful embedding of multiple documents."""
+ mock_client = mock_openai.return_value
+ embeddings = [
+ [0.1, 0.2],
+ [0.3, 0.4],
+ [0.5, 0.6]
+ ]
+ mock_response = self._make_mock_response(embeddings)
+ mock_client.embeddings.create.return_value = mock_response
+
+ result, skipped = self.embedder.transform(self.corpus[[0, 1, 2]])
+
+ self.assertIsNotNone(result)
+ self.assertIsNone(skipped)
+ assert_array_equal(result.X, embeddings)
+ self.assertEqual(len(self.embedder._cache._cache_dict), 3)
+
+ @patch('openai.OpenAI')
+ def test_success_shapes(self, mock_openai):
+ """Test that output shapes are correct."""
+ corpus = self.corpus[:5]
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.1, 0.2, 0.3]] * 5)
+ mock_client.embeddings.create.return_value = mock_response
+
+ result, skipped = self.embedder.transform(corpus)
+
+ self.assertEqual(result.X.shape, (len(corpus), 3))
+ # Check that new features were added
+ self.assertEqual(len(result.domain.variables),
+ len(self.corpus.domain.variables) + 3)
+ # Verify feature names
+ feature_names = [v.name for v in result.domain.attributes]
+ self.assertIn("Dim1", feature_names)
+ self.assertIn("Dim2", feature_names)
+ self.assertIn("Dim3", feature_names)
+
+ @patch('openai.OpenAI')
+ def test_persistent_caching(self, mock_openai):
+ """Test that cache persists across embedder instances."""
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.5, 0.6, 0.7]])
+ mock_client.embeddings.create.return_value = mock_response
+
+ # First transform - cache should be empty
+ self.assertEqual(len(self.embedder._cache._cache_dict), 0)
+ self.embedder.transform(self.corpus[[0]])
+ self.assertEqual(len(self.embedder._cache._cache_dict), 1)
+
+ # Create a new embedder instance - cache should still have the data
+ new_embedder = OAIDocumentEmbedder(
+ base_url=self.base_url,
+ api_key=self.api_key,
+ model=self.model
+ )
+ self.assertEqual(len(new_embedder._cache._cache_dict), 1)
+
+ @patch('openai.OpenAI')
+ def test_cache_avoids_duplicate_api_calls(self, mock_openai):
+ """Test that cached documents don't trigger additional API calls."""
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.1, 0.2, 0.3]])
+ mock_client.embeddings.create.return_value = mock_response
+
+ # First transform - should call API
+ self.embedder.transform(self.corpus[[0]])
+ self.assertEqual(mock_client.embeddings.create.call_count, 1)
+
+ # Second transform with same document - should NOT call API again
+ self.embedder.transform(self.corpus[[0]])
+ self.assertEqual(mock_client.embeddings.create.call_count, 1)
+
+ @patch('openai.OpenAI')
+ def test_cache_partial_hits(self, mock_openai):
+ """Test that only uncached documents trigger API calls."""
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.1, 0.2, 0.3]])
+ mock_client.embeddings.create.return_value = mock_response
+
+ # Embed first document
+ self.embedder.transform(self.corpus[[0]])
+ self.assertEqual(mock_client.embeddings.create.call_count, 1)
+
+ # Embed second document - should only embed the new one
+ self.embedder.transform(self.corpus[[1]])
+ self.assertEqual(mock_client.embeddings.create.call_count, 2)
+
+ # Embed both again - should not call API at all
+ self.embedder.transform(self.corpus[[0, 1]])
+ self.assertEqual(mock_client.embeddings.create.call_count, 2)
+
+ @patch('openai.OpenAI')
+ def test_different_models_different_caches(self, mock_openai):
+ """Test that different models use different caches."""
+ embedder1 = OAIDocumentEmbedder(
+ base_url=self.base_url,
+ api_key=self.api_key,
+ model="model-a"
+ )
+ embedder2 = OAIDocumentEmbedder(
+ base_url=self.base_url,
+ api_key=self.api_key,
+ model="model-b"
+ )
+
+ self.assertNotEqual(embedder1._cache._cache_file_path, embedder2._cache._cache_file_path)
+
+ @patch('openai.OpenAI')
+ def test_different_base_urls_different_caches(self, mock_openai):
+ """Test that different base URLs use different caches."""
+ embedder1 = OAIDocumentEmbedder(
+ base_url="https://api.openai.com/v1",
+ api_key=self.api_key,
+ model="text-embedding-ada-002"
+ )
+ embedder2 = OAIDocumentEmbedder(
+ base_url="https://api.anthropic.com/v1",
+ api_key=self.api_key,
+ model="claude-embed"
+ )
+
+ self.assertNotEqual(embedder1._cache._cache_file_path, embedder2._cache._cache_file_path)
+
+ @patch('openai.OpenAI')
+ def test_progress_callback(self, mock_openai):
+ """Test that progress callback is called during embedding."""
+ mock_client = mock_openai.return_value
+ embeddings = [[0.1, 0.2] for _ in range(5)]
+ mock_response = self._make_mock_response(embeddings)
+ mock_client.embeddings.create.return_value = mock_response
+
+ callback = MagicMock()
+ self.embedder.transform(self.corpus[:5], callback=callback)
+
+ # Progress callback should have been called
+ callback.assert_called()
+ # The last call should indicate completion
+ last_call = callback.call_args_list[-1]
+ self.assertEqual(last_call[0][0], 1.0)
+
+ @patch('openai.OpenAI')
+ def test_api_error_propagates(self, mock_openai):
+ """Test that API errors are propagated."""
+ mock_client = mock_openai.return_value
+ mock_client.embeddings.create.side_effect = Exception("API Error")
+
+ with self.assertRaises(Exception):
+ self.embedder.transform(self.corpus[[0]])
+
+ @patch('openai.OpenAI')
+ def test_response_as_list(self, mock_openai):
+ """Test handling of response when it's already a list (not an object)."""
+ mock_client = mock_openai.return_value
+ # Simulate response being a list directly
+ mock_embedding = self._make_mock_embedding([0.1, 0.2, 0.3])
+ mock_client.embeddings.create.return_value = [mock_embedding]
+
+ result, skipped = self.embedder.transform(self.corpus[[0]])
+
+ self.assertIsNotNone(result)
+ self.assertIsNone(skipped)
+ assert_array_equal(result.X, [[0.1, 0.2, 0.3]])
+
+ @patch('openai.OpenAI')
+ def test_multidimensional_embedding_flattened(self, mock_openai):
+ """Test that multidimensional embeddings are flattened."""
+ mock_client = mock_openai.return_value
+ # Embedding with more than 1 dimension (e.g., 2D array)
+ mock_embedding = self._make_mock_embedding([[0.1, 0.2], [0.3, 0.4]])
+ mock_response = self._make_mock_response([[0.1, 0.2], [0.3, 0.4]])
+
+ # Override to return a 2D embedding for the first document
+ mock_embedding_2d = MagicMock()
+ mock_embedding_2d.embedding = [[0.1, 0.2], [0.3, 0.4]]
+ mock_response_2d = MagicMock()
+ mock_response_2d.data = [mock_embedding_2d]
+ mock_client.embeddings.create.return_value = mock_response_2d
+
+ result, skipped = self.embedder.transform(self.corpus[[0]])
+
+ # Should be flattened to 1D
+ self.assertEqual(result.X.shape[1], 4)
+ assert_array_equal(result.X[0], [0.1, 0.2, 0.3, 0.4])
+
+ @patch('openai.OpenAI')
+ def test_corpus_extend_attributes(self, mock_openai):
+ """Test that corpus is extended with correct feature attributes."""
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.1, 0.2]])
+ mock_client.embeddings.create.return_value = mock_response
+
+ result, skipped = self.embedder.transform(self.corpus[[0]])
+
+ # Check that embedding features have correct attributes
+ for var in result.domain.attributes:
+ if var.name.startswith("Dim"):
+ self.assertTrue(var.attributes.get("embedding-feature", False))
+ self.assertTrue(var.attributes.get("hidden", False))
+
+ @patch('openai.OpenAI')
+ def test_transform_returns_none_skipped(self, mock_openai):
+ """Test that skipped corpus is always None for OAIDocumentEmbedder."""
+ mock_client = mock_openai.return_value
+ mock_response = self._make_mock_response([[0.1, 0.2]])
+ mock_client.embeddings.create.return_value = mock_response
+
+ result, skipped = self.embedder.transform(self.corpus[[0]])
+
+ self.assertIsNone(skipped)
+
+
+if __name__ == "__main__":
+ unittest.main()
+
+
if __name__ == "__main__":
unittest.main()
diff --git a/orangecontrib/text/vectorization/document_embedder.py b/orangecontrib/text/vectorization/document_embedder.py
index ae1fe7893..7d5e4c9df 100644
--- a/orangecontrib/text/vectorization/document_embedder.py
+++ b/orangecontrib/text/vectorization/document_embedder.py
@@ -6,9 +6,12 @@
import sys
import warnings
import zlib
-from typing import Any, Optional, Tuple
+import re
+from typing import Any, Optional, Tuple, Callable
import numpy as np
+import openai
+
from Orange.misc.server_embedder import ServerEmbedderCommunicator
from Orange.misc.utils.embedder_utils import EmbedderCache
from Orange.util import dummy_callback
@@ -164,6 +167,140 @@ async def _encode_data_instance(self, data_instance: Any) -> Optional[bytes]:
return json_string.encode('utf-8', 'replace')
+def url_to_safe_filename(url: str) -> str:
+ """
+ Convert an URL into a safe, cross-platform single filesystem filename.
+ Args:
+ url: The input URL string
+
+ Returns:
+ A sanitized, valid single-file filename.
+ """
+ if not url or not url.strip():
+ raise ValueError("'url' cannot be empty")
+
+ # Replace all Windows/POSIX invalid characters and control chars with underscore
+ safe = re.sub(r'[<>:"/\\|?*]', '_', url)
+ safe = re.sub(r'[\x00-\x1f]', '_', safe)
+ # Handle Windows reserved names (CON, PRN, etc.)
+ safe = re.sub(
+ r'^(CON|PRN|AUX|NUL|COM[1-9]|LPT[1-3])$', r'_\1', safe, flags=re.IGNORECASE
+ )
+ return safe
+
+
+class OAIDocumentEmbedder(BaseVectorizer):
+ def __init__(self, base_url, api_key, model):
+ self.base_url = base_url
+ self.api_key = api_key
+ self.model = model
+ cache_name = url_to_safe_filename(f"{base_url}_{model}")
+ self._cache = EmbedderCache(cache_name)
+
+ def _transform(self, corpus, source_dict, callback=dummy_callback):
+ texts = list(corpus.documents)
+ results = [None] * len(texts)
+ # Collect all cached results
+ cache = self._cache
+ query = []
+ indices = []
+
+ for i, txt in enumerate(texts):
+ r = cache.get_cached_result_or_none(cache.md5_hash(txt.encode("utf-8")))
+ if r is not None:
+ results[i] = r
+ else:
+ query.append(txt)
+ indices.append(i)
+
+ callback(0.0)
+ embs = openai_get_embeddings(
+ query, self.api_key, self.base_url, self.model,
+ progress_callback=lambda a, b: callback(a/b)
+ )
+ embs = embs.tolist()
+ # Update cache and results list
+ for i, r, txt in zip(indices, embs, query):
+ cache.add(cache.md5_hash(txt.encode("utf-8"),), r)
+ results[i] = r
+ cache.persist_cache()
+ embs = np.array(results)
+ if results:
+ dim = embs.shape[1]
+ new_corpus = corpus.extend_attributes(
+ embs,
+ feature_names=["Dim{}".format(i + 1) for i in range(dim)],
+ var_attrs={
+ "embedding-feature": True,
+ "hidden": True,
+ }
+ )
+ else:
+ new_corpus = corpus
+ return new_corpus, None
+
+
+def openai_get_embeddings(
+ texts: list[str],
+ api_key: str,
+ base_url: Optional[str] = None,
+ model: str = "gpt-4o",
+ batch_size: int = 20,
+ progress_callback: Optional[Callable[[int, int], None]] = None,
+) -> np.ndarray:
+ """Generate embeddings for texts using an OpenAI-compatible API.
+
+ Processed in concurrent batches, each containing up to ``batch_size``
+ texts.
+
+ Args:
+ texts: List of texts
+ api_key: OpenAI-compatible API key.
+ base_url: Base URL for the OpenAI-compatible API. If None, uses the
+ default OpenAI endpoint.
+ model: Model name to use for embeddings. Defaults to "gpt-4o".
+ batch_size: Maximum number of concurrent API requests. Defaults to 10.
+ progress_callback: Optional callable invoked as
+ ``callback(completed_count, total_count)`` after each
+ batch finishes. Use ``None`` to skip progress reporting.
+
+ Returns:
+ Array of embedding vectors, one per query text. The order matches
+ the order of the input texts list.
+
+ Raises:
+ openai.OpenAIError: If the API call fails.
+ """
+ client = openai.OpenAI(api_key=api_key, base_url=base_url)
+ embeddings = [[]] * len(texts)
+
+ total = len(texts)
+ batch_number = 0
+
+ for start in range(0, total, batch_size):
+ batch_number += 1
+ batch_end = min(start + batch_size, total)
+ batch_indices = range(start, batch_end)
+ batch = texts[start:batch_end]
+ resp = client.embeddings.create(
+ input=batch, model=model, encoding_format="float",
+ )
+ if isinstance(resp, list):
+ resp = resp
+ else:
+ resp = resp.data
+ for i, r in zip(batch_indices, resp):
+ emb = np.array(r.embedding)
+ if emb.ndim > 1:
+ emb = emb.flatten()
+ embeddings[i] = emb
+ # Notify progress callback: (completed, total, batch_number)
+ if progress_callback is not None:
+ completed = min(batch_number * batch_size, total)
+ progress_callback(completed, total)
+ return np.array(embeddings)
+
+
if __name__ == '__main__':
with DocumentEmbedder(language='en', aggregator='Max') as embedder:
embedder.clear_cache()
diff --git a/orangecontrib/text/widgets/icons/hide-password.svg b/orangecontrib/text/widgets/icons/hide-password.svg
new file mode 100644
index 000000000..ebf67c0c3
--- /dev/null
+++ b/orangecontrib/text/widgets/icons/hide-password.svg
@@ -0,0 +1,10 @@
+
diff --git a/orangecontrib/text/widgets/icons/show-password.svg b/orangecontrib/text/widgets/icons/show-password.svg
new file mode 100644
index 000000000..890dcc777
--- /dev/null
+++ b/orangecontrib/text/widgets/icons/show-password.svg
@@ -0,0 +1,11 @@
+
diff --git a/orangecontrib/text/widgets/owdocumentembedding.py b/orangecontrib/text/widgets/owdocumentembedding.py
index 5445e3db8..140513faf 100644
--- a/orangecontrib/text/widgets/owdocumentembedding.py
+++ b/orangecontrib/text/widgets/owdocumentembedding.py
@@ -1,11 +1,20 @@
+import json
+import enum
+import os
from typing import Dict, Optional, Any
-from AnyQt.QtCore import Qt
+import openai
+
+from AnyQt.QtCore import Qt, QSettings
from AnyQt.QtWidgets import QVBoxLayout, QPushButton, QStyle
+
from Orange.misc.utils.embedder_utils import EmbeddingConnectionError
-from Orange.widgets import gui
+from Orange.widgets import gui, settings
from Orange.widgets.settings import Setting
+from Orange.widgets.utils import qname
+from Orange.widgets.utils.settings import QSettings_writeArray, QSettings_readArray
from Orange.widgets.widget import Msg, Output, OWWidget
+from orangecanvas.utils import findf
from orangecontrib.text.corpus import Corpus
from orangecontrib.text.language import (
@@ -15,13 +24,14 @@
AGGREGATORS,
AGGREGATORS_ITEMS,
DocumentEmbedder,
- LANGUAGES,
+ LANGUAGES, OAIDocumentEmbedder,
)
from orangecontrib.text.vectorization.sbert import SBERT
from orangecontrib.text.widgets.utils.owbasevectorizer import (
OWBaseVectorizer,
Vectorizer,
)
+from orangecontrib.text.widgets.utils.llmmodelwidget import LLMModelWidget
class EmbeddingVectorizer(Vectorizer):
@@ -33,6 +43,35 @@ def _transform(self, callback):
self.skipped_documents = skipped
+class Methods(enum.Enum):
+ SBERT = 0
+ FastText = 1
+ OpenAIEmbedder = 2
+
+
+Providers = {
+ "ollama": "http://localhost:11434/v1",
+ "llama.cpp": "http://localhost:9931/v1",
+}
+
+
+def is_localhost(url: str) -> bool:
+ """
+ Check if the given URL points to localhost.
+
+ Args:
+ url (str): The URL to check
+
+ Returns:
+ bool: True if the URL points to localhost, False otherwise
+ """
+ from urllib.parse import urlparse
+ parsed_url = urlparse(url)
+ localhost_hosts = {'localhost', '127.0.0.1', '::1'}
+ return parsed_url.hostname in localhost_hosts
+
+ServiceName = "Orange: API key"
+
class OWDocumentEmbedding(OWBaseVectorizer):
name = "Document Embedding"
description = "Document embedding using pretrained models."
@@ -43,7 +82,7 @@ class OWDocumentEmbedding(OWBaseVectorizer):
buttons_area_orientation = Qt.Vertical
settings_version = 3
- Methods = [SBERT, DocumentEmbedder]
+ Methods = [SBERT, DocumentEmbedder, OAIDocumentEmbedder]
class Outputs(OWBaseVectorizer.Outputs):
skipped = Output("Skipped documents", Corpus)
@@ -54,6 +93,9 @@ class Error(OWWidget.Error):
"another vectorizer."
)
unexpected_error = Msg("Embedding error: {}")
+ authentication_error = Msg("Authentication error: {}")
+ api_spec_error = Msg("API call error: {}")
+ connection_error = Msg("Connection error. Check that the service is running and is accessible. {}")
class Warning(OWWidget.Warning):
unsuccessful_embeddings = Msg("Some embeddings were unsuccessful.")
@@ -62,6 +104,10 @@ class Warning(OWWidget.Warning):
language: str = Setting(default=DEFAULT_LANGUAGE, schema_only=True)
aggregator: str = Setting(default="Mean")
+ base_url = Setting("", schema_only=True)
+ api_key = ""
+ model = Setting("", schema_only=True)
+
def __init__(self):
super().__init__()
self.cancel_button = QPushButton(
@@ -73,6 +119,13 @@ def __init__(self):
# it should be only set when setting loaded from schema/workflow
self.__pending_language = self.language
+ DefaultModels = [
+ {"base_url": "ollama", "model": "snowflake-arctic-embed:22m"},
+ {"base_url": "ollama", "model": "granite-embedding:30m"},
+ {"base_url": "ollama", "model": "all-minilm"},
+ {"base_url": "https://api.openai.com/v1", "model": ""}
+ ]
+
def create_configuration_layout(self):
layout = QVBoxLayout()
rbtns = gui.radioButtons(None, self, "method", callback=self.on_change)
@@ -80,7 +133,7 @@ def create_configuration_layout(self):
gui.appendRadioButton(rbtns, "Multilingual SBERT")
gui.appendRadioButton(rbtns, "fastText:")
- ibox = gui.indentedBox(rbtns)
+ self.fast_text_controls = ibox = gui.indentedBox(rbtns)
self.language_cb = gui.comboBox(
ibox,
self,
@@ -103,6 +156,19 @@ def create_configuration_layout(self):
callback=self.on_change,
searchable=True,
)
+ gui.appendRadioButton(rbtns, "Other (OpenAI API compatible):")
+ self.oai_controls = ibox = gui.indentedBox(rbtns)
+ self.llmapiwidget = LLMModelWidget(keyringServiceName=ServiceName)
+ ibox.layout().addWidget(self.llmapiwidget)
+
+ items = self._load_history()
+ items = items + self.DefaultModels
+ self.llmapiwidget.setHistory(items)
+ if self.base_url and self.model:
+ self.llmapiwidget.setBaseUrl(self.base_url)
+ self.llmapiwidget.setModelId(self.model)
+ self.api_key = self.llmapiwidget.apiKey()
+ self.llmapiwidget.changed.connect(self.on_api_param_change)
return layout
@OWBaseVectorizer.Inputs.corpus
@@ -122,14 +188,82 @@ def set_data(self, corpus):
super().set_data(corpus)
def update_method(self):
- disabled = self.method == 0
- self.aggregator_cb.setDisabled(disabled)
- self.language_cb.setDisabled(disabled)
+ method = Methods(self.method)
+ self.fast_text_controls.setEnabled(method == Methods.FastText)
+ self.oai_controls.setEnabled(method == Methods.OpenAIEmbedder)
self.vectorizer = EmbeddingVectorizer(self.init_method(), self.corpus)
+ def on_change(self):
+ if Methods(self.method) != Methods.OpenAIEmbedder:
+ self.Error.api_spec_error.clear()
+ self.Error.authentication_error.clear()
+ self.Error.connection_error.clear()
+ super().on_change()
+
+ def on_api_param_change(self):
+ self.base_url = self.llmapiwidget.baseUrl()
+ self.model = self.llmapiwidget.modelId()
+ self.api_key = self.llmapiwidget.apiKey()
+ if self.base_url and self.model:
+ self._save_history_item(
+ {"base_url": self.base_url, "model": self.model,
+ "has_key": bool(self.api_key)}
+ )
+ self.on_change()
+
+ @classmethod
+ def _local_settings(cls) -> QSettings:
+ """Return a QSettings instance with local persistent QSettings for `cls`."""
+ filename = "{}.ini".format(qname(cls))
+ fname = os.path.join(settings.widget_settings_dir(versioned=False), filename)
+ return QSettings(fname, QSettings.IniFormat)
+
+ def _save_history_item(self, item: LLMModelWidget.Item):
+ settings = self._local_settings()
+ items = self._load_history()
+ # find/replace item in stored history
+ existing = findf(items, lambda it: it["base_url"] == item["base_url"] and it["model"] == item["model"])
+ if existing:
+ items.remove(existing)
+ items.insert(0, item)
+ QSettings_writeArray(settings, "endpoints", items)
+
+ def _load_history(self) -> list[LLMModelWidget.Item]:
+ settings = self._local_settings()
+ items = QSettings_readArray(settings, "endpoints", {
+ "base_url": str, "model": str, "has_key": bool
+ })
+ items = [item for item in items if item["base_url"].strip() and item["model"].strip()]
+ return items
+
+ def set_base_url(self, url):
+ self.llmapiwidget.setBaseUrl(url)
+
+ def set_api_key(self, key):
+ self.llmapiwidget.setApiKey(key)
+
+ def set_model(self, model):
+ self.llmapiwidget.setModelId(model)
+
def init_method(self):
- params = dict(language=self.language, aggregator=self.aggregator)
- kwargs = ({}, params)[self.method]
+ method = Methods(self.method)
+ match method:
+ case Methods.SBERT:
+ kwargs = {}
+ case Methods.FastText:
+ kwargs = dict(language=self.language, aggregator=self.aggregator)
+ case Methods.OpenAIEmbedder:
+ base_url = self.llmapiwidget.baseUrl()
+ api_key = self.llmapiwidget.apiKey()
+ model = self.llmapiwidget.modelId()
+ base_url = Providers.get(base_url, base_url)
+ if is_localhost(base_url) and not api_key.strip():
+ # local providers probably do not need a key but
+ # `openai.Client` still complains about it.
+ api_key = "sk-no-key-required"
+ kwargs = dict(base_url=base_url, api_key=api_key, model=model)
+ case _:
+ raise NameError
return self.Methods[self.method](**kwargs)
@gui.deferred
@@ -148,11 +282,26 @@ def on_done(self, result):
super().on_done(result)
def on_exception(self, ex: Exception):
+ def oaie_message(ex: openai.APIStatusError) -> str:
+ """Extract message from openai error"""
+ try:
+ return json.loads(ex.response.content)["error"]["message"]
+ except (json.JSONDecodeError, KeyError, AttributeError):
+ return ex.message
self.cancel_button.setDisabled(True)
if isinstance(ex, EmbeddingConnectionError):
self.Error.no_connection()
+ elif isinstance(ex, openai.AuthenticationError):
+ self.Error.authentication_error(oaie_message(ex))
+ elif isinstance(ex, openai.BadRequestError):
+ self.Error.api_spec_error(oaie_message(ex))
+ elif isinstance(ex, openai.APIStatusError):
+ self.Error.api_spec_error(oaie_message(ex))
+ elif isinstance(ex, openai.APIConnectionError):
+ ex = ex.__cause__ if ex.__cause__ is not None else ex
+ self.Error.connection_error(str(ex))
else:
- self.Error.unexpected_error(str(ex))
+ self.Error.unexpected_error(str(ex), exc_info=ex)
self.cancel()
def cancel(self):
@@ -174,17 +323,22 @@ def migrate_settings(cls, settings: Dict[str, Any], version: Optional[int]):
settings["language"] = LANG2ISO[settings["language"]]
def send_report(self):
- if self.method == 0:
- self.report_items((
- ("Embedder", "Multilingual SBERT"),
- ))
- if self.method == 1:
- items = (
- ("Embedder", "fastText"),
- ("Language", ISO2LANG[self.language]),
- ("Aggregator", self.aggregator),
- )
- self.report_items(items)
+ match Methods(self.method):
+ case Methods.SBERT:
+ self.report_items((
+ ("Embedder", "Multilingual SBERT"),
+ ))
+ case Methods.FastText:
+ self.report_items((
+ ("Embedder", "fastText"),
+ ("Language", ISO2LANG[self.language]),
+ ("Aggregator", self.aggregator),
+ ))
+ case Methods.OpenAIEmbedder:
+ self.report_items ((
+ ("Base Api", self.base_url),
+ ("Model", self.model),
+ ))
if __name__ == "__main__":
diff --git a/orangecontrib/text/widgets/tests/test_owdocumentembedding.py b/orangecontrib/text/widgets/tests/test_owdocumentembedding.py
index d79333a6b..f4d1f22bc 100644
--- a/orangecontrib/text/widgets/tests/test_owdocumentembedding.py
+++ b/orangecontrib/text/widgets/tests/test_owdocumentembedding.py
@@ -1,5 +1,8 @@
+import json
import unittest
-from unittest.mock import Mock, patch
+from unittest.mock import Mock, patch, MagicMock
+
+import openai
import numpy as np
from AnyQt.QtWidgets import QComboBox, QRadioButton
@@ -9,6 +12,7 @@
from orangecontrib.text.language import DEFAULT_LANGUAGE, ISO2LANG
from orangecontrib.text.tests.test_documentembedder import PATCH_METHOD, make_dummy_post
+from orangecontrib.text.vectorization import document_embedder
from orangecontrib.text.vectorization.document_embedder import (
DocumentEmbedder,
LANGUAGES,
@@ -243,5 +247,109 @@ def test_migrate_settings(self):
self.assertEqual(iso_lang, widget.language)
+class TestOWDocumentEmbeddingOAI(WidgetTest):
+ def setUp(self):
+ super().setUp()
+ self.widget = self.create_widget(OWDocumentEmbedding, stored_settings={
+ "base_url": "localhost:8000", "model": "embedder", "method": 2
+ })
+ self.corpus = Corpus.from_file('deerwester')
+ self.larger_corpus = Corpus.from_file('book-excerpts')
+
+ # Disable embedder caches.
+ cache = MagicMock()
+ cache.md5_hash = lambda _: b""
+ cache.get_cached_result_or_none = lambda _: None
+ self._patch = patch.object(
+ document_embedder, "EmbedderCache", MagicMock(return_value=cache)
+ )
+ self._patch.__enter__()
+
+ def tearDown(self):
+ self._patch.__exit__(None, None, None)
+ super().tearDown()
+
+ @patch("openai.OpenAI")
+ def test_openai_output(self, mock_openai):
+ """Test that OpenAI embedder (method 2) produces correct output."""
+ # Mock the OpenAI client response
+ mock_embedding = MagicMock()
+ mock_embedding.embedding = np.arange(EMB_DIM, dtype=float).tolist()
+ mock_response = MagicMock()
+ mock_response.data = [mock_embedding] * len(self.corpus)
+ mock_client = MagicMock()
+ mock_client.embeddings.create.return_value = mock_response
+ mock_openai.return_value = mock_embedding
+
+ # Select OpenAI embedder (third radio button, index 2)
+ self.widget.findChildren(QRadioButton)[2].click()
+ self.send_signal("Corpus", self.corpus)
+ self.wait_until_finished()
+ result = self.get_output(self.widget.Outputs.corpus)
+ self.assertIsNotNone(result)
+ self.assertIsInstance(result, Corpus)
+ self.assertEqual(len(self.corpus), len(result))
+
+ @patch("openai.OpenAI")
+ def test_openai_authentication_error(self, mock_openai):
+ """Test authentication error handling for OpenAI embedder."""
+ mock_response = MagicMock()
+ mock_response.content = json.dumps(
+ {"error": {"message": "Invalid API key"}}
+ ).encode()
+ error = openai.AuthenticationError(
+ "Invalid API key", response=mock_response, body=None
+ )
+ mock_client = MagicMock()
+ mock_client.embeddings.create.side_effect = error
+ mock_openai.return_value = mock_client
+
+ self.widget.findChildren(QRadioButton)[2].click()
+ self.send_signal("Corpus", self.corpus)
+ self.wait_until_finished()
+ self.assertIsNone(self.get_output(self.widget.Outputs.corpus))
+ self.assertTrue(self.widget.Error.authentication_error.is_shown())
+
+ @patch("openai.OpenAI")
+ def test_openai_api_spec_error(self, mock_openai):
+ """Test API spec error handling for OpenAI embedder."""
+ mock_response = MagicMock()
+ mock_response.content = json.dumps(
+ {"error": {"message": "Invalid model: bad-model"}}
+ ).encode()
+ error = openai.BadRequestError(
+ "Bad request", response=mock_response, body=None
+ )
+ mock_client = MagicMock()
+ mock_client.embeddings.create.side_effect = error
+ mock_openai.return_value = mock_client
+
+ self.widget.findChildren(QRadioButton)[2].click()
+ self.send_signal("Corpus", self.corpus)
+ self.wait_until_finished()
+ self.assertIsNone(self.get_output(self.widget.Outputs.corpus))
+ self.assertTrue(self.widget.Error.api_spec_error.is_shown())
+
+ @patch("openai.OpenAI")
+ def test_openai_connection_error(self, mock_openai):
+ """Test connection error handling for OpenAI embedder."""
+ mock_request = MagicMock()
+ error = openai.APIConnectionError(message="Connection refused", request=mock_request)
+ mock_client = MagicMock()
+ mock_client.embeddings.create.side_effect = error
+ mock_openai.return_value = mock_client
+
+ self.widget.findChildren(QRadioButton)[2].click()
+ self.send_signal("Corpus", self.corpus)
+ self.wait_until_finished()
+ self.assertIsNone(self.get_output(self.widget.Outputs.corpus))
+ self.assertTrue(self.widget.Error.connection_error.is_shown())
+
+ def test_report(self):
+ self.widget.findChildren(QRadioButton)[2].click()
+ self.send_signal(self.widget.Inputs.corpus, self.corpus)
+ self.widget.send_report()
+
+
if __name__ == "__main__":
unittest.main()
diff --git a/orangecontrib/text/widgets/utils/llmmodelwidget.py b/orangecontrib/text/widgets/utils/llmmodelwidget.py
new file mode 100644
index 000000000..bf7f86e9b
--- /dev/null
+++ b/orangecontrib/text/widgets/utils/llmmodelwidget.py
@@ -0,0 +1,335 @@
+import logging
+
+from typing import Mapping, Any, TypedDict, Iterable
+
+import keyring
+from more_itertools import unique_everseen
+
+from AnyQt.QtCore import (
+ Qt, QEvent, QModelIndex, QAbstractItemModel, QObject, Signal,
+)
+from AnyQt.QtGui import QStandardItemModel, QStandardItem
+from AnyQt.QtWidgets import QWidget, QFormLayout, QApplication, QLineEdit
+
+from orangecanvas.utils import group_by_all
+from orangecontrib.text.widgets.utils.passwordedit import PasswordEdit
+
+from Orange.widgets.utils.combobox import TextEditCombo
+
+ApiKeyRole = Qt.UserRole + 41
+#: Flag indicating if a user already entered api key for a provider.
+#: Used to avoid premature calls to `keyring.get_password`.
+HasApiKeyRole = Qt.UserRole + 42
+#: Stores model id
+ModelRole = Qt.UserRole + 43
+
+log = logging.getLogger(__name__)
+
+
+class TextCombo(TextEditCombo):
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.__insertPolicy = self.insertPolicy()
+ self.setLineEdit(QLineEdit(self))
+
+ def setLineEdit(self, edit: QLineEdit) -> None:
+ if edit == self.lineEdit():
+ return
+ try:
+ old = self.lineEdit()
+ old.returnPressed.disconnect(self.__le_rp_before)
+ old.returnPressed.disconnect(self.__le_rp_after)
+ except TypeError:
+ pass
+ edit.returnPressed.connect(self.__le_rp_before)
+ super().setLineEdit(edit)
+ edit.returnPressed.connect(self.__le_rp_after)
+
+ def __le_rp_before(self):
+ # Disable insertion before ComboBox can process returnPressed signal
+ # from the line edit
+ self.__insertPolicy = self.insertPolicy()
+ self.setInsertPolicy(TextCombo.NoInsert)
+
+ def __le_rp_after(self):
+ # Re-enable insertion
+ self.setInsertPolicy(self.__insertPolicy)
+ # Move focus to trigger TextEditCombo.__on_editingFinished
+ self.focusNextChild()
+
+ def insertItem(self, index: int, text: str, userData = None):
+ """Reimplemented.
+
+ QComboBox does not properly insert under rootModelIndex when model
+ is a QStandardItemModel.
+ This only works if called from `TextEditCombo.__on_editingFinished`
+ """
+ model = self.model()
+ if model is None or not text.strip():
+ return
+ root = self.rootModelIndex()
+ if isinstance(model, QStandardItemModel):
+ item = Item({Qt.DisplayRole: text, Qt.UserRole: userData})
+ ritem = model.itemFromIndex(root)
+ if ritem is None:
+ ritem = model.invisibleRootItem()
+ count = ritem.rowCount()
+ ritem.insertRow(index, item)
+ if count == 0: # Need to update current state if count was 0 before
+ self.setCurrentIndex(0)
+ else:
+ super().insertItem(index, text, userData)
+
+ def addItem(self, text, userData = None):
+ self.insertItem(self.count(), text, userData)
+
+
+class Item(QStandardItem):
+ def __init__(self, data: Mapping[int, Any]):
+ super().__init__()
+ for role, value in data.items():
+ self.setData(value, role)
+
+
+def move_up_helper(model: QStandardItemModel, parent: QModelIndex, index: int):
+ """Move the `index` row in model to first position."""
+ if index < 1:
+ return
+ root = model.itemFromIndex(parent)
+ if root is None:
+ root = model.invisibleRootItem()
+ if 0 <= index < root.rowCount():
+ row = root.takeRow(index)
+ root.insertRow(0, row)
+
+
+class LLMModelWidget(QWidget):
+ """
+ A widget form for entering llm provider endpoint url with api key and
+ model selection. The api key is stored using `keyring`
+
+ Parameters:
+ keyringServiceName:
+ The keyring service name under which the entered api key is stored.
+ """
+ #: Signal emitted when the data entered by the user changes.
+ changed = Signal()
+ #: Signal emitted when widget loses focus or Enter/Return is pressed.
+ editingFinished = Signal()
+
+ class Item(TypedDict):
+ base_url: str
+ model: str
+ has_key: bool
+
+ def __init__(self, *args, keyringServiceName: str, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.__edited: bool = False
+ self.keyringServiceName = keyringServiceName
+ self.setContentsMargins(0, 0, 0, 0)
+ form = QFormLayout(
+ formAlignment=Qt.AlignLeft,
+ labelAlignment=Qt.AlignLeft,
+ fieldGrowthPolicy=QFormLayout.FieldGrowthPolicy.AllNonFixedFieldsGrow,
+ )
+ form.setContentsMargins(0, 0, 0, 0)
+ self.base_url_cb = TextCombo(
+ insertPolicy=TextCombo.InsertAtTop
+ )
+ self._model = QStandardItemModel()
+ self.base_url_cb.setModel(self._model)
+
+ self.base_url_cb.editingFinished.connect(self._on_base_url_edited)
+ self.base_url_cb.editTextChanged.connect(self._mark_edited)
+ self.base_url_cb.currentIndexChanged.connect(self._on_base_url_index_change)
+ self.api_key_le = PasswordEdit(
+ echoMode=PasswordEdit.EchoMode.PasswordEchoOnEdit,
+ placeholderText="sk-...",
+ )
+
+ self.api_key_le.editingFinished.connect(self._on_api_key_edit)
+ self.api_key_le.returnPressed.connect(self.api_key_le.focusNextChild)
+ self.api_key_le.textEdited.connect(self._mark_edited)
+ self.model_cb = TextCombo(
+ placeholderText="model",
+ insertPolicy=TextCombo.InsertAtTop
+ )
+ self.model_cb.setModel(self._model)
+ self.model_cb.editingFinished.connect(self._on_model_id_edit)
+ self.model_cb.editTextChanged.connect(self._mark_edited)
+
+ form.addRow("API Base Url:", self.base_url_cb)
+ form.addRow("Api Key:", self.api_key_le)
+ form.addRow("Model:", self.model_cb)
+ self.base_url_cb.installEventFilter(self)
+ self.api_key_le.installEventFilter(self)
+ self.model_cb.installEventFilter(self)
+ self.setLayout(form)
+
+ def setHistory(self, history: list[Item]):
+ """
+ Set the history.
+ """
+ model = self._model
+ model.clear()
+ items = unique_everseen(history, key=lambda item: (item["base_url"], item["model"]))
+ items_by_url = group_by_all(items, key=lambda item: item["base_url"])
+ for i, (base_url, items) in enumerate(items_by_url):
+ item = Item({
+ Qt.DisplayRole: base_url,
+ HasApiKeyRole: items[0].get("has_key"),
+ })
+ item.setFlags(Qt.ItemIsEditable | Qt.ItemIsEnabled | Qt.ItemIsSelectable)
+ model.appendRow(item)
+ for j, itm in enumerate(items):
+ item.appendRow(Item({
+ Qt.DisplayRole: itm["model"],
+ ModelRole: itm["model"],
+ }))
+
+ if model.rowCount():
+ self.model_cb.setRootModelIndex(self._model.index(0, 0))
+ self.model_cb.setCurrentIndex(0)
+
+ self._on_base_url_change()
+
+ def history(self) -> list[Item]:
+ """Return the history"""
+ model = self.base_url_cb.model()
+ items = []
+ roles = {Qt.DisplayRole: "base_url", ModelRole: "model", HasApiKeyRole: "has_key"}
+
+ def values(model: QAbstractItemModel, midx: QModelIndex, roles: Iterable[int]) -> dict:
+ return {role: model.data(midx, role) for role in roles}
+
+ for i in range(model.rowCount()):
+ midx = model.index(i, 0)
+ endp = values(model, midx, roles.keys())
+ for j in range(model.rowCount(midx)):
+ vals = {
+ **endp,
+ **{ModelRole: model.index(j, 0, midx).data(Qt.DisplayRole)},
+ }
+ vals = {roles[r]: vals[r] for r in roles}
+ items.append(vals)
+ return items
+
+ def baseUrl(self) -> str:
+ """Return the base api url."""
+ return self.base_url_cb.currentText()
+
+ def setBaseUrl(self, url: str) -> None:
+ """Set the base api url."""
+ current = self.base_url_cb.currentText()
+ if current == url:
+ return
+ self.base_url_cb.setText(url)
+ self._on_base_url_change()
+
+ def _on_base_url_change(self):
+ model = self.base_url_cb.model()
+ index = self.base_url_cb.currentIndex()
+ if index != 0:
+ move_up_helper(model, QModelIndex(), index)
+ self.base_url_cb.setCurrentIndex(0)
+
+ url = self.base_url_cb.currentText()
+ midx = model.index(self.base_url_cb.currentIndex(), 0)
+ self.model_cb.setRootModelIndex(midx)
+ if self.model_cb.currentIndex() == -1 and self.model_cb.count():
+ self.model_cb.setCurrentIndex(0)
+ api_key = ""
+ if self.base_url_cb.currentData(HasApiKeyRole):
+ api_key = self._get_secret(url) or ""
+ self.__set_api_key(api_key, store=False)
+
+ def _on_base_url_edited(self):
+ self._on_base_url_change()
+ self.__emit_changed()
+
+ def _on_base_url_index_change(self, index):
+ model = self.base_url_cb.model()
+ self.model_cb.setRootModelIndex(model.index(index, 0))
+ self.model_cb.setCurrentIndex(0)
+
+ def _mark_edited(self):
+ self.__edited = True
+
+ def apiKey(self) -> str:
+ """Return the api key."""
+ return self.api_key_le.text()
+
+ def setApiKey(self, key: str) -> None:
+ """Set the api key."""
+ if key == self.apiKey():
+ return
+ self.__set_api_key(key)
+
+ def __set_api_key(self, key, store=True):
+ base_url = self.baseUrl()
+ if key and base_url and store:
+ self._store_secret(base_url, key)
+ self.api_key_le.setText(key)
+ index = self.base_url_cb.currentIndex()
+ self.base_url_cb.setItemData(index, key, ApiKeyRole)
+ self.base_url_cb.setItemData(index, bool(key), HasApiKeyRole)
+
+ def _on_api_key_edit(self):
+ # coming from QLineEdit.editingFinished which triggers on focus out
+ # after setText(text) even when the text is not modified.
+ modified = self.api_key_le.isModified()
+ self.__set_api_key(self.apiKey())
+ if modified:
+ self.__emit_changed()
+
+ def modelId(self) -> str:
+ """Return the model id"""
+ return self.model_cb.currentText()
+
+ def setModelId(self, modelId: str) -> None:
+ """Set model id"""
+ if modelId == self.model_cb.currentText():
+ return
+ self.model_cb.setText(modelId)
+ self._on_model_id_change()
+
+ def _on_model_id_change(self):
+ # move the model to first index
+ index = self.model_cb.currentIndex()
+ model = self.model_cb.model()
+ rootmidx = self.model_cb.rootModelIndex()
+ if index != 0:
+ move_up_helper(model, rootmidx, index)
+ self.model_cb.setCurrentIndex(0)
+
+ def _on_model_id_edit(self) -> None:
+ self._on_model_id_change()
+ self.__emit_changed()
+
+ def eventFilter(self, recv: QObject, event: QEvent) -> bool:
+ if event.type() == QEvent.FocusOut and event.reason() == Qt.FocusReason.TabFocusReason:
+ newfocus = QApplication.focusWidget()
+ if self.__edited and not self.isAncestorOf(newfocus):
+ self.__emit_editingFinished()
+ return super().eventFilter(recv, event)
+
+ def __emit_changed(self):
+ self.__edited = True
+ self.changed.emit()
+
+ def __emit_editingFinished(self):
+ self.__edited = False
+ self.editingFinished.emit()
+
+ def _get_secret(self, service):
+ try:
+ return keyring.get_password(self.keyringServiceName, service)
+ except Exception:
+ log.exception("Failed to get secret for '%r'.", service)
+ return None
+
+ def _store_secret(self, service, password):
+ try:
+ keyring.set_password(self.keyringServiceName, service, password)
+ except Exception:
+ log.exception("Failed to set secret for '%s'.", service)
\ No newline at end of file
diff --git a/orangecontrib/text/widgets/utils/passwordedit.py b/orangecontrib/text/widgets/utils/passwordedit.py
new file mode 100644
index 000000000..a19e42827
--- /dev/null
+++ b/orangecontrib/text/widgets/utils/passwordedit.py
@@ -0,0 +1,40 @@
+from AnyQt.QtCore import QTimer
+from AnyQt.QtWidgets import QLineEdit, QAction
+
+from orangewidget.utils import load_styled_icon
+
+
+class PasswordEdit(QLineEdit):
+ """Password entry widget with reveal/conceal action."""
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.setEchoMode(PasswordEdit.EchoMode.PasswordEchoOnEdit)
+ self.timer = QTimer(singleShot=True, interval=5000)
+ self.timer.timeout.connect(self.concealText)
+ self._icons = [
+ load_styled_icon(__package__, "../icons/show-password.svg"),
+ load_styled_icon(__package__, "../icons/hide-password.svg"),
+ ]
+ self._action = ac = QAction("Show password", self)
+ ac.setIcon(self._icons[0])
+ ac.triggered.connect(self.toggleEchoMode)
+ self.addAction(ac, QLineEdit.TrailingPosition)
+
+ def toggleEchoMode(self):
+ """Toggle echo mode"""
+ if self.echoMode() == PasswordEdit.EchoMode.Normal:
+ self.concealText()
+ else:
+ self.revealText()
+
+ def revealText(self):
+ """Temporarily reveal password."""
+ self._action.setIcon(self._icons[1])
+ self.setEchoMode(PasswordEdit.EchoMode.Normal)
+ self.timer.start()
+
+ def concealText(self):
+ """Conceal the password."""
+ self._action.setIcon(self._icons[0])
+ self.setEchoMode(PasswordEdit.EchoMode.PasswordEchoOnEdit)
+ self.timer.stop()
diff --git a/orangecontrib/text/widgets/utils/tests/test_llmmodelwidget.py b/orangecontrib/text/widgets/utils/tests/test_llmmodelwidget.py
new file mode 100644
index 000000000..9e0e5723a
--- /dev/null
+++ b/orangecontrib/text/widgets/utils/tests/test_llmmodelwidget.py
@@ -0,0 +1,291 @@
+import unittest
+from unittest.mock import patch
+
+from AnyQt.QtWidgets import QLineEdit, QApplication, QWidget
+from AnyQt.QtCore import QEvent, Qt
+from AnyQt.QtGui import QStandardItemModel, QFocusEvent
+from AnyQt.QtTest import QTest, QSignalSpy
+
+from Orange.widgets.tests.base import GuiTest
+
+from orangecontrib.text.widgets.utils.llmmodelwidget import (
+ LLMModelWidget,
+ TextCombo,
+)
+
+
+class TestTextCombo(GuiTest):
+ """Tests for the TextCombo class."""
+
+ def setUp(self):
+ super().setUp()
+ self.combo = TextCombo()
+ self.model = QStandardItemModel()
+ self.combo.setModel(self.model)
+
+ def tearDown(self) -> None:
+ del self.combo
+ del self.model
+ super().tearDown()
+
+ def test_add_item(self):
+ self.combo.addItem("first item", "user-data-1")
+ self.assertEqual(self.combo.count(), 1)
+ self.assertEqual(self.combo.currentText(), "first item")
+
+ def test_insert_item_child(self):
+ self.combo.addItem("root item")
+ self.assertEqual(self.combo.count(), 1)
+ self.assertEqual(self.combo.currentText(), "root item")
+ self.combo.setRootModelIndex(self.model.index(0, 0))
+ self.combo.addItem("child item")
+ self.assertEqual(self.combo.count(), 1)
+ self.assertEqual(self.combo.currentText(), "child item")
+ self.combo.insertItem(0, "child item 1")
+ self.assertEqual(self.combo.count(), 2)
+ self.assertEqual(self.combo.currentText(), "child item")
+ self.assertEqual(self.combo.currentIndex(), 1)
+ self.assertEqual(self.model.rowCount(), 1)
+
+ def test_insert_item_empty_text(self):
+ """Empty text should not insert."""
+ initial_count = self.combo.count()
+ self.combo.insertItem(0, "", "data")
+ self.assertEqual(self.combo.count(), initial_count)
+
+ def test_insert_item_whitespace_only(self):
+ """Whitespace-only text should not insert."""
+ initial_count = self.combo.count()
+ self.combo.insertItem(0, " ", "data")
+ self.assertEqual(self.combo.count(), initial_count)
+
+
+def enter_text(widget: QLineEdit, text: str, enter: bool=True):
+ widget.selectAll()
+ QTest.keyClick(widget, Qt.Key.Key_Delete)
+ QTest.keyClicks(widget, text)
+ if enter:
+ QTest.keyClick(widget, Qt.Key.Key_Return)
+
+
+def send_focus_out(widget: QWidget, reason =Qt.TabFocusReason):
+ event = QFocusEvent(QEvent.FocusOut, reason)
+ QApplication.sendEvent(widget, event)
+
+
+class TestLLMModelWidget(GuiTest):
+ """Tests for the LLMModelWidget widget."""
+
+ def setUp(self):
+ super().setUp()
+ self.widget = LLMModelWidget(keyringServiceName="test-service")
+
+ def tearDown(self) -> None:
+ del self.widget
+ super().tearDown()
+
+ def test_initial_state(self):
+ """Test widget initial state."""
+ self.assertEqual(self.widget.baseUrl(), "")
+ self.assertEqual(self.widget.apiKey(), "")
+ self.assertEqual(self.widget.modelId(), "")
+
+ def test_edit(self):
+ """Test user simulated editing."""
+ spy = QSignalSpy(self.widget.editingFinished)
+ le = self.widget.base_url_cb.lineEdit()
+ enter_text(le, "http://localhost", enter=False)
+ self.assertEqual(self.widget.baseUrl(), "http://localhost")
+ self.assertEqual(len(spy), 0)
+ enter_text(self.widget.model_cb.lineEdit(), "model", enter=False)
+ self.assertEqual(self.widget.modelId(), "model")
+ self.assertEqual(len(spy), 0)
+ # move focus out of the widget
+ QApplication.sendEvent(
+ self.widget.model_cb, QFocusEvent(QEvent.FocusOut, Qt.TabFocusReason)
+ )
+ self.assertEqual(len(spy), 1)
+
+ @patch("keyring.get_password", return_value=None)
+ def test_set_base_url(self, mock_get):
+ """Test setting base URL after history is set."""
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ )
+ ]
+ self.widget.setHistory(items)
+ spy = QSignalSpy(self.widget.changed)
+ self.widget.setBaseUrl("https://api.example.com")
+ self.assertEqual(self.widget.baseUrl(), "https://api.example.com")
+ self.assertEqual(len(spy), 0)
+
+ @patch("keyring.set_password")
+ @patch("keyring.get_password", return_value=None)
+ def test_api_key_persisted_to_keyring(self, mock_get, mock_set):
+ """Test that API key is stored in keyring."""
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ )
+ ]
+ self.widget.setHistory(items)
+ self.widget.setBaseUrl("https://api.example.com")
+ spy = QSignalSpy(self.widget.changed)
+ self.widget.setApiKey("sk-secret-key")
+ mock_set.assert_called_once_with("test-service", "https://api.example.com", "sk-secret-key")
+ self.assertEqual(self.widget.apiKey(), "sk-secret-key")
+ mock_set.reset_mock()
+ enter_text(self.widget.api_key_le, "sk-key")
+ mock_set.assert_called_once_with("test-service", "https://api.example.com", "sk-key")
+ self.assertEqual(len(spy), 1)
+
+ @patch("keyring.get_password", return_value=None)
+ def test_set_model_id(self, mock_get):
+ """Test setting model ID after history is set."""
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ )
+ ]
+ self.widget.setHistory(items)
+ self.widget.setModelId("llmmodel")
+ self.assertEqual(self.widget.modelId(), "llmmodel")
+ # get_password must not be called when has_key is False
+ mock_get.assert_not_called()
+
+ @patch("keyring.get_password", return_value="sk-stored-key")
+ def test_set_history_with_api_key(self, mock_get):
+ """Test setting history where an API key is stored in keyring."""
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=True,
+ )
+ ]
+ self.widget.setHistory(items)
+ self.assertEqual(self.widget.apiKey(), "sk-stored-key")
+ mock_get.assert_called_once_with("test-service", "https://api.example.com")
+
+ @patch("keyring.get_password", return_value=None)
+ def test_set_history_multiple_urls(self, mock_get):
+ """Test setting history with multiple different base URLs."""
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ ),
+ LLMModelWidget.Item(
+ base_url="https://api.foo.com",
+ model="llmmodel-2",
+ has_key=False,
+ ),
+ ]
+ self.widget.setHistory(items)
+ history = self.widget.history()
+ self.assertEqual(len(history), 2)
+
+ urls = [h["base_url"] for h in history]
+ self.assertEqual(urls[0], "https://api.example.com")
+ self.assertEqual(urls[1], "https://api.foo.com")
+
+ @patch("keyring.get_password")
+ def test_changed_signal_on_api_key_edit(self, mock_get):
+ """Test that changed signal is emitted when API key is edited via line edit."""
+ mock_get.return_value = None
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ )
+ ]
+ self.widget.setHistory(items)
+ self.widget.setBaseUrl("https://api.example.com")
+ spy = QSignalSpy(self.widget.changed)
+ # Simulate editing the API key line edit
+ enter_text(self.widget.api_key_le, "sk-new-key")
+ self.assertEqual(len(spy), 1)
+
+ @patch("keyring.get_password")
+ def test_changed_signal_on_model_id_change(self, mock_get):
+ """Test that changed signal is emitted when model ID changes via line edit."""
+ mock_get.return_value = None
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ ),
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="gpt-3.5-turbo",
+ has_key=False,
+ ),
+ ]
+ self.widget.setHistory(items)
+
+ spy = QSignalSpy(self.widget.changed)
+ # Simulate editing the model combobox directly
+ self.widget.model_cb.setCurrentIndex(0)
+
+ # QTest.keyClick(..., Qt.KeyEnter ) in enter_text doubles returnPressed
+ # emit, this does not happen in normal event dispatch. Manually
+ # send focus out event to compensate.
+ enter_text(self.widget.model_cb.lineEdit(), "llmmodel-2", enter=False)
+ send_focus_out(self.widget.model_cb)
+ self.assertEqual(self.widget.modelId(), "llmmodel-2", )
+ self.assertEqual(len(spy), 1)
+
+ @patch("keyring.set_password")
+ @patch("keyring.get_password", return_value=None)
+ def test_api_key_not_stored_without_base_url(self, mock_get, mock_set):
+ """Test that API key is not stored if there's no base URL."""
+ self.widget.setApiKey("sk-secret-key")
+ mock_set.assert_not_called()
+
+ @patch("keyring.get_password", return_value=None)
+ def test_combobox_populated(self, mock_get):
+ """Test that base URL and model combobox are populated after setHistory."""
+ items = [
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel",
+ has_key=False,
+ ),
+ LLMModelWidget.Item(
+ base_url="https://api.example.com",
+ model="llmmodel-1",
+ has_key=False,
+ ),
+ LLMModelWidget.Item(
+ base_url="https://api.foo.com",
+ model="llmmodel-3",
+ has_key=False,
+ ),
+ LLMModelWidget.Item(
+ base_url="https://api.foo.com",
+ model="llmmodel-4",
+ has_key=False,
+ ),
+ ]
+ self.widget.setHistory(items)
+ # The base URL combobox should have 2 entries
+ self.assertEqual(self.widget.base_url_cb.count(), 2)
+ self.assertEqual(self.widget.baseUrl(), "https://api.example.com")
+ # Model combobox must have 2 entries
+ self.assertEqual(self.widget.model_cb.count(), 2)
+ self.assertEqual(self.widget.modelId(), "llmmodel")
+ self.assertEqual(self.widget.model_cb.itemText(1), "llmmodel-1")
+
+
+if __name__ == "__main__":
+ unittest.main()
\ No newline at end of file
diff --git a/orangecontrib/text/widgets/utils/tests/test_passwordedit.py b/orangecontrib/text/widgets/utils/tests/test_passwordedit.py
new file mode 100644
index 000000000..ab3abd9f1
--- /dev/null
+++ b/orangecontrib/text/widgets/utils/tests/test_passwordedit.py
@@ -0,0 +1,40 @@
+from orangewidget.tests.base import GuiTest
+from orangecontrib.text.widgets.utils.passwordedit import PasswordEdit
+
+
+class TestPasswordEdit(GuiTest):
+ """Tests for the PasswordEdit class."""
+
+ def setUp(self):
+ super().setUp()
+ self.edit = PasswordEdit()
+
+ def tearDown(self) -> None:
+ del self.edit
+ super().tearDown()
+
+ def test_initial_echo_mode(self):
+ self.assertEqual(self.edit.echoMode(), PasswordEdit.EchoMode.PasswordEchoOnEdit)
+
+ def test_set_and_get_text(self):
+ self.edit.setText("secret-key-123")
+ self.assertEqual(self.edit.text(), "secret-key-123")
+
+ def test_toggle_echo_mode(self):
+ self.edit.toggleEchoMode()
+ self.assertEqual(self.edit.echoMode(), PasswordEdit.EchoMode.Normal)
+ self.edit.toggleEchoMode()
+ self.assertEqual(self.edit.echoMode(), PasswordEdit.EchoMode.PasswordEchoOnEdit)
+
+ def test_reveal_conceal_text(self):
+ """revealText should set Normal echo mode."""
+ self.edit.revealText()
+ self.assertEqual(self.edit.echoMode(), PasswordEdit.EchoMode.Normal)
+ self.edit.concealText()
+ self.assertEqual(self.edit.echoMode(), PasswordEdit.EchoMode.PasswordEchoOnEdit)
+
+ def test_toggleEchoMode_normal_to_password(self):
+ """toggleEchoMode should switch from Normal to PasswordEchoOnEdit."""
+ self.edit.revealText()
+ self.edit.toggleEchoMode()
+ self.assertEqual(self.edit.echoMode(), PasswordEdit.EchoMode.PasswordEchoOnEdit)
diff --git a/requirements.txt b/requirements.txt
index fec3772c7..1058b4545 100644
--- a/requirements.txt
+++ b/requirements.txt
@@ -5,11 +5,14 @@ conllu
docx2txt>=0.6
gensim>=4.3.3
httpx!=0.23.1 # temporary fix - semantic search fail (but only in tests)
+keyring
langdetect
lemmagen3
+more_itertools
nltk>=3.9.1
numpy
odfpy>=1.3.5
+openai
Orange3 >=3.38.1
orange-widget-base >=4.25.0
orange-canvas-core >=0.2.5