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