diff --git a/Makefile b/Makefile index 0f6fe9b4..e13fc653 100644 --- a/Makefile +++ b/Makefile @@ -8,7 +8,7 @@ all: venv format check test build .PHONY: format format: venv - .venv/bin/black -tpy312 -tpy313 -tpy314 typeagent test tools gmail $(FLAGS) + .venv/bin/black -tpy312 -tpy313 -tpy314 typeagent test tools gmail demo $(FLAGS) .PHONY: check check: venv diff --git a/make.bat b/make.bat index d96b04c1..d0e0bd8c 100644 --- a/make.bat +++ b/make.bat @@ -26,7 +26,7 @@ goto help :format if not exist ".venv\" call make.bat venv echo Formatting code... -.venv\Scripts\black typeagent test tools +.venv\Scripts\black typeagent test tools gmail demo goto end :check diff --git a/test/test_interfaces.py b/test/test_interfaces.py index f7fb3bfa..8a7fbce7 100644 --- a/test/test_interfaces.py +++ b/test/test_interfaces.py @@ -293,7 +293,7 @@ def test_text_range_equality_with_logical_equivalence(): # Test with non-TextRange object assert point_range != "not a TextRange" - assert point_range != None + assert point_range is not None def test_text_range_equality_both_explicit_ends(): diff --git a/tools/ingest_vtt.py b/tools/ingest_vtt.py index ef311db4..080e1665 100644 --- a/tools/ingest_vtt.py +++ b/tools/ingest_vtt.py @@ -50,12 +50,6 @@ def create_arg_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser( description="Ingest WebVTT transcript files into a database for querying", formatter_class=argparse.RawDescriptionHelpFormatter, - epilog=""" -Examples: - %(prog)s input.vtt --database transcript.db - %(prog)s file1.vtt file2.vtt -d transcript.db --name "Combined Transcript" - %(prog)s lecture.vtt -d lecture.db --merge - """, ) parser.add_argument( @@ -90,6 +84,13 @@ def create_arg_parser() -> argparse.ArgumentParser: help="Batch size for knowledge extraction (default: from settings)", ) + parser.add_argument( + "--embedding-name", + type=str, + default=None, + help="Embedding model name (default: text-embedding-ada-002)", + ) + parser.add_argument( "-v", "--verbose", action="store_true", help="Show verbose output" ) @@ -135,6 +136,7 @@ async def ingest_vtt_files( merge_consecutive: bool = False, verbose: bool = False, batchsize: int | None = None, + embedding_name: str | None = None, ) -> None: """Ingest one or more VTT files into a database.""" @@ -197,7 +199,7 @@ async def ingest_vtt_files( if verbose: print("Setting up conversation settings...") try: - embedding_model = AsyncEmbeddingModel() + embedding_model = AsyncEmbeddingModel(model_name=embedding_name) settings = ConversationSettings(embedding_model) # Create storage provider explicitly with the database @@ -441,6 +443,7 @@ def main(): name=args.name, merge_consecutive=args.merge, batchsize=args.batchsize, + embedding_name=args.embedding_name, verbose=args.verbose, ) ) diff --git a/typeagent/knowpro/convsettings.py b/typeagent/knowpro/convsettings.py index bab5657c..a785e0b7 100644 --- a/typeagent/knowpro/convsettings.py +++ b/typeagent/knowpro/convsettings.py @@ -59,36 +59,28 @@ def __init__( # Storage provider will be created lazily if not provided self._storage_provider: IStorageProvider | None = storage_provider - self._storage_provider_created = storage_provider is not None @property def storage_provider(self) -> IStorageProvider: - if not self._storage_provider_created: + if self._storage_provider is None: raise RuntimeError( - "Storage provider not initialized. Use await ConversationSettings.get_storage_provider() " + "Storage provider not initialized. " + "Use await ConversationSettings.get_storage_provider() " "or provide storage_provider in constructor." ) - assert ( - self._storage_provider is not None - ), "Storage provider should be set when _storage_provider_created is True" return self._storage_provider @storage_provider.setter def storage_provider(self, value: IStorageProvider) -> None: self._storage_provider = value - self._storage_provider_created = True async def get_storage_provider(self) -> IStorageProvider: """Get or create the storage provider asynchronously.""" - if not self._storage_provider_created: + if self._storage_provider is None: from ..storage.memory import MemoryStorageProvider self._storage_provider = MemoryStorageProvider( message_text_settings=self.message_text_index_settings, related_terms_settings=self.related_term_index_settings, ) - self._storage_provider_created = True - assert ( - self._storage_provider is not None - ), "Storage provider should be set after creation" return self._storage_provider diff --git a/typeagent/storage/sqlite/provider.py b/typeagent/storage/sqlite/provider.py index c946918c..f43869c1 100644 --- a/typeagent/storage/sqlite/provider.py +++ b/typeagent/storage/sqlite/provider.py @@ -40,7 +40,6 @@ def __init__( conversation_id: str = "default", message_type: type[TMessage] = None, # type: ignore semantic_ref_type: type[interfaces.SemanticRef] = None, # type: ignore - conversation_index_settings=None, message_text_index_settings: MessageTextIndexSettings | None = None, related_term_index_settings: RelatedTermIndexSettings | None = None, ): @@ -50,14 +49,11 @@ def __init__( self.semantic_ref_type = semantic_ref_type # Settings with defaults (require embedding settings) - self.conversation_index_settings = conversation_index_settings or {} if message_text_index_settings is None: # Create default embedding settings if not provided - from ...aitools.embeddings import AsyncEmbeddingModel from ...aitools.vectorbase import TextEmbeddingIndexSettings - model = AsyncEmbeddingModel() - embedding_settings = TextEmbeddingIndexSettings(model) + embedding_settings = TextEmbeddingIndexSettings() message_text_index_settings = MessageTextIndexSettings(embedding_settings) self.message_text_index_settings = message_text_index_settings