Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions src/typeagent/storage/sqlite/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Won't defaulting the score to 1.0 might hide more valid results? Should this be 0 or some really small value?


FOREIGN KEY (semref_id) REFERENCES SemanticRefs(semref_id) ON DELETE CASCADE
);
Expand Down Expand Up @@ -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"
)
Comment on lines +284 to +288

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

unlikely but seems like a simple enough change



def init_db_schema(db: sqlite3.Connection) -> None:
"""Initialize the database schema with all required tables."""
cursor = db.cursor()
Expand All @@ -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)
Expand Down
55 changes: 28 additions & 27 deletions src/typeagent/storage/sqlite/semrefindex.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Maybe make this a constant and reuse this in/from schema.py when initializing the SemanticRefIndex table.



class SqliteTermToSemanticRefIndex(ITermToSemanticRefIndex):
"""SQLite-backed implementation of term to semantic ref index."""

Expand All @@ -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
Expand All @@ -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,
)

Expand All @@ -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."""
Expand All @@ -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
Expand Down Expand Up @@ -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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same const reuse here and on line #157

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,
)

Expand Down
73 changes: 73 additions & 0 deletions tests/test_semrefindex.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
IMessage,
ISemanticRefCollection,
ITermToSemanticRefIndex,
ScoredSemanticRefOrdinal,
Topic,
)
from typeagent.knowpro.knowledge_schema import (
Expand Down Expand Up @@ -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
Comment on lines +475 to +477

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()
Loading