From 5c4799164c342399906b131d38a680d55b8a82d7 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 11 Dec 2025 10:17:45 +0200 Subject: [PATCH] Escape when getting by URI --- .../haiku/rag/store/repositories/document.py | 8 +++++- tests/test_document.py | 27 +++++++++++++++++++ 2 files changed, 34 insertions(+), 1 deletion(-) diff --git a/haiku_rag_slim/haiku/rag/store/repositories/document.py b/haiku_rag_slim/haiku/rag/store/repositories/document.py index 52172ccf..f9a2c037 100644 --- a/haiku_rag_slim/haiku/rag/store/repositories/document.py +++ b/haiku_rag_slim/haiku/rag/store/repositories/document.py @@ -6,6 +6,11 @@ from haiku.rag.store.engine import DocumentRecord, Store from haiku.rag.store.models.document import Document +def _escape_sql_string(value: str) -> str: + """Escape single quotes in SQL string literals.""" + return value.replace("'", "''") + + class DocumentRepository: """Repository for Document operations.""" @@ -161,9 +166,10 @@ class DocumentRepository: async def get_by_uri(self, uri: str) -> Document | None: """Get a document by its URI.""" + escaped_uri = _escape_sql_string(uri) results = list( self.store.documents_table.search() - .where(f"uri = '{uri}'") + .where(f"uri = '{escaped_uri}'") .limit(1) .to_pydantic(DocumentRecord) ) diff --git a/tests/test_document.py b/tests/test_document.py index ba884696..1cf25678 100644 --- a/tests/test_document.py +++ b/tests/test_document.py @@ -173,3 +173,30 @@ def test_document_get_docling_document_no_id_no_cache(): # Each call parses fresh (different objects) assert doc1 is not doc2 + + +@pytest.mark.asyncio +async def test_document_get_by_uri_with_special_characters( + qa_corpus: Dataset, temp_db_path +): + """Test get_by_uri handles URIs with special characters like single quotes.""" + store = Store(temp_db_path, create=True) + doc_repo = DocumentRepository(store) + + first_doc = qa_corpus[0] + document_text = first_doc["document_extracted"] + + doc_with_quote = Document( + content=document_text, + uri="Hamish and Andy's Gap Year", + metadata={"source": "test"}, + ) + + created_doc = await doc_repo.create(doc_with_quote) + + retrieved = await doc_repo.get_by_uri("Hamish and Andy's Gap Year") + assert retrieved is not None + assert retrieved.id == created_doc.id + assert retrieved.uri == "Hamish and Andy's Gap Year" + + store.close()