Add uri to Documents
This commit is contained in:
parent
e24d3cdb0a
commit
1ff17c2f50
4 changed files with 92 additions and 35 deletions
|
|
@ -24,6 +24,7 @@ class Store:
|
||||||
CREATE TABLE IF NOT EXISTS documents (
|
CREATE TABLE IF NOT EXISTS documents (
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
content TEXT NOT NULL,
|
content TEXT NOT NULL,
|
||||||
|
uri TEXT,
|
||||||
metadata TEXT DEFAULT '{}',
|
metadata TEXT DEFAULT '{}',
|
||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ class Document(BaseModel):
|
||||||
|
|
||||||
id: int | None = None
|
id: int | None = None
|
||||||
content: str
|
content: str
|
||||||
|
uri: str | None = None
|
||||||
metadata: dict = {}
|
metadata: dict = {}
|
||||||
created_at: datetime = Field(default_factory=datetime.now)
|
created_at: datetime = Field(default_factory=datetime.now)
|
||||||
updated_at: datetime = Field(default_factory=datetime.now)
|
updated_at: datetime = Field(default_factory=datetime.now)
|
||||||
|
|
|
||||||
|
|
@ -30,11 +30,12 @@ class DocumentRepository(BaseRepository[Document]):
|
||||||
# Insert the document
|
# Insert the document
|
||||||
cursor.execute(
|
cursor.execute(
|
||||||
"""
|
"""
|
||||||
INSERT INTO documents (content, metadata, created_at, updated_at)
|
INSERT INTO documents (content, uri, metadata, created_at, updated_at)
|
||||||
VALUES (?, ?, ?, ?)
|
VALUES (?, ?, ?, ?, ?)
|
||||||
""",
|
""",
|
||||||
(
|
(
|
||||||
entity.content,
|
entity.content,
|
||||||
|
entity.uri,
|
||||||
json.dumps(entity.metadata),
|
json.dumps(entity.metadata),
|
||||||
entity.created_at,
|
entity.created_at,
|
||||||
entity.updated_at,
|
entity.updated_at,
|
||||||
|
|
@ -65,7 +66,7 @@ class DocumentRepository(BaseRepository[Document]):
|
||||||
cursor = self.store._connection.cursor()
|
cursor = self.store._connection.cursor()
|
||||||
cursor.execute(
|
cursor.execute(
|
||||||
"""
|
"""
|
||||||
SELECT id, content, metadata, created_at, updated_at
|
SELECT id, content, uri, metadata, created_at, updated_at
|
||||||
FROM documents WHERE id = ?
|
FROM documents WHERE id = ?
|
||||||
""",
|
""",
|
||||||
(entity_id,),
|
(entity_id,),
|
||||||
|
|
@ -75,12 +76,43 @@ class DocumentRepository(BaseRepository[Document]):
|
||||||
if row is None:
|
if row is None:
|
||||||
return 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 {}
|
metadata = json.loads(metadata_json) if metadata_json else {}
|
||||||
|
|
||||||
return Document(
|
return Document(
|
||||||
id=document_id,
|
id=document_id,
|
||||||
content=content,
|
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,
|
metadata=metadata,
|
||||||
created_at=created_at,
|
created_at=created_at,
|
||||||
updated_at=updated_at,
|
updated_at=updated_at,
|
||||||
|
|
@ -103,11 +135,12 @@ class DocumentRepository(BaseRepository[Document]):
|
||||||
cursor.execute(
|
cursor.execute(
|
||||||
"""
|
"""
|
||||||
UPDATE documents
|
UPDATE documents
|
||||||
SET content = ?, metadata = ?, updated_at = ?
|
SET content = ?, uri = ?, metadata = ?, updated_at = ?
|
||||||
WHERE id = ?
|
WHERE id = ?
|
||||||
""",
|
""",
|
||||||
(
|
(
|
||||||
entity.content,
|
entity.content,
|
||||||
|
entity.uri,
|
||||||
json.dumps(entity.metadata),
|
json.dumps(entity.metadata),
|
||||||
entity.updated_at,
|
entity.updated_at,
|
||||||
entity.id,
|
entity.id,
|
||||||
|
|
@ -150,7 +183,7 @@ class DocumentRepository(BaseRepository[Document]):
|
||||||
raise ValueError("Store connection is not available")
|
raise ValueError("Store connection is not available")
|
||||||
|
|
||||||
cursor = self.store._connection.cursor()
|
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 = []
|
params = []
|
||||||
|
|
||||||
if limit is not None:
|
if limit is not None:
|
||||||
|
|
@ -166,12 +199,13 @@ class DocumentRepository(BaseRepository[Document]):
|
||||||
|
|
||||||
documents = []
|
documents = []
|
||||||
for row in rows:
|
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 {}
|
metadata = json.loads(metadata_json) if metadata_json else {}
|
||||||
documents.append(
|
documents.append(
|
||||||
Document(
|
Document(
|
||||||
id=document_id,
|
id=document_id,
|
||||||
content=content,
|
content=content,
|
||||||
|
uri=uri,
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
created_at=created_at,
|
created_at=created_at,
|
||||||
updated_at=updated_at,
|
updated_at=updated_at,
|
||||||
|
|
|
||||||
|
|
@ -12,52 +12,61 @@ async def test_create_document_with_chunks(qa_corpus: Dataset):
|
||||||
# Create an in-memory store and repository
|
# Create an in-memory store and repository
|
||||||
store = Store(":memory:")
|
store = Store(":memory:")
|
||||||
doc_repo = DocumentRepository(store)
|
doc_repo = DocumentRepository(store)
|
||||||
|
|
||||||
# Get the first document from the corpus
|
# Get the first document from the corpus
|
||||||
first_doc = qa_corpus[0]
|
first_doc = qa_corpus[0]
|
||||||
document_text = first_doc["document_extracted"]
|
document_text = first_doc["document_extracted"]
|
||||||
|
|
||||||
# Create a Document instance
|
# Create a Document instance
|
||||||
document = Document(
|
document = Document(
|
||||||
content=document_text,
|
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
|
# Create the document with chunks in the database
|
||||||
created_document = await doc_repo.create(document)
|
created_document = await doc_repo.create(document)
|
||||||
|
|
||||||
# Verify the document was created
|
# Verify the document was created
|
||||||
assert created_document.id is not None
|
assert created_document.id is not None
|
||||||
assert created_document.content == document_text
|
assert created_document.content == document_text
|
||||||
|
|
||||||
# Check that chunks were created in the database
|
# Check that chunks were created in the database
|
||||||
if store._connection is not None:
|
if store._connection is not None:
|
||||||
cursor = store._connection.cursor()
|
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]
|
chunk_count = cursor.fetchone()[0]
|
||||||
|
|
||||||
assert chunk_count > 0
|
assert chunk_count > 0
|
||||||
|
|
||||||
# Check that embeddings were created
|
# Check that embeddings were created
|
||||||
cursor.execute("""
|
cursor.execute(
|
||||||
|
"""
|
||||||
SELECT COUNT(*) FROM chunk_embeddings ce
|
SELECT COUNT(*) FROM chunk_embeddings ce
|
||||||
JOIN chunks c ON c.id = ce.chunk_id
|
JOIN chunks c ON c.id = ce.chunk_id
|
||||||
WHERE c.document_id = ?
|
WHERE c.document_id = ?
|
||||||
""", (created_document.id,))
|
""",
|
||||||
|
(created_document.id,),
|
||||||
|
)
|
||||||
embedding_count = cursor.fetchone()[0]
|
embedding_count = cursor.fetchone()[0]
|
||||||
|
|
||||||
assert embedding_count == chunk_count
|
assert embedding_count == chunk_count
|
||||||
|
|
||||||
# Verify chunk metadata contains order information
|
# 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()
|
chunk_metadata = cursor.fetchall()
|
||||||
|
|
||||||
for i, (metadata_json,) in enumerate(chunk_metadata):
|
for i, (metadata_json,) in enumerate(chunk_metadata):
|
||||||
import json
|
import json
|
||||||
|
|
||||||
metadata = json.loads(metadata_json)
|
metadata = json.loads(metadata_json)
|
||||||
assert "order" in metadata
|
assert "order" in metadata
|
||||||
assert metadata["order"] == i
|
assert metadata["order"] == i
|
||||||
|
|
||||||
store.close()
|
store.close()
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -67,40 +76,52 @@ async def test_document_repository_crud(qa_corpus: Dataset):
|
||||||
# Create an in-memory store and repository
|
# Create an in-memory store and repository
|
||||||
store = Store(":memory:")
|
store = Store(":memory:")
|
||||||
doc_repo = DocumentRepository(store)
|
doc_repo = DocumentRepository(store)
|
||||||
|
|
||||||
# Get the first document from the corpus
|
# Get the first document from the corpus
|
||||||
first_doc = qa_corpus[0]
|
first_doc = qa_corpus[0]
|
||||||
document_text = first_doc["document_extracted"]
|
document_text = first_doc["document_extracted"]
|
||||||
|
|
||||||
# Create a document
|
# Create a document with URI
|
||||||
|
test_uri = "file:///path/to/test.txt"
|
||||||
document = Document(
|
document = Document(
|
||||||
content=document_text,
|
content=document_text, uri=test_uri, metadata={"source": "test"}
|
||||||
metadata={"source": "test"}
|
|
||||||
)
|
)
|
||||||
created_document = await doc_repo.create(document)
|
created_document = await doc_repo.create(document)
|
||||||
|
|
||||||
# Test get_by_id
|
# Test get_by_id
|
||||||
assert created_document.id is not None
|
assert created_document.id is not None
|
||||||
retrieved_document = await doc_repo.get_by_id(created_document.id)
|
retrieved_document = await doc_repo.get_by_id(created_document.id)
|
||||||
assert retrieved_document is not None
|
assert retrieved_document is not None
|
||||||
assert retrieved_document.content == document_text
|
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)
|
# Test update (should regenerate chunks)
|
||||||
retrieved_document.content = "Updated content for testing"
|
retrieved_document.content = "Updated content for testing"
|
||||||
updated_document = await doc_repo.update(retrieved_document)
|
updated_document = await doc_repo.update(retrieved_document)
|
||||||
assert updated_document.content == "Updated content for testing"
|
assert updated_document.content == "Updated content for testing"
|
||||||
|
|
||||||
# Test list_all
|
# Test list_all
|
||||||
all_documents = await doc_repo.list_all()
|
all_documents = await doc_repo.list_all()
|
||||||
assert len(all_documents) == 1
|
assert len(all_documents) == 1
|
||||||
assert all_documents[0].id == created_document.id
|
assert all_documents[0].id == created_document.id
|
||||||
|
|
||||||
# Test delete
|
# Test delete
|
||||||
deleted = await doc_repo.delete(created_document.id)
|
deleted = await doc_repo.delete(created_document.id)
|
||||||
assert deleted is True
|
assert deleted is True
|
||||||
|
|
||||||
# Verify document is gone
|
# Verify document is gone
|
||||||
retrieved_document = await doc_repo.get_by_id(created_document.id)
|
retrieved_document = await doc_repo.get_by_id(created_document.id)
|
||||||
assert retrieved_document is None
|
assert retrieved_document is None
|
||||||
|
|
||||||
store.close()
|
store.close()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue