diff --git a/src/haiku/rag/store/engine.py b/src/haiku/rag/store/engine.py index 180eaa19..449b702b 100644 --- a/src/haiku/rag/store/engine.py +++ b/src/haiku/rag/store/engine.py @@ -24,6 +24,7 @@ class Store: CREATE TABLE IF NOT EXISTS documents ( id INTEGER PRIMARY KEY AUTOINCREMENT, content TEXT NOT NULL, + uri TEXT, metadata TEXT DEFAULT '{}', created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP diff --git a/src/haiku/rag/store/models/document.py b/src/haiku/rag/store/models/document.py index 00878212..89539c65 100644 --- a/src/haiku/rag/store/models/document.py +++ b/src/haiku/rag/store/models/document.py @@ -10,6 +10,7 @@ class Document(BaseModel): id: int | None = None content: str + uri: str | None = None metadata: dict = {} created_at: datetime = Field(default_factory=datetime.now) updated_at: datetime = Field(default_factory=datetime.now) diff --git a/src/haiku/rag/store/repositories/document.py b/src/haiku/rag/store/repositories/document.py index 027b76b7..83cb183d 100644 --- a/src/haiku/rag/store/repositories/document.py +++ b/src/haiku/rag/store/repositories/document.py @@ -30,11 +30,12 @@ class DocumentRepository(BaseRepository[Document]): # Insert the document cursor.execute( """ - INSERT INTO documents (content, metadata, created_at, updated_at) - VALUES (?, ?, ?, ?) + INSERT INTO documents (content, uri, metadata, created_at, updated_at) + VALUES (?, ?, ?, ?, ?) """, ( entity.content, + entity.uri, json.dumps(entity.metadata), entity.created_at, entity.updated_at, @@ -65,7 +66,7 @@ class DocumentRepository(BaseRepository[Document]): cursor = self.store._connection.cursor() cursor.execute( """ - SELECT id, content, metadata, created_at, updated_at + SELECT id, content, uri, metadata, created_at, updated_at FROM documents WHERE id = ? """, (entity_id,), @@ -75,12 +76,43 @@ class DocumentRepository(BaseRepository[Document]): if row is None: return None - document_id, content, metadata_json, created_at, updated_at = row + document_id, content, uri, metadata_json, created_at, updated_at = row metadata = json.loads(metadata_json) if metadata_json else {} return Document( id=document_id, content=content, + uri=uri, + metadata=metadata, + created_at=created_at, + updated_at=updated_at, + ) + + async def get_by_uri(self, uri: str) -> Document | None: + """Get a document by its URI.""" + if self.store._connection is None: + raise ValueError("Store connection is not available") + + cursor = self.store._connection.cursor() + cursor.execute( + """ + SELECT id, content, uri, metadata, created_at, updated_at + FROM documents WHERE uri = ? + """, + (uri,), + ) + + row = cursor.fetchone() + if row is None: + return None + + document_id, content, uri, metadata_json, created_at, updated_at = row + metadata = json.loads(metadata_json) if metadata_json else {} + + return Document( + id=document_id, + content=content, + uri=uri, metadata=metadata, created_at=created_at, updated_at=updated_at, @@ -103,11 +135,12 @@ class DocumentRepository(BaseRepository[Document]): cursor.execute( """ UPDATE documents - SET content = ?, metadata = ?, updated_at = ? + SET content = ?, uri = ?, metadata = ?, updated_at = ? WHERE id = ? """, ( entity.content, + entity.uri, json.dumps(entity.metadata), entity.updated_at, entity.id, @@ -150,7 +183,7 @@ class DocumentRepository(BaseRepository[Document]): raise ValueError("Store connection is not available") cursor = self.store._connection.cursor() - query = "SELECT id, content, metadata, created_at, updated_at FROM documents ORDER BY created_at DESC" + query = "SELECT id, content, uri, metadata, created_at, updated_at FROM documents ORDER BY created_at DESC" params = [] if limit is not None: @@ -166,12 +199,13 @@ class DocumentRepository(BaseRepository[Document]): documents = [] for row in rows: - document_id, content, metadata_json, created_at, updated_at = row + document_id, content, uri, metadata_json, created_at, updated_at = row metadata = json.loads(metadata_json) if metadata_json else {} documents.append( Document( id=document_id, content=content, + uri=uri, metadata=metadata, created_at=created_at, updated_at=updated_at, diff --git a/tests/test_document.py b/tests/test_document.py index c6bc5485..b1b3210d 100644 --- a/tests/test_document.py +++ b/tests/test_document.py @@ -12,52 +12,61 @@ async def test_create_document_with_chunks(qa_corpus: Dataset): # Create an in-memory store and repository store = Store(":memory:") doc_repo = DocumentRepository(store) - + # Get the first document from the corpus first_doc = qa_corpus[0] document_text = first_doc["document_extracted"] - + # Create a Document instance document = Document( content=document_text, - metadata={"source": "qa_corpus", "topic": first_doc.get("document_topic", "")} + metadata={"source": "qa_corpus", "topic": first_doc.get("document_topic", "")}, ) - + # Create the document with chunks in the database created_document = await doc_repo.create(document) - + # Verify the document was created assert created_document.id is not None assert created_document.content == document_text - + # Check that chunks were created in the database if store._connection is not None: cursor = store._connection.cursor() - cursor.execute("SELECT COUNT(*) FROM chunks WHERE document_id = ?", (created_document.id,)) + cursor.execute( + "SELECT COUNT(*) FROM chunks WHERE document_id = ?", (created_document.id,) + ) chunk_count = cursor.fetchone()[0] - + assert chunk_count > 0 - + # Check that embeddings were created - cursor.execute(""" + cursor.execute( + """ SELECT COUNT(*) FROM chunk_embeddings ce JOIN chunks c ON c.id = ce.chunk_id WHERE c.document_id = ? - """, (created_document.id,)) + """, + (created_document.id,), + ) embedding_count = cursor.fetchone()[0] - + assert embedding_count == chunk_count - + # Verify chunk metadata contains order information - cursor.execute("SELECT metadata FROM chunks WHERE document_id = ? ORDER BY id", (created_document.id,)) + cursor.execute( + "SELECT metadata FROM chunks WHERE document_id = ? ORDER BY id", + (created_document.id,), + ) chunk_metadata = cursor.fetchall() - + for i, (metadata_json,) in enumerate(chunk_metadata): import json + metadata = json.loads(metadata_json) assert "order" in metadata assert metadata["order"] == i - + store.close() @@ -67,40 +76,52 @@ async def test_document_repository_crud(qa_corpus: Dataset): # Create an in-memory store and repository store = Store(":memory:") doc_repo = DocumentRepository(store) - + # Get the first document from the corpus first_doc = qa_corpus[0] document_text = first_doc["document_extracted"] - - # Create a document + + # Create a document with URI + test_uri = "file:///path/to/test.txt" document = Document( - content=document_text, - metadata={"source": "test"} + content=document_text, uri=test_uri, metadata={"source": "test"} ) created_document = await doc_repo.create(document) - + # Test get_by_id assert created_document.id is not None retrieved_document = await doc_repo.get_by_id(created_document.id) assert retrieved_document is not None assert retrieved_document.content == document_text - + assert retrieved_document.uri == test_uri + + # Test get_by_uri + retrieved_by_uri = await doc_repo.get_by_uri(test_uri) + assert retrieved_by_uri is not None + assert retrieved_by_uri.id == created_document.id + assert retrieved_by_uri.content == document_text + assert retrieved_by_uri.uri == test_uri + + # Test get_by_uri with non-existent URI + non_existent = await doc_repo.get_by_uri("file:///non/existent.txt") + assert non_existent is None + # Test update (should regenerate chunks) retrieved_document.content = "Updated content for testing" updated_document = await doc_repo.update(retrieved_document) assert updated_document.content == "Updated content for testing" - + # Test list_all all_documents = await doc_repo.list_all() assert len(all_documents) == 1 assert all_documents[0].id == created_document.id - + # Test delete deleted = await doc_repo.delete(created_document.id) assert deleted is True - + # Verify document is gone retrieved_document = await doc_repo.get_by_id(created_document.id) assert retrieved_document is None - - store.close() \ No newline at end of file + + store.close()