diff --git a/src/typeagent/storage/sqlite/schema.py b/src/typeagent/storage/sqlite/schema.py index 99117c24..549a66af 100644 --- a/src/typeagent/storage/sqlite/schema.py +++ b/src/typeagent/storage/sqlite/schema.py @@ -59,6 +59,7 @@ CREATE TABLE IF NOT EXISTS SemanticRefIndex ( term TEXT NOT NULL, -- lowercased, not-unique/normalized semref_id INTEGER NOT NULL, + score REAL NOT NULL DEFAULT 1.0, FOREIGN KEY (semref_id) REFERENCES SemanticRefs(semref_id) ON DELETE CASCADE ); @@ -278,6 +279,15 @@ def _set_conversation_metadata( ) +def _migrate_semantic_ref_index_score(cursor: sqlite3.Cursor) -> None: + """Add the score column to SemanticRefIndex tables created before it existed.""" + columns = [row[1] for row in cursor.execute("PRAGMA table_info(SemanticRefIndex)")] + if "score" not in columns: + cursor.execute( + "ALTER TABLE SemanticRefIndex ADD COLUMN score REAL NOT NULL DEFAULT 1.0" + ) + + def init_db_schema(db: sqlite3.Connection) -> None: """Initialize the database schema with all required tables.""" cursor = db.cursor() @@ -287,6 +297,7 @@ def init_db_schema(db: sqlite3.Connection) -> None: cursor.execute(MESSAGES_SCHEMA) cursor.execute(SEMANTIC_REFS_SCHEMA) cursor.execute(SEMANTIC_REF_INDEX_SCHEMA) + _migrate_semantic_ref_index_score(cursor) cursor.execute(MESSAGE_TEXT_INDEX_SCHEMA) cursor.execute(PROPERTY_INDEX_SCHEMA) cursor.execute(RELATED_TERMS_ALIASES_SCHEMA) diff --git a/src/typeagent/storage/sqlite/semrefindex.py b/src/typeagent/storage/sqlite/semrefindex.py index f462955e..125a374a 100644 --- a/src/typeagent/storage/sqlite/semrefindex.py +++ b/src/typeagent/storage/sqlite/semrefindex.py @@ -18,6 +18,14 @@ ) +def _split_ordinal( + ordinal: SemanticRefOrdinal | ScoredSemanticRefOrdinal, +) -> tuple[SemanticRefOrdinal, float]: + if isinstance(ordinal, ScoredSemanticRefOrdinal): + return ordinal.semantic_ref_ordinal, ordinal.score + return ordinal, 1.0 + + class SqliteTermToSemanticRefIndex(ITermToSemanticRefIndex): """SQLite-backed implementation of term to semantic ref index.""" @@ -44,19 +52,15 @@ async def add_term( term = self._prepare_term(term) - # Extract semref_id from the ordinal - if isinstance(semantic_ref_ordinal, ScoredSemanticRefOrdinal): - semref_id = semantic_ref_ordinal.semantic_ref_ordinal - else: - semref_id = semantic_ref_ordinal + semref_id, score = _split_ordinal(semantic_ref_ordinal) cursor = self.db.cursor() cursor.execute( """ - INSERT OR IGNORE INTO SemanticRefIndex (term, semref_id) - VALUES (?, ?) + INSERT OR IGNORE INTO SemanticRefIndex (term, semref_id, score) + VALUES (?, ?, ?) """, - (term, semref_id), + (term, semref_id, score), ) return term @@ -72,15 +76,12 @@ async def add_terms_batch( if not term: continue term = self._prepare_term(term) - if isinstance(ordinal, ScoredSemanticRefOrdinal): - semref_id = ordinal.semantic_ref_ordinal - else: - semref_id = ordinal - rows.append((term, semref_id)) + semref_id, score = _split_ordinal(ordinal) + rows.append((term, semref_id, score)) if rows: cursor = self.db.cursor() cursor.executemany( - "INSERT OR IGNORE INTO SemanticRefIndex (term, semref_id) VALUES (?, ?)", + "INSERT OR IGNORE INTO SemanticRefIndex (term, semref_id, score) VALUES (?, ?, ?)", rows, ) @@ -98,16 +99,13 @@ async def lookup_term(self, term: str) -> list[ScoredSemanticRefOrdinal] | None: term = self._prepare_term(term) cursor = self.db.cursor() cursor.execute( - "SELECT semref_id FROM SemanticRefIndex WHERE term = ?", + "SELECT semref_id, score FROM SemanticRefIndex WHERE term = ? ORDER BY rowid", (term,), ) - - # Return as ScoredSemanticRefOrdinal with default score of 1.0 - results = [] - for row in cursor.fetchall(): - semref_id = row[0] - results.append(ScoredSemanticRefOrdinal(semref_id, 1.0)) - return results + return [ + ScoredSemanticRefOrdinal(semref_id, score) + for semref_id, score in cursor.fetchall() + ] async def clear(self) -> None: """Clear all terms from the semantic ref index.""" @@ -118,15 +116,16 @@ async def serialize(self) -> TermToSemanticRefIndexData: """Serialize the index data for compatibility with in-memory version.""" cursor = self.db.cursor() cursor.execute( - "SELECT term, semref_id FROM SemanticRefIndex ORDER BY term, semref_id" + "SELECT term, semref_id, score FROM SemanticRefIndex " + "ORDER BY term, semref_id, rowid" ) # Group by term term_to_semrefs: dict[str, list[ScoredSemanticRefOrdinalData]] = {} - for term, semref_id in cursor.fetchall(): + for term, semref_id, score in cursor.fetchall(): if term not in term_to_semrefs: term_to_semrefs[term] = [] - scored_ref = ScoredSemanticRefOrdinal(semref_id, 1.0) + scored_ref = ScoredSemanticRefOrdinal(semref_id, score) term_to_semrefs[term].append(scored_ref.serialize()) # Convert to the expected format @@ -155,15 +154,17 @@ async def deserialize(self, data: TermToSemanticRefIndexData) -> None: for semref_ordinal_data in item["semanticRefOrdinals"]: if isinstance(semref_ordinal_data, dict): semref_id = semref_ordinal_data["semanticRefOrdinal"] + score = semref_ordinal_data.get("score", 1.0) else: # Fallback for direct integer semref_id = semref_ordinal_data - insertion_data.append((term, semref_id)) + score = 1.0 + insertion_data.append((term, semref_id, score)) # Bulk insert all the data if insertion_data: cursor.executemany( - "INSERT OR IGNORE INTO SemanticRefIndex (term, semref_id) VALUES (?, ?)", + "INSERT OR IGNORE INTO SemanticRefIndex (term, semref_id, score) VALUES (?, ?, ?)", insertion_data, ) diff --git a/tests/test_semrefindex.py b/tests/test_semrefindex.py index f12de683..87006288 100644 --- a/tests/test_semrefindex.py +++ b/tests/test_semrefindex.py @@ -18,6 +18,7 @@ IMessage, ISemanticRefCollection, ITermToSemanticRefIndex, + ScoredSemanticRefOrdinal, Topic, ) from typeagent.knowpro.knowledge_schema import ( @@ -414,3 +415,75 @@ async def test_semantic_ref_index_serialize_empty( serialized = await legacy_semantic_ref_index.serialize() assert "items" in serialized assert serialized["items"] == [] + + +@pytest.mark.asyncio +async def test_semantic_ref_index_preserves_scores( + semantic_ref_index: ITermToSemanticRefIndex, needs_auth: None +) -> None: + """Scores must round-trip identically on both backends (issue #321).""" + await semantic_ref_index.add_term("cat", ScoredSemanticRefOrdinal(1, 0.25)) + await semantic_ref_index.add_term("cat", 2) # bare ordinal -> 1.0 + await semantic_ref_index.add_terms_batch( + [("dog", ScoredSemanticRefOrdinal(3, 0.5)), ("dog", 1)] + ) + + cat = await semantic_ref_index.lookup_term("cat") + assert cat is not None + assert [(r.semantic_ref_ordinal, r.score) for r in cat] == [(1, 0.25), (2, 1.0)] + dog = await semantic_ref_index.lookup_term("dog") + assert dog is not None + assert [(r.semantic_ref_ordinal, r.score) for r in dog] == [(3, 0.5), (1, 1.0)] + + data = await semantic_ref_index.serialize() + scores = { + item["term"]: { + o["semanticRefOrdinal"]: o["score"] for o in item["semanticRefOrdinals"] + } + for item in data["items"] + } + assert scores == {"cat": {1: 0.25, 2: 1.0}, "dog": {3: 0.5, 1: 1.0}} + + +@pytest.mark.asyncio +async def test_semantic_ref_index_deserialize_preserves_scores( + semantic_ref_index: ITermToSemanticRefIndex, needs_auth: None +) -> None: + source = TermToSemanticRefIndex() + await source.add_term("cat", ScoredSemanticRefOrdinal(1, 0.25)) + await semantic_ref_index.deserialize(await source.serialize()) + + cat = await semantic_ref_index.lookup_term("cat") + assert cat is not None + assert [(r.semantic_ref_ordinal, r.score) for r in cat] == [(1, 0.25)] + + +@pytest.mark.asyncio +async def test_semantic_ref_index_duplicate_pairs_match_across_backends( + semantic_ref_index: ITermToSemanticRefIndex, needs_auth: None +) -> None: + """Repeated (term, semref) pairs behave the same on both backends.""" + await semantic_ref_index.add_term("cat", 1) + await semantic_ref_index.add_term("cat", 1) + cat = await semantic_ref_index.lookup_term("cat") + assert cat is not None + assert len(cat) == 2 + + +def test_init_db_schema_adds_score_column_to_legacy_semref_index() -> None: + """DBs created before the score column existed are migrated in place.""" + import sqlite3 + + from typeagent.storage.sqlite.schema import init_db_schema + + db = sqlite3.connect(":memory:") + db.execute( + "CREATE TABLE SemanticRefIndex (term TEXT NOT NULL, semref_id INTEGER NOT NULL)" + ) + db.execute("INSERT INTO SemanticRefIndex VALUES ('cat', 1)") + init_db_schema(db) + assert db.execute( + "SELECT term, semref_id, score FROM SemanticRefIndex" + ).fetchall() == [("cat", 1, 1.0)] + init_db_schema(db) # idempotent + db.close()