Set/get haiku version to db
This commit is contained in:
parent
58a04db556
commit
b05e5f5408
7 changed files with 46 additions and 108 deletions
|
|
@ -7,7 +7,7 @@ You can perform your own evaluations using as example the script found at
|
||||||
|
|
||||||
## Recall
|
## Recall
|
||||||
|
|
||||||
In order to calculate recall, we load the `News Stories` from `repliqa_3` which is 1035 documents and index them in a sqlite db. Subsequently, we run a search over the `question` field for each row of the dataset and check whether we match the document that answers the question.
|
In order to calculate recall, we load the `News Stories` from `repliqa_3` which is 1035 documents and index them. Subsequently, we run a search over the `question` field for each row of the dataset and check whether we match the document that answers the question.
|
||||||
|
|
||||||
|
|
||||||
The recall obtained is ~0.73 for matching in the top result, raising to ~0.75 for the top 3 results.
|
The recall obtained is ~0.73 for matching in the top result, raising to ~0.75 for the top 3 results.
|
||||||
|
|
|
||||||
|
|
@ -60,9 +60,6 @@ class Store:
|
||||||
settings_repo = SettingsRepository(self)
|
settings_repo = SettingsRepository(self)
|
||||||
settings_repo.validate_config_compatibility()
|
settings_repo.validate_config_compatibility()
|
||||||
|
|
||||||
current_version = metadata.version("haiku.rag")
|
|
||||||
self.set_user_version(current_version)
|
|
||||||
|
|
||||||
def create_or_update_db(self):
|
def create_or_update_db(self):
|
||||||
"""Create the database tables."""
|
"""Create the database tables."""
|
||||||
|
|
||||||
|
|
@ -98,6 +95,10 @@ class Store:
|
||||||
[SettingsRecord(id="settings", settings=json.dumps(settings_data))]
|
[SettingsRecord(id="settings", settings=json.dumps(settings_data))]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Set current version in settings
|
||||||
|
current_version = metadata.version("haiku.rag")
|
||||||
|
self.set_haiku_version(current_version)
|
||||||
|
|
||||||
# Check if we need to perform upgrades
|
# Check if we need to perform upgrades
|
||||||
try:
|
try:
|
||||||
existing_settings = list(
|
existing_settings = list(
|
||||||
|
|
@ -105,7 +106,7 @@ class Store:
|
||||||
)
|
)
|
||||||
if existing_settings:
|
if existing_settings:
|
||||||
console = Console()
|
console = Console()
|
||||||
db_version = self.get_user_version()
|
db_version = self.get_haiku_version()
|
||||||
# Future: Add upgrade logic here similar to SQLite version
|
# Future: Add upgrade logic here similar to SQLite version
|
||||||
console.print(
|
console.print(
|
||||||
f"[green]LanceDB store initialized (version: {db_version})[/green]"
|
f"[green]LanceDB store initialized (version: {db_version})[/green]"
|
||||||
|
|
@ -113,7 +114,7 @@ class Store:
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def get_user_version(self) -> str:
|
def get_haiku_version(self) -> str:
|
||||||
"""Returns the user version stored in settings."""
|
"""Returns the user version stored in settings."""
|
||||||
try:
|
try:
|
||||||
settings_records = list(
|
settings_records = list(
|
||||||
|
|
@ -130,7 +131,7 @@ class Store:
|
||||||
pass
|
pass
|
||||||
return "0.0.0"
|
return "0.0.0"
|
||||||
|
|
||||||
def set_user_version(self, version: str) -> None:
|
def set_haiku_version(self, version: str) -> None:
|
||||||
"""Updates the user version in settings."""
|
"""Updates the user version in settings."""
|
||||||
try:
|
try:
|
||||||
settings_records = list(
|
settings_records = list(
|
||||||
|
|
@ -145,7 +146,7 @@ class Store:
|
||||||
settings["version"] = version
|
settings["version"] = version
|
||||||
# Update the record
|
# Update the record
|
||||||
self.settings_table.update(
|
self.settings_table.update(
|
||||||
where="id = 1", values={"settings": json.dumps(settings)}
|
where="id = 'settings'", values={"settings": json.dumps(settings)}
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Create new settings record
|
# Create new settings record
|
||||||
|
|
|
||||||
|
|
@ -140,29 +140,23 @@ class ChunkRepository:
|
||||||
|
|
||||||
return created_chunks
|
return created_chunks
|
||||||
|
|
||||||
async def delete_all(self) -> bool:
|
async def delete_all(self) -> None:
|
||||||
"""Delete all chunks from the database."""
|
"""Delete all chunks from the database."""
|
||||||
try:
|
count = len(
|
||||||
count = len(
|
list(
|
||||||
list(
|
self.store.chunks_table.search()
|
||||||
self.store.chunks_table.search()
|
.limit(1)
|
||||||
.limit(1)
|
.to_pydantic(self.store.ChunkRecord)
|
||||||
.to_pydantic(self.store.ChunkRecord)
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
if count > 0:
|
)
|
||||||
# Drop and recreate table to clear all data
|
if count > 0:
|
||||||
self.store.db.drop_table("chunks")
|
# Drop and recreate table to clear all data
|
||||||
self.store.chunks_table = self.store.db.create_table(
|
self.store.db.drop_table("chunks")
|
||||||
"chunks", schema=self.store.ChunkRecord
|
self.store.chunks_table = self.store.db.create_table(
|
||||||
)
|
"chunks", schema=self.store.ChunkRecord
|
||||||
# Create FTS index on the new table
|
)
|
||||||
self.store.chunks_table.create_fts_index("content", replace=True)
|
# Create FTS index on the new table
|
||||||
|
self.store.chunks_table.create_fts_index("content", replace=True)
|
||||||
return True
|
|
||||||
return False
|
|
||||||
except Exception:
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def delete_by_document_id(self, document_id: str) -> bool:
|
async def delete_by_document_id(self, document_id: str) -> bool:
|
||||||
"""Delete all chunks for a document."""
|
"""Delete all chunks for a document."""
|
||||||
|
|
|
||||||
|
|
@ -15,13 +15,12 @@ if TYPE_CHECKING:
|
||||||
class DocumentRepository:
|
class DocumentRepository:
|
||||||
"""Repository for Document operations."""
|
"""Repository for Document operations."""
|
||||||
|
|
||||||
def __init__(self, store: Store, chunk_repository=None) -> None:
|
def __init__(self, store: Store) -> None:
|
||||||
self.store = store
|
self.store = store
|
||||||
# Avoid circular import by using late import if not provided
|
|
||||||
if chunk_repository is None:
|
|
||||||
from haiku.rag.store.repositories.chunk import ChunkRepository
|
|
||||||
|
|
||||||
chunk_repository = ChunkRepository(store)
|
from haiku.rag.store.repositories.chunk import ChunkRepository
|
||||||
|
|
||||||
|
chunk_repository = ChunkRepository(store)
|
||||||
self.chunk_repository = chunk_repository
|
self.chunk_repository = chunk_repository
|
||||||
|
|
||||||
async def create(self, entity: Document) -> Document:
|
async def create(self, entity: Document) -> Document:
|
||||||
|
|
@ -169,33 +168,26 @@ class DocumentRepository:
|
||||||
else datetime.now(),
|
else datetime.now(),
|
||||||
)
|
)
|
||||||
|
|
||||||
async def delete_all(self) -> bool:
|
async def delete_all(self) -> None:
|
||||||
"""Delete all documents from the database."""
|
"""Delete all documents from the database."""
|
||||||
try:
|
# Delete all chunks first
|
||||||
# Delete all chunks first
|
from haiku.rag.store.repositories.chunk import ChunkRepository
|
||||||
from haiku.rag.store.repositories.chunk import ChunkRepository
|
|
||||||
|
|
||||||
chunk_repo = ChunkRepository(self.store)
|
chunk_repo = ChunkRepository(self.store)
|
||||||
await chunk_repo.delete_all()
|
await chunk_repo.delete_all()
|
||||||
|
|
||||||
# Get count before deletion
|
# Get count before deletion
|
||||||
count = len(
|
count = len(
|
||||||
list(
|
list(
|
||||||
self.store.documents_table.search()
|
self.store.documents_table.search().limit(1).to_pydantic(DocumentRecord)
|
||||||
.limit(1)
|
)
|
||||||
.to_pydantic(DocumentRecord)
|
)
|
||||||
)
|
if count > 0:
|
||||||
|
# Drop and recreate table to clear all data
|
||||||
|
self.store.db.drop_table("documents")
|
||||||
|
self.store.documents_table = self.store.db.create_table(
|
||||||
|
"documents", schema=DocumentRecord
|
||||||
)
|
)
|
||||||
if count > 0:
|
|
||||||
# Drop and recreate table to clear all data
|
|
||||||
self.store.db.drop_table("documents")
|
|
||||||
self.store.documents_table = self.store.db.create_table(
|
|
||||||
"documents", schema=DocumentRecord
|
|
||||||
)
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
except Exception:
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def _create_with_docling(
|
async def _create_with_docling(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -32,37 +32,6 @@ def get_default_data_dir() -> Path:
|
||||||
return data_path
|
return data_path
|
||||||
|
|
||||||
|
|
||||||
def semantic_version_to_int(version: str) -> int:
|
|
||||||
"""Convert a semantic version string to an integer.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
version: Semantic version string.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Integer representation of semantic version.
|
|
||||||
"""
|
|
||||||
major, minor, patch = version.split(".")
|
|
||||||
major = int(major) << 16
|
|
||||||
minor = int(minor) << 8
|
|
||||||
patch = int(patch)
|
|
||||||
return major + minor + patch
|
|
||||||
|
|
||||||
|
|
||||||
def int_to_semantic_version(version: int) -> str:
|
|
||||||
"""Convert an integer to a semantic version string.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
version: Integer representation of semantic version.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Semantic version string.
|
|
||||||
"""
|
|
||||||
major = version >> 16
|
|
||||||
minor = (version >> 8) & 255
|
|
||||||
patch = version & 255
|
|
||||||
return f"{major}.{minor}.{patch}"
|
|
||||||
|
|
||||||
|
|
||||||
async def is_up_to_date() -> tuple[bool, Version, Version]:
|
async def is_up_to_date() -> tuple[bool, Version, Version]:
|
||||||
"""Check whether haiku.rag is current.
|
"""Check whether haiku.rag is current.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -144,7 +144,7 @@ async def main():
|
||||||
await populate_db()
|
await populate_db()
|
||||||
|
|
||||||
console.print("Running retrieval benchmarks...", style="bold blue")
|
console.print("Running retrieval benchmarks...", style="bold blue")
|
||||||
await run_match_benchmark()
|
# await run_match_benchmark()
|
||||||
|
|
||||||
console.print("\nRunning QA benchmarks...", style="bold yellow")
|
console.print("\nRunning QA benchmarks...", style="bold yellow")
|
||||||
await run_qa_benchmark()
|
await run_qa_benchmark()
|
||||||
|
|
|
||||||
|
|
@ -1,22 +1,4 @@
|
||||||
from haiku.rag.utils import (
|
from haiku.rag.utils import text_to_docling_document
|
||||||
int_to_semantic_version,
|
|
||||||
semantic_version_to_int,
|
|
||||||
text_to_docling_document,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_sqlite_user_version():
|
|
||||||
version = "0.1.5"
|
|
||||||
assert semantic_version_to_int(version) == 261
|
|
||||||
assert int_to_semantic_version(261) == version
|
|
||||||
|
|
||||||
version = "0.0.0"
|
|
||||||
assert semantic_version_to_int(version) == 0
|
|
||||||
assert int_to_semantic_version(0) == version
|
|
||||||
|
|
||||||
version = "255.255.255"
|
|
||||||
assert semantic_version_to_int(version) == 16777215
|
|
||||||
assert int_to_semantic_version(16777215) == version
|
|
||||||
|
|
||||||
|
|
||||||
def test_text_to_docling_document():
|
def test_text_to_docling_document():
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue