From 0e93d53f0036d3f2b28f99994c9cc3c6ff7f321b Mon Sep 17 00:00:00 2001 From: harshkedia177 Date: Tue, 10 Mar 2026 01:30:58 +0530 Subject: [PATCH 1/2] refactor: optimize embedding process and enhance pipeline logging - Introduced a new `_embed_node_list` function to streamline the embedding of graph nodes, improving code reusability. - Updated the `embed_graph` and `embed_nodes` functions to utilize the new embedding method, enhancing clarity and maintainability. - Increased the default batch size for embeddings from 64 to 256 for improved performance. - Added logging to track the timing of various phases in the pipeline, providing better insights into processing durations. - Cleaned up the Kuzu backend to manage embeddings more efficiently, including a mechanism to wipe the embeddings table before bulk loading. - Adjusted tests to reflect changes in embedding parameters and batch sizes. --- src/axon/core/embeddings/embedder.py | 108 ++++++++--------- src/axon/core/ingestion/pipeline.py | 163 ++++++++++++++------------ src/axon/core/ingestion/processes.py | 3 +- src/axon/core/storage/kuzu_backend.py | 48 +++++--- tests/core/test_embedder.py | 10 +- 5 files changed, 173 insertions(+), 159 deletions(-) diff --git a/src/axon/core/embeddings/embedder.py b/src/axon/core/embeddings/embedder.py index 0ed8ef55a..eab3221ad 100644 --- a/src/axon/core/embeddings/embedder.py +++ b/src/axon/core/embeddings/embedder.py @@ -11,21 +11,28 @@ from __future__ import annotations +import logging import threading from typing import TYPE_CHECKING from axon.core.embeddings.text import build_class_method_index, generate_text from axon.core.graph.graph import KnowledgeGraph -from axon.core.graph.model import NodeLabel +from axon.core.graph.model import GraphNode, NodeLabel from axon.core.storage.base import NodeEmbedding if TYPE_CHECKING: from fastembed import TextEmbedding +logger = logging.getLogger(__name__) _model_cache: dict[str, "TextEmbedding"] = {} _model_lock = threading.Lock() +# BGE-small max sequence is 512 tokens (~2000 chars). Truncating long +# descriptions avoids wasting tokenisation and padding time on text that +# the model would discard anyway. +_MAX_TEXT_CHARS = 2000 + def _get_model(model_name: str) -> "TextEmbedding": cached = _model_cache.get(model_name) @@ -36,7 +43,9 @@ def _get_model(model_name: str) -> "TextEmbedding": if cached is not None: return cached from fastembed import TextEmbedding - model = TextEmbedding(model_name=model_name) + # threads=0 lets ONNX Runtime use all CPU cores for intra-op + # parallelism, which significantly speeds up batch inference. + model = TextEmbedding(model_name=model_name, threads=0) _model_cache[model_name] = model return model @@ -76,10 +85,42 @@ def embed_query(query: str, model_name: str = _DEFAULT_MODEL) -> list[float] | N return None +def _embed_node_list( + nodes: list[GraphNode], + graph: KnowledgeGraph, + model_name: str, + batch_size: int, +) -> list[NodeEmbedding]: + """Shared implementation for embedding a list of graph nodes. + + Generates text descriptions, truncates to model context window, + encodes via fastembed, and returns :class:`NodeEmbedding` objects. + """ + class_method_idx = build_class_method_index(graph) + + texts: list[str] = [] + valid_nodes: list[GraphNode] = [] + for node in nodes: + text = generate_text(node, graph, class_method_idx) + if text and text.strip(): + texts.append(text[:_MAX_TEXT_CHARS]) + valid_nodes.append(node) + + if not texts: + return [] + + logger.info("Embedding %d texts (batch_size=%d) …", len(texts), batch_size) + model = _get_model(model_name) + return [ + NodeEmbedding(node_id=node.id, embedding=vector.tolist()) + for node, vector in zip(valid_nodes, model.embed(texts, batch_size=batch_size)) + ] + + def embed_graph( graph: KnowledgeGraph, model_name: str = "BAAI/bge-small-en-v1.5", - batch_size: int = 64, + batch_size: int = 256, ) -> list[NodeEmbedding]: """Generate embeddings for all embeddable nodes in the graph. @@ -91,7 +132,7 @@ def embed_graph( graph: The knowledge graph whose nodes should be embedded. model_name: The fastembed model identifier. Defaults to ``"BAAI/bge-small-en-v1.5"``. - batch_size: Number of texts to encode per batch. Defaults to 64. + batch_size: Number of texts to encode per batch. Defaults to 256. Returns: A list of :class:`NodeEmbedding` instances, one per embeddable node, @@ -99,75 +140,22 @@ def embed_graph( Python ``list[float]``. """ all_nodes = [n for n in graph.iter_nodes() if n.label in EMBEDDABLE_LABELS] - if not all_nodes: return [] - - class_method_idx = build_class_method_index(graph) - - texts: list[str] = [] - nodes = [] - for node in all_nodes: - text = generate_text(node, graph, class_method_idx) - if text and text.strip(): - texts.append(text) - nodes.append(node) - - if not texts: - return [] - - model = _get_model(model_name) - vectors = list(model.embed(texts, batch_size=batch_size)) - - results: list[NodeEmbedding] = [] - for node, vector in zip(nodes, vectors): - results.append( - NodeEmbedding( - node_id=node.id, - embedding=vector.tolist(), - ) - ) - - return results + return _embed_node_list(all_nodes, graph, model_name, batch_size) def embed_nodes( graph: KnowledgeGraph, node_ids: set[str], model_name: str = "BAAI/bge-small-en-v1.5", - batch_size: int = 64, + batch_size: int = 256, ) -> list[NodeEmbedding]: """Like :func:`embed_graph`, but only for the given *node_ids*.""" if not node_ids: return [] - nodes = [graph.get_node(nid) for nid in node_ids] nodes = [n for n in nodes if n is not None and n.label in EMBEDDABLE_LABELS] - if not nodes: return [] - - class_method_idx = build_class_method_index(graph) - - texts: list[str] = [] - valid_nodes = [] - for node in nodes: - text = generate_text(node, graph, class_method_idx) - if text and text.strip(): - texts.append(text) - valid_nodes.append(node) - - if not texts: - return [] - - model = _get_model(model_name) - embeddings: list[NodeEmbedding] = [] - for node, vector in zip(valid_nodes, model.embed(texts, batch_size=batch_size)): - embeddings.append( - NodeEmbedding( - node_id=node.id, - embedding=vector.tolist(), - ) - ) - - return embeddings + return _embed_node_list(nodes, graph, model_name, batch_size) diff --git a/src/axon/core/ingestion/pipeline.py b/src/axon/core/ingestion/pipeline.py index 04436b14b..29732baf1 100644 --- a/src/axon/core/ingestion/pipeline.py +++ b/src/axon/core/ingestion/pipeline.py @@ -24,6 +24,7 @@ import time from collections.abc import Callable from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path @@ -120,72 +121,84 @@ def run_pipeline( """ start = time.monotonic() result = PipelineResult() + log = logging.getLogger(__name__) + + phase_times: dict[str, float] = {} def report(phase: str, pct: float) -> None: if progress_callback is not None: progress_callback(phase, pct) - report("Walking files", 0.0) - gitignore = load_gitignore(repo_path) - files = walk_repo(repo_path, gitignore) - result.files = len(files) - report("Walking files", 1.0) + @contextmanager + def _timed(phase_name: str): + """Context manager that logs and records phase wall-clock time.""" + report(phase_name, 0.0) + t0 = time.monotonic() + try: + yield + finally: + elapsed = time.monotonic() - t0 + phase_times[phase_name] = elapsed + log.info("Phase %-30s %.2fs", phase_name, elapsed) + report(phase_name, 1.0) + + with _timed("Walking files"): + gitignore = load_gitignore(repo_path) + files = walk_repo(repo_path, gitignore) + result.files = len(files) graph = KnowledgeGraph() - report("Processing structure", 0.0) - process_structure(files, graph) - report("Processing structure", 1.0) + with _timed("Processing structure"): + process_structure(files, graph) - report("Parsing code", 0.0) - parse_data = process_parsing(files, graph) - report("Parsing code", 1.0) + with _timed("Parsing code"): + parse_data = process_parsing(files, graph) - report("Resolving imports", 0.0) - process_imports(parse_data, graph, parallel=True) - report("Resolving imports", 1.0) + with _timed("Resolving imports"): + process_imports(parse_data, graph, parallel=True) - shared_labels = ( - NodeLabel.FUNCTION, NodeLabel.METHOD, NodeLabel.CLASS, - NodeLabel.INTERFACE, NodeLabel.TYPE_ALIAS, - ) - shared_name_index = build_name_index(graph, shared_labels) - heritage_labels = {NodeLabel.CLASS, NodeLabel.INTERFACE} - heritage_name_index: dict[str, list[str]] = {} - for name, ids in shared_name_index.items(): - filtered = [ - nid for nid in ids - if (n := graph.get_node(nid)) is not None and n.label in heritage_labels - ] - if filtered: - heritage_name_index[name] = filtered - - report("Resolving relationships", 0.0) - with ThreadPoolExecutor(max_workers=3) as pool: - calls_f = pool.submit( - process_calls, parse_data, graph, - name_index=shared_name_index, parallel=False, collect=True, - ) - heritage_f = pool.submit( - process_heritage, parse_data, graph, - name_index=heritage_name_index, parallel=False, collect=True, - ) - types_f = pool.submit( - process_types, parse_data, graph, - name_index=shared_name_index, parallel=False, collect=True, + with _timed("Building indexes"): + shared_labels = ( + NodeLabel.FUNCTION, NodeLabel.METHOD, NodeLabel.CLASS, + NodeLabel.INTERFACE, NodeLabel.TYPE_ALIAS, ) + shared_name_index = build_name_index(graph, shared_labels) + heritage_labels = {NodeLabel.CLASS, NodeLabel.INTERFACE} + heritage_name_index: dict[str, list[str]] = {} + for name, ids in shared_name_index.items(): + filtered = [ + nid for nid in ids + if (n := graph.get_node(nid)) is not None and n.label in heritage_labels + ] + if filtered: + heritage_name_index[name] = filtered + + with _timed("Resolving relationships"): + with ThreadPoolExecutor(max_workers=3) as pool: + calls_f = pool.submit( + process_calls, parse_data, graph, + name_index=shared_name_index, parallel=False, collect=True, + ) + heritage_f = pool.submit( + process_heritage, parse_data, graph, + name_index=heritage_name_index, parallel=False, collect=True, + ) + types_f = pool.submit( + process_types, parse_data, graph, + name_index=shared_name_index, parallel=False, collect=True, + ) - _write_collected_edges(calls_f.result() or [], graph) + _write_collected_edges(calls_f.result() or [], graph) - heritage_edges, heritage_patches = heritage_f.result() - _write_collected_edges(heritage_edges, graph) - for patch in heritage_patches: - node = graph.get_node(patch.node_id) - if node is not None: - node.properties[patch.key] = patch.value + heritage_edges, heritage_patches = heritage_f.result() + _write_collected_edges(heritage_edges, graph) + for patch in heritage_patches: + node = graph.get_node(patch.node_id) + if node is not None: + node.properties[patch.key] = patch.value - _write_collected_edges(types_f.result() or [], graph) - report("Resolving relationships", 1.0) + _write_collected_edges(types_f.result() or [], graph) coupling_file_nodes = graph.get_nodes_by_label(NodeLabel.FILE) @@ -194,41 +207,35 @@ def report(phase: str, pct: float) -> None: resolve_coupling, graph, repo_path, file_nodes=coupling_file_nodes, ) - report("Detecting communities", 0.0) - result.clusters = process_communities(graph) - report("Detecting communities", 1.0) + with _timed("Detecting communities"): + result.clusters = process_communities(graph) - report("Detecting execution flows", 0.0) - result.processes = process_processes(graph) - report("Detecting execution flows", 1.0) + with _timed("Detecting execution flows"): + result.processes = process_processes(graph) - report("Finding dead code", 0.0) - result.dead_code = process_dead_code(graph) - report("Finding dead code", 1.0) + with _timed("Finding dead code"): + result.dead_code = process_dead_code(graph) - report("Analyzing git history", 0.0) - coupling_edges = coupling_future.result() - _write_collected_edges(coupling_edges, graph) - result.coupled_pairs = len(coupling_edges) - report("Analyzing git history", 1.0) + with _timed("Analyzing git history"): + coupling_edges = coupling_future.result() + _write_collected_edges(coupling_edges, graph) + result.coupled_pairs = len(coupling_edges) result.symbols = sum(1 for n in graph.iter_nodes() if n.label in _SYMBOL_LABELS) result.relationships = graph.relationship_count if storage is not None: - report("Loading to storage", 0.0) - storage.bulk_load(graph) - report("Loading to storage", 1.0) + with _timed("Loading to storage"): + storage.bulk_load(graph) if embeddings: try: - report("Generating embeddings", 0.0) - node_embeddings = embed_graph(graph) - storage.store_embeddings(node_embeddings) - result.embeddings = len(node_embeddings) - report("Generating embeddings", 1.0) + with _timed("Generating embeddings"): + node_embeddings = embed_graph(graph) + storage.store_embeddings(node_embeddings) + result.embeddings = len(node_embeddings) except Exception: - logging.getLogger(__name__).warning( + log.warning( "Embedding phase failed — search will use FTS only", exc_info=True, ) @@ -236,6 +243,14 @@ def report(phase: str, pct: float) -> None: result.duration_seconds = time.monotonic() - start + # Log phase breakdown summary + if phase_times: + log.info("─── Phase timing breakdown ───") + for phase, elapsed in phase_times.items(): + pct = (elapsed / result.duration_seconds * 100) if result.duration_seconds > 0 else 0 + log.info(" %-30s %6.1fs (%4.1f%%)", phase, elapsed, pct) + log.info(" %-30s %6.1fs", "TOTAL", result.duration_seconds) + return graph, result def reindex_files( diff --git a/src/axon/core/ingestion/processes.py b/src/axon/core/ingestion/processes.py index 72b51cd63..2ed995b90 100644 --- a/src/axon/core/ingestion/processes.py +++ b/src/axon/core/ingestion/processes.py @@ -61,8 +61,7 @@ def _is_entry_point(node: GraphNode, graph: KnowledgeGraph) -> bool: if _matches_framework_pattern(node): return True - incoming_calls = graph.get_incoming(node.id, RelType.CALLS) - if incoming_calls: + if graph.has_incoming(node.id, RelType.CALLS): return False if node.is_exported: diff --git a/src/axon/core/storage/kuzu_backend.py b/src/axon/core/storage/kuzu_backend.py index 8c34905c2..ab97105ba 100644 --- a/src/axon/core/storage/kuzu_backend.py +++ b/src/axon/core/storage/kuzu_backend.py @@ -9,6 +9,7 @@ from __future__ import annotations import csv +import json import hashlib import json import logging @@ -124,6 +125,7 @@ def __init__(self) -> None: self._db: kuzu.Database | None = None self._conn: kuzu.Connection | None = None self._lock = threading.Lock() + self._embeddings_clean: bool = False def _require_conn(self) -> kuzu.Connection: if self._conn is None: @@ -884,6 +886,14 @@ def bulk_load(self, graph: KnowledgeGraph) -> None: except Exception: pass + # Wipe embeddings table too — avoids redundant per-batch DELETE + # queries inside _bulk_store_embeddings_csv later. + try: + conn.execute("MATCH (e:Embedding) DELETE e") + except Exception: + pass + self._embeddings_clean = True + if not self._bulk_load_nodes_csv(graph): self.add_nodes(list(graph.iter_nodes())) @@ -893,13 +903,13 @@ def bulk_load(self, graph: KnowledgeGraph) -> None: self.rebuild_fts_indexes() def rebuild_fts_indexes(self) -> None: - """Drop and recreate all FTS indexes. + """Drop and recreate FTS indexes on searchable tables only. - Must be called after any bulk data change so the BM25 indexes - reflect the current node contents. + Skips structural tables (Folder, Community, Process) that lack + meaningful content/signature fields — saves ~30% of FTS rebuild time. """ conn = self._require_conn() - for table in _NODE_TABLE_NAMES: + for table in _SEARCHABLE_TABLES: idx_name = f"{table.lower()}_fts" try: conn.execute(f"CALL DROP_FTS_INDEX('{table}', '{idx_name}')") @@ -1001,22 +1011,24 @@ def _bulk_store_embeddings_csv(self, embeddings: list[NodeEmbedding]) -> bool: """ conn = self._require_conn() try: - current_ids = [emb.node_id for emb in embeddings] - for i in range(0, len(current_ids), 500): - batch = current_ids[i:i + 500] - try: - conn.execute( - "MATCH (e:Embedding) WHERE e.node_id IN $ids DETACH DELETE e", - parameters={"ids": batch}, - ) - except Exception: - pass + # Skip DELETE if bulk_load already wiped the table + if not self._embeddings_clean: + current_ids = [emb.node_id for emb in embeddings] + for i in range(0, len(current_ids), 500): + batch = current_ids[i:i + 500] + try: + conn.execute( + "MATCH (e:Embedding) WHERE e.node_id IN $ids DETACH DELETE e", + parameters={"ids": batch}, + ) + except Exception: + pass self._csv_copy("Embedding", [ - [emb.node_id, - "[" + ",".join(str(v) for v in emb.embedding) + "]"] + [emb.node_id, json.dumps(emb.embedding)] for emb in embeddings ]) + self._embeddings_clean = False return True except Exception: logger.debug("CSV bulk_store_embeddings failed, falling back", exc_info=True) @@ -1062,9 +1074,9 @@ def _create_schema(self) -> None: self._create_fts_indexes() def _create_fts_indexes(self) -> None: - """Create FTS indexes for every node table (idempotent).""" + """Create FTS indexes for searchable node tables (idempotent).""" conn = self._require_conn() - for table in _NODE_TABLE_NAMES: + for table in _SEARCHABLE_TABLES: idx_name = f"{table.lower()}_fts" try: conn.execute( diff --git a/tests/core/test_embedder.py b/tests/core/test_embedder.py index a836267c6..f16479896 100644 --- a/tests/core/test_embedder.py +++ b/tests/core/test_embedder.py @@ -288,7 +288,7 @@ def test_default_model_name(self, mock_te_cls: MagicMock, sample_graph: Knowledg embed_graph(sample_graph) - mock_te_cls.assert_called_once_with(model_name="BAAI/bge-small-en-v1.5") + mock_te_cls.assert_called_once_with(model_name="BAAI/bge-small-en-v1.5", threads=0) @patch("fastembed.TextEmbedding") def test_custom_model_name(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: @@ -300,7 +300,7 @@ def test_custom_model_name(self, mock_te_cls: MagicMock, sample_graph: Knowledge embed_graph(sample_graph, model_name="BAAI/bge-base-en-v1.5") - mock_te_cls.assert_called_once_with(model_name="BAAI/bge-base-en-v1.5") + mock_te_cls.assert_called_once_with(model_name="BAAI/bge-base-en-v1.5", threads=0) @patch("fastembed.TextEmbedding") def test_custom_batch_size_passed_to_embed( @@ -394,7 +394,7 @@ def test_many_nodes_all_embedded(self, mock_te_cls: MagicMock) -> None: assert all(len(r.embedding) == 3 for r in results) @patch("fastembed.TextEmbedding") - def test_default_batch_size_is_64(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: + def test_default_batch_size_is_256(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: mock_model = MagicMock() mock_model.embed.return_value = iter( [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] @@ -404,8 +404,8 @@ def test_default_batch_size_is_64(self, mock_te_cls: MagicMock, sample_graph: Kn embed_graph(sample_graph) embed_call = mock_model.embed.call_args - assert embed_call.kwargs.get("batch_size") == 64 or ( - len(embed_call.args) > 1 and embed_call.args[1] == 64 + assert embed_call.kwargs.get("batch_size") == 256 or ( + len(embed_call.args) > 1 and embed_call.args[1] == 256 ) From e92fee073d42cb3276e3feb272306a5bb1dd118c Mon Sep 17 00:00:00 2001 From: harshkedia177 Date: Tue, 10 Mar 2026 14:54:13 +0530 Subject: [PATCH 2/2] feat(embeddings): add dimension validation to NodeEmbedding --- src/axon/core/storage/base.py | 8 +++ tests/core/test_embedder.py | 99 ++++++++++++++++++----------------- 2 files changed, 58 insertions(+), 49 deletions(-) diff --git a/src/axon/core/storage/base.py b/src/axon/core/storage/base.py index 847c9495f..3110e84c5 100644 --- a/src/axon/core/storage/base.py +++ b/src/axon/core/storage/base.py @@ -26,6 +26,8 @@ class SearchResult: label: str = "" snippet: str = "" +EMBEDDING_DIMENSIONS = 384 + @dataclass class NodeEmbedding: """An embedding vector associated with a graph node.""" @@ -33,6 +35,12 @@ class NodeEmbedding: node_id: str embedding: list[float] = field(default_factory=list) + def __post_init__(self) -> None: + if self.embedding and len(self.embedding) != EMBEDDING_DIMENSIONS: + raise ValueError( + f"Expected {EMBEDDING_DIMENSIONS}d embedding, got {len(self.embedding)}d" + ) + @runtime_checkable class StorageBackend(Protocol): """Protocol that every Axon storage backend must implement. diff --git a/tests/core/test_embedder.py b/tests/core/test_embedder.py index f16479896..53502f288 100644 --- a/tests/core/test_embedder.py +++ b/tests/core/test_embedder.py @@ -8,7 +8,14 @@ from axon.core.embeddings.embedder import EMBEDDABLE_LABELS, _get_model, embed_graph, embed_nodes from axon.core.graph.graph import KnowledgeGraph from axon.core.graph.model import GraphNode, GraphRelationship, NodeLabel, RelType, generate_id -from axon.core.storage.base import NodeEmbedding +from axon.core.storage.base import EMBEDDING_DIMENSIONS, NodeEmbedding + +# Helper to create 384d mock vectors with a distinguishable base value. +_D = EMBEDDING_DIMENSIONS + +def _mock_vec(base: float = 0.1) -> np.ndarray: + """Return a 384d numpy array filled with *base*.""" + return np.array([base] * _D) @pytest.fixture(autouse=True) @@ -128,9 +135,7 @@ class TestEmbedGraphBasic: @patch("fastembed.TextEmbedding") def test_returns_node_embeddings(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model results = embed_graph(sample_graph) @@ -143,9 +148,7 @@ def test_embedding_vectors_are_lists_of_float( self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph ) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model results = embed_graph(sample_graph) @@ -159,26 +162,22 @@ def test_embedding_values_match( self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph ) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model results = embed_graph(sample_graph) # We should get two results with the two mock vectors embeddings = [r.embedding for r in results] - assert [0.1, 0.2, 0.3] in embeddings or pytest.approx([0.1, 0.2, 0.3]) in embeddings - assert [0.4, 0.5, 0.6] in embeddings or pytest.approx([0.4, 0.5, 0.6]) in embeddings + assert [0.1] * _D in embeddings or pytest.approx([0.1] * _D) in embeddings + assert [0.4] * _D in embeddings or pytest.approx([0.4] * _D) in embeddings @patch("fastembed.TextEmbedding") def test_node_ids_are_correct( self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph ) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model results = embed_graph(sample_graph) @@ -194,9 +193,7 @@ def test_skips_folder_nodes( self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph ) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model results = embed_graph(sample_graph) @@ -211,7 +208,7 @@ def test_skips_community_and_process( embeddable_count = 7 # FILE, FUNCTION, CLASS, METHOD, INTERFACE, TYPE_ALIAS, ENUM mock_model = MagicMock() mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]) for _ in range(embeddable_count)] + [_mock_vec(0.1) for _ in range(embeddable_count)] ) mock_te_cls.return_value = mock_model @@ -230,7 +227,7 @@ def test_all_embeddable_labels_included( embeddable_count = 7 mock_model = MagicMock() mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]) for _ in range(embeddable_count)] + [_mock_vec(0.1) for _ in range(embeddable_count)] ) mock_te_cls.return_value = mock_model @@ -281,9 +278,7 @@ class TestEmbedGraphModelConfig: @patch("fastembed.TextEmbedding") def test_default_model_name(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model embed_graph(sample_graph) @@ -293,9 +288,7 @@ def test_default_model_name(self, mock_te_cls: MagicMock, sample_graph: Knowledg @patch("fastembed.TextEmbedding") def test_custom_model_name(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model embed_graph(sample_graph, model_name="BAAI/bge-base-en-v1.5") @@ -307,9 +300,7 @@ def test_custom_batch_size_passed_to_embed( self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph ) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model embed_graph(sample_graph, batch_size=32) @@ -332,9 +323,7 @@ def test_generate_text_called_for_each_node( ) -> None: mock_gen_text.return_value = "mock text" mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model embed_graph(sample_graph) @@ -352,9 +341,7 @@ def test_generated_texts_passed_to_model( ) -> None: mock_gen_text.side_effect = ["text for foo", "text for Bar"] mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model embed_graph(sample_graph) @@ -383,22 +370,20 @@ def test_many_nodes_all_embedded(self, mock_te_cls: MagicMock) -> None: mock_model = MagicMock() mock_model.embed.return_value = iter( - [np.array([float(i), float(i + 1), float(i + 2)]) for i in range(count)] + [_mock_vec(float(i)) for i in range(count)] ) mock_te_cls.return_value = mock_model results = embed_graph(graph, batch_size=16) assert len(results) == count - # Each embedding should have 3 dimensions - assert all(len(r.embedding) == 3 for r in results) + # Each embedding should have 384 dimensions + assert all(len(r.embedding) == _D for r in results) @patch("fastembed.TextEmbedding") def test_default_batch_size_is_256(self, mock_te_cls: MagicMock, sample_graph: KnowledgeGraph) -> None: mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model embed_graph(sample_graph) @@ -459,8 +444,8 @@ def test_embedding_alignment( ) # Two distinguishable embedding vectors. - embedding_a = np.array([1.0, 0.0, 0.0]) - embedding_b = np.array([0.0, 1.0, 0.0]) + embedding_a = _mock_vec(1.0) + embedding_b = _mock_vec(2.0) mock_model = MagicMock() mock_te_cls.return_value = mock_model @@ -484,7 +469,7 @@ def test_embeds_only_requested_nodes(self, mock_te_cls: MagicMock) -> None: node_ids = {generate_id(NodeLabel.FUNCTION, "src/a.py", "func_a")} mock_model = MagicMock() - mock_model.embed.return_value = iter([np.array([0.1, 0.2, 0.3])]) + mock_model.embed.return_value = iter([_mock_vec(0.1)]) mock_te_cls.return_value = mock_model results = embed_nodes(graph, node_ids) @@ -522,9 +507,7 @@ def test_embeds_both_requested_nodes(self, mock_te_cls: MagicMock) -> None: node_ids = {id_a, id_b} mock_model = MagicMock() - mock_model.embed.return_value = iter( - [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])] - ) + mock_model.embed.return_value = iter([_mock_vec(0.1), _mock_vec(0.4)]) mock_te_cls.return_value = mock_model results = embed_nodes(graph, node_ids) @@ -539,7 +522,7 @@ def test_embedding_values_are_correct(self, mock_te_cls: MagicMock) -> None: id_a = generate_id(NodeLabel.FUNCTION, "src/a.py", "func_a") mock_model = MagicMock() - mock_model.embed.return_value = iter([np.array([1.0, 2.0, 3.0])]) + mock_model.embed.return_value = iter([_mock_vec(1.0)]) mock_te_cls.return_value = mock_model results = embed_nodes(graph, {id_a}) @@ -547,4 +530,22 @@ def test_embedding_values_are_correct(self, mock_te_cls: MagicMock) -> None: assert len(results) == 1 assert results[0].node_id == id_a assert isinstance(results[0].embedding, list) - assert results[0].embedding == pytest.approx([1.0, 2.0, 3.0]) + assert results[0].embedding == pytest.approx([1.0] * _D) + + +class TestNodeEmbeddingValidation: + def test_valid_dimensions_accepted(self) -> None: + emb = NodeEmbedding(node_id="fn:a", embedding=[0.1] * 384) + assert len(emb.embedding) == 384 + + def test_wrong_dimensions_rejected(self) -> None: + with pytest.raises(ValueError, match="Expected 384"): + NodeEmbedding(node_id="fn:a", embedding=[0.1] * 768) + + def test_empty_embedding_accepted(self) -> None: + emb = NodeEmbedding(node_id="fn:a", embedding=[]) + assert emb.embedding == [] + + def test_default_factory_accepted(self) -> None: + emb = NodeEmbedding(node_id="fn:a") + assert emb.embedding == []