Commit 0d61c1f7 authored by Kantz's avatar Kantz
Browse files

moving retrieval to retrieval store

parent 8c122e25
# Package marker for services
from app.deterministic_services.vector_store import Source, SourceID
from app.deterministic_services.retrieval_store import Source, SourceID
__all__ = ["Source", "SourceID"]
\ No newline at end of file
......@@ -2,11 +2,13 @@ from __future__ import annotations
import os
import unittest
from unittest.mock import patch
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.deterministic_services.vector_store import (
from app.deterministic_services import Source as PackageSource, SourceID as PackageSourceID
from app.deterministic_services.retrieval_store import (
Retrieved,
Source,
SourceID,
......@@ -14,6 +16,8 @@ from app.deterministic_services.vector_store import (
expand_neighbor_children,
merge_retrieval_groups,
merge_sources,
retrieve,
retrieve_with_subsections,
select_dominant_scope,
)
......@@ -52,6 +56,10 @@ def _mk_retrieved(
class VectorStorePipelineUnitTest(unittest.TestCase):
def test_package_exports_point_to_retrieval_models(self) -> None:
self.assertIs(PackageSource, Source)
self.assertIs(PackageSourceID, SourceID)
def test_merge_sources_prefers_task_childs_on_duplicate(self) -> None:
child_direct = Source(
source_id=SourceID(
......@@ -238,6 +246,47 @@ class VectorStorePipelineUnitTest(unittest.TestCase):
merged = merge_sources([first], [second])
self.assertEqual(len(merged), 2)
def test_retrieve_sorts_sources_by_score(self) -> None:
groups = {
"children_direct": [
_mk_retrieved("low", 0.2, 1, 1, 1),
_mk_retrieved("high", 0.9, 1, 1, 1),
]
}
with patch("app.deterministic_services.retrieval_store._use_subsection_retrieval", return_value=False):
with patch("app.deterministic_services.retrieval_store.run_child_retrieval_pipeline", return_value=groups):
result = retrieve(
pg_url="postgresql://unused",
embedder=object(), # type: ignore[arg-type]
query="query",
k=2,
)
self.assertEqual([source.source_id.title for source in result], ["high", "low"])
def test_retrieve_with_subsections_matches_retrieve_without_refs(self) -> None:
expected = [
Source(
source_id=SourceID(title="Child", doc_type="child"),
retrieved_as="children_direct",
source_type="child",
score=0.7,
markdown="md",
)
]
with patch("app.deterministic_services.retrieval_store._use_subsection_retrieval", return_value=False):
with patch("app.deterministic_services.retrieval_store.retrieve", return_value=expected):
result = retrieve_with_subsections(
pg_url="postgresql://unused",
embedder=object(), # type: ignore[arg-type]
query="query",
subsection_refs=None,
)
self.assertEqual(result, expected)
if __name__ == "__main__":
unittest.main()
Supports Markdown
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment