Use resolve_document()
This commit is contained in:
parent
764b62ee20
commit
122da834d7
3 changed files with 63 additions and 33 deletions
|
|
@ -12,7 +12,6 @@ def build_namespace(
|
||||||
client: Any, config: Any, context: Any, loop: asyncio.AbstractEventLoop
|
client: Any, config: Any, context: Any, loop: asyncio.AbstractEventLoop
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Build execution namespace with haiku.rag functions injected."""
|
"""Build execution namespace with haiku.rag functions injected."""
|
||||||
from haiku.rag.store.repositories.document import _escape_sql_string
|
|
||||||
|
|
||||||
def run_async(coro: Any) -> Any:
|
def run_async(coro: Any) -> Any:
|
||||||
"""Run async coroutine from sync context using thread-safe scheduling."""
|
"""Run async coroutine from sync context using thread-safe scheduling."""
|
||||||
|
|
@ -58,37 +57,15 @@ def build_namespace(
|
||||||
|
|
||||||
def get_document(id_or_title: str) -> str | None:
|
def get_document(id_or_title: str) -> str | None:
|
||||||
async def _get() -> str | None:
|
async def _get() -> str | None:
|
||||||
doc = await client.get_document_by_id(id_or_title)
|
doc = await client.resolve_document(id_or_title)
|
||||||
if doc:
|
return doc.content if doc else None
|
||||||
return doc.content
|
|
||||||
safe_input = _escape_sql_string(id_or_title)
|
|
||||||
docs = await client.list_documents(filter=f"title = '{safe_input}'")
|
|
||||||
if docs and docs[0].id:
|
|
||||||
full_doc = await client.get_document_by_id(docs[0].id)
|
|
||||||
return full_doc.content if full_doc else None
|
|
||||||
docs = await client.list_documents(filter=f"uri = '{safe_input}'")
|
|
||||||
if docs and docs[0].id:
|
|
||||||
full_doc = await client.get_document_by_id(docs[0].id)
|
|
||||||
return full_doc.content if full_doc else None
|
|
||||||
return None
|
|
||||||
|
|
||||||
return run_async(_get())
|
return run_async(_get())
|
||||||
|
|
||||||
def get_docling_document(id_or_title: str) -> Any:
|
def get_docling_document(id_or_title: str) -> Any:
|
||||||
async def _get() -> Any:
|
async def _get() -> Any:
|
||||||
doc = await client.get_document_by_id(id_or_title)
|
doc = await client.resolve_document(id_or_title)
|
||||||
if doc:
|
return doc.get_docling_document() if doc else None
|
||||||
return doc.get_docling_document()
|
|
||||||
safe_input = _escape_sql_string(id_or_title)
|
|
||||||
docs = await client.list_documents(filter=f"title = '{safe_input}'")
|
|
||||||
if docs and docs[0].id:
|
|
||||||
full_doc = await client.get_document_by_id(docs[0].id)
|
|
||||||
return full_doc.get_docling_document() if full_doc else None
|
|
||||||
docs = await client.list_documents(filter=f"uri = '{safe_input}'")
|
|
||||||
if docs and docs[0].id:
|
|
||||||
full_doc = await client.get_document_by_id(docs[0].id)
|
|
||||||
return full_doc.get_docling_document() if full_doc else None
|
|
||||||
return None
|
|
||||||
|
|
||||||
return run_async(_get())
|
return run_async(_get())
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -739,6 +739,30 @@ class HaikuRAG:
|
||||||
"""
|
"""
|
||||||
return await self.document_repository.get_by_uri(uri)
|
return await self.document_repository.get_by_uri(uri)
|
||||||
|
|
||||||
|
async def resolve_document(self, id_or_title: str) -> Document | None:
|
||||||
|
"""Resolve a document by ID, title, or URI (in that order).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
id_or_title: Document ID, title, or URI to look up.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The Document instance if found, None otherwise.
|
||||||
|
"""
|
||||||
|
doc = await self.get_document_by_id(id_or_title)
|
||||||
|
if doc:
|
||||||
|
return doc
|
||||||
|
|
||||||
|
safe_input = _escape_sql_string(id_or_title)
|
||||||
|
docs = await self.list_documents(filter=f"title = '{safe_input}'")
|
||||||
|
if docs and docs[0].id:
|
||||||
|
return await self.get_document_by_id(docs[0].id)
|
||||||
|
|
||||||
|
docs = await self.list_documents(filter=f"uri = '{safe_input}'")
|
||||||
|
if docs and docs[0].id:
|
||||||
|
return await self.get_document_by_id(docs[0].id)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
async def update_document(
|
async def update_document(
|
||||||
self,
|
self,
|
||||||
document_id: str,
|
document_id: str,
|
||||||
|
|
@ -1328,12 +1352,7 @@ class HaikuRAG:
|
||||||
if documents:
|
if documents:
|
||||||
loaded_docs = []
|
loaded_docs = []
|
||||||
for doc_ref in documents:
|
for doc_ref in documents:
|
||||||
doc = await self.get_document_by_id(doc_ref)
|
doc = await self.resolve_document(doc_ref)
|
||||||
if not doc:
|
|
||||||
safe_ref = _escape_sql_string(doc_ref)
|
|
||||||
docs = await self.list_documents(filter=f"title = '{safe_ref}'")
|
|
||||||
if docs and docs[0].id:
|
|
||||||
doc = await self.get_document_by_id(docs[0].id)
|
|
||||||
if doc:
|
if doc:
|
||||||
loaded_docs.append(doc)
|
loaded_docs.append(doc)
|
||||||
context.documents = loaded_docs if loaded_docs else None
|
context.documents = loaded_docs if loaded_docs else None
|
||||||
|
|
|
||||||
|
|
@ -85,6 +85,40 @@ async def test_client_document_crud(qa_corpus: Dataset, temp_db_path):
|
||||||
assert deleted_again is False
|
assert deleted_again is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_client_resolve_document(temp_db_path):
|
||||||
|
"""Test resolve_document finds documents by ID, title, or URI."""
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
doc = await client.create_document(
|
||||||
|
content="Test content",
|
||||||
|
uri="test://resolve-test",
|
||||||
|
title="Resolve Test Doc",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Resolve by ID
|
||||||
|
by_id = await client.resolve_document(doc.id)
|
||||||
|
assert by_id is not None
|
||||||
|
assert by_id.id == doc.id
|
||||||
|
|
||||||
|
# Resolve by title
|
||||||
|
by_title = await client.resolve_document("Resolve Test Doc")
|
||||||
|
assert by_title is not None
|
||||||
|
assert by_title.id == doc.id
|
||||||
|
|
||||||
|
# Resolve by URI
|
||||||
|
by_uri = await client.resolve_document("test://resolve-test")
|
||||||
|
assert by_uri is not None
|
||||||
|
assert by_uri.id == doc.id
|
||||||
|
|
||||||
|
# Not found returns None
|
||||||
|
not_found = await client.resolve_document("nonexistent")
|
||||||
|
assert not_found is None
|
||||||
|
|
||||||
|
# SQL injection is escaped
|
||||||
|
injection = "x' OR title LIKE '%"
|
||||||
|
injected = await client.resolve_document(injection)
|
||||||
|
assert injected is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.vcr()
|
@pytest.mark.vcr()
|
||||||
async def test_client_update_document(qa_corpus: Dataset, temp_db_path):
|
async def test_client_update_document(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test updating document with individual parameters."""
|
"""Test updating document with individual parameters."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue