diff --git a/src/extraction_review/config.py b/src/extraction_review/config.py index 2d47d4b..4be6cf8 100644 --- a/src/extraction_review/config.py +++ b/src/extraction_review/config.py @@ -48,10 +48,16 @@ class SplitCategory(BaseModel): description: str +class SplittingStrategy(BaseModel): + """Strategy for document splitting""" + + allow_uncategorized: bool = False + + class SplitSettings(BaseModel): """Settings for document splitting.""" - pass + splitting_strategy: SplittingStrategy = SplittingStrategy() class SplitConfig(BaseModel): diff --git a/src/extraction_review/testing_utils/_deterministic.py b/src/extraction_review/testing_utils/_deterministic.py index f4ca1d1..598d616 100644 --- a/src/extraction_review/testing_utils/_deterministic.py +++ b/src/extraction_review/testing_utils/_deterministic.py @@ -197,3 +197,19 @@ def _generate_value(schema: Any, rng: random.Random, depth: int) -> Any: return _generate_value(option, rng, depth + 1) return generate_text_blob(rng.randint(0, 1_000_000), sentences=1) + + +def categorize_pages( + content: bytes, categories: list[str], seed: int +) -> dict[str, list[int]]: + rng = random.Random(seed) + page_size = rng.randint(1, 50) + categorized_pages: dict[str, list[int]] = {c: [] for c in categories} + i = 0 + j = 0 + while j + page_size < len(content): + i += 1 + category = rng.choice(categories) + categorized_pages[category].append(i) + j += page_size + return categorized_pages diff --git a/src/extraction_review/testing_utils/server.py b/src/extraction_review/testing_utils/server.py index bdb72ed..ab4f16a 100644 --- a/src/extraction_review/testing_utils/server.py +++ b/src/extraction_review/testing_utils/server.py @@ -8,11 +8,12 @@ from typing import Any, Callable, Dict, Optional, Sequence import httpx import respx +from .agent_data import FakeAgentDataNamespace from .classify import FakeClassifyNamespace from .extract import FakeExtractNamespace from .files import FakeFilesNamespace from .parse import FakeParseNamespace -from .agent_data import FakeAgentDataNamespace +from .split import FakeSplitNamespace Handler = Callable[[httpx.Request], httpx.Response] @@ -31,14 +32,23 @@ class FakeLlamaCloudServer: download_base_url: Optional[str] = None, default_project_id: str = "proj-test", default_organization_id: str = "org-test", + default_user_id: str = "user-test", ) -> None: self.base_urls = tuple(base_urls or (self.DEFAULT_BASE_URL,)) - selected = namespaces or ("files", "extract", "parse", "classify", "agent_data") + selected = namespaces or ( + "files", + "extract", + "parse", + "classify", + "agent_data", + "split", + ) self._namespace_names = {name.lower() for name in selected} self._upload_base_url = upload_base_url or self.DEFAULT_UPLOAD_BASE self._download_base_url = download_base_url or self.DEFAULT_DOWNLOAD_BASE self.default_project_id = default_project_id self.default_organization_id = default_organization_id + self.default_user_id = default_user_id self.router = respx.MockRouter(assert_all_called=False) self._installed = False self._registered = False @@ -52,6 +62,7 @@ class FakeLlamaCloudServer: self.parse = FakeParseNamespace(server=self) self.classify = FakeClassifyNamespace(server=self, files=self.files) self.agent_data = FakeAgentDataNamespace(server=self) + self.split = FakeSplitNamespace(server=self) # Context management ---------------------------------------------- def install(self) -> "FakeLlamaCloudServer": @@ -168,6 +179,8 @@ class FakeLlamaCloudServer: self.classify.register() if "agent_data" in self._namespace_names: self.agent_data.register() + if "split" in self._namespace_names: + self.split.register() self._registered = True diff --git a/src/extraction_review/testing_utils/split.py b/src/extraction_review/testing_utils/split.py new file mode 100644 index 0000000..7eda325 --- /dev/null +++ b/src/extraction_review/testing_utils/split.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +import httpx +from llama_cloud.types.beta.split_category import SplitCategory +from llama_cloud.types.beta.split_category_param import SplitCategoryParam +from llama_cloud.types.beta.split_create_response import SplitCreateResponse +from llama_cloud.types.beta.split_document_input import SplitDocumentInput +from llama_cloud.types.beta.split_get_response import SplitGetResponse +from llama_cloud.types.beta.split_result_response import SplitResultResponse +from llama_cloud.types.beta.split_segment_response import SplitSegmentResponse +from pydantic.dataclasses import dataclass + +from ._deterministic import categorize_pages, utcnow +from .files import StoredFile + +if TYPE_CHECKING: + from .server import FakeLlamaCloudServer + + +@dataclass +class SplitRequest: + categories: list[SplitCategoryParam] + file_id: str + stored_file: StoredFile + + +class FakeSplitNamespace: + def __init__(self, *, server: "FakeLlamaCloudServer") -> None: + self._server = server + self._jobs: dict[str, SplitGetResponse] = {} + self.routes: dict[str, Any] = {} + self._allowed_input_types = ("file_id",) + self._page_size = 50 + + def _validate_split_request( + self, request: httpx.Request + ) -> httpx.Response | SplitRequest: + payload = self._server.json(request) + document_input = payload.get("document_input") + if not document_input: + response = {"detail": "the document_input field should be non-null"} + return self._server.json_response(response, status_code=400) + input_type = document_input.get("type", "file_id") + if input_type not in self._allowed_input_types: + response = { + "detail": f"document_input.type {input_type} is invalid. Allowed input types: {', '.join(self._allowed_input_types)}" + } + return self._server.json_response(response, status_code=400) + input_value = document_input.get("value") + if input_value is None: + response = {"detail": "Missing document_input.value field"} + return self._server.json_response(response, status_code=400) + categories = payload.get("categories", []) + if not categories: + response = {"detail": "categories field should be non-null and non-empty"} + return self._server.json_response(response, status_code=400) + stored_file = self._server.files.get(input_value) + if stored_file is None: + response = {"detail": f"file with ID {input_value} not found"} + return self._server.json_response(response, status_code=404) + return SplitRequest( + categories=categories, file_id=input_value, stored_file=stored_file + ) + + def _create_split_job(self, request: httpx.Request) -> httpx.Response: + validated = self._validate_split_request(request) + if isinstance(validated, httpx.Response): + return validated + categorized = categorize_pages( + validated.stored_file.content, + [category["name"] for category in validated.categories], + 0, + ) + result = SplitResultResponse(segments=[]) + for c in categorized: + result.segments.append( + SplitSegmentResponse( + category=c, confidence_category="high", pages=categorized[c] + ) + ) + job_id = self._server.new_id("split-") + job = SplitGetResponse( + id=job_id, + categories=[ + SplitCategory(name=c["name"], description=c.get("description")) + for c in validated.categories + ], + document_input=SplitDocumentInput(type="file_id", value=validated.file_id), + project_id=self._server.default_project_id, + user_id=self._server.default_user_id, + status="completed", + result=result, + created_at=utcnow(), + updated_at=utcnow(), + error_message=None, + ) + self._jobs[job_id] = job + response = SplitCreateResponse( + id=job_id, + categories=[ + SplitCategory(name=c["name"], description=c.get("description")) + for c in validated.categories + ], + document_input=SplitDocumentInput(type="file_id", value=validated.file_id), + project_id=self._server.default_project_id, + user_id=self._server.default_user_id, + status="pending", + error_message=None, + ) + return self._server.json_response(response.model_dump(), status_code=200) + + def _get_split_job_result(self, request: httpx.Request) -> httpx.Response: + job_id = request.url.path.split("/")[-1] + job = self._jobs.get(job_id) + if job is not None: + return self._server.json_response(job.model_dump()) + return self._server.json_response( + {"detail": f"job with ID {job_id} does not exist"}, status_code=404 + ) + + def register(self) -> None: + server = self._server + create_route = server.add_route( + "POST", + "/api/v1/beta/split/jobs", + self._create_split_job, + namespace="split", + alias="create", + ) + self.routes["create"] = create_route + get_route = server.add_route( + "GET", + "/api/v1/beta/split/jobs/{split_job_id}", + self._get_split_job_result, + namespace="split", + alias="get", + ) + self.routes["get"] = get_route diff --git a/tests/testing_utils/test_split.py b/tests/testing_utils/test_split.py new file mode 100644 index 0000000..770c4df --- /dev/null +++ b/tests/testing_utils/test_split.py @@ -0,0 +1,111 @@ +import pytest +from extraction_review.testing_utils import FakeLlamaCloudServer +from llama_cloud import APIStatusError, AsyncLlamaCloud + + +@pytest.fixture +def server(): + with FakeLlamaCloudServer() as srv: + yield srv + + +@pytest.fixture() +def client() -> AsyncLlamaCloud: + return AsyncLlamaCloud(api_key="fake-api-key") + + +@pytest.mark.asyncio +async def test_split_end_to_end( + server: FakeLlamaCloudServer, client: AsyncLlamaCloud +) -> None: + file_id = server.files.preload(path="tests/files/test.pdf") + split_job = await client.beta.split.create( + categories=[ + {"name": "hello", "description": ""}, + {"name": "world", "description": ""}, + ], + document_input={"type": "file_id", "value": file_id}, + ) + assert split_job.id.startswith("split-") + assert split_job.status == "pending" + cts = [c.name for c in split_job.categories] + cts.sort() + assert cts == ["hello", "world"] + split_result = await client.beta.split.wait_for_completion( + split_job_id=split_job.id, + ) + assert split_result.result is not None + cts = [s.category for s in split_result.result.segments] + cts.sort() + assert cts == [ + "hello", + "world", + ] + assert any(len(s.pages) > 0 for s in split_result.result.segments) + + +@pytest.mark.asyncio +async def test_split_no_categories_raises_bad_request( + server: FakeLlamaCloudServer, client: AsyncLlamaCloud +) -> None: + file_id = server.files.preload(path="tests/files/test.pdf") + # should fail because there are no categories + with pytest.raises(APIStatusError) as exc_info: + await client.beta.split.create( + categories=[], + document_input={"type": "file_id", "value": file_id}, + ) + assert exc_info.value.status_code == 400 + assert ( + "categories field should be non-null and non-empty" + in exc_info.value.message + ) + + +@pytest.mark.asyncio +async def test_split_invalid_document_input_type_raises_bad_request( + server: FakeLlamaCloudServer, client: AsyncLlamaCloud +) -> None: + file_id = server.files.preload(path="tests/files/test.pdf") + # should fail because the only allowed document_input type is file_id + with pytest.raises(APIStatusError) as exc_info: + await client.beta.split.create( + categories=[ + {"name": "hello", "description": ""}, + {"name": "world", "description": ""}, + ], + document_input={"type": "file", "value": file_id}, + ) + assert exc_info.value.status_code == 400 + assert ( + "document_input.type file is invalid. Allowed input types: file_id" + in exc_info.value.message + ) + + +@pytest.mark.asyncio +async def test_split_non_existing_file_id_raises_notfound( + server: FakeLlamaCloudServer, client: AsyncLlamaCloud +) -> None: + file_id = "file-doesnotexist" + with pytest.raises(APIStatusError) as exc_info: + await client.beta.split.create( + categories=[ + {"name": "hello", "description": ""}, + {"name": "world", "description": ""}, + ], + document_input={"type": "file_id", "value": file_id}, + ) + assert exc_info.value.status_code == 404 + assert f"file with ID {file_id} not found" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_split_non_existing_job_id_raises_notfound( + server: FakeLlamaCloudServer, client: AsyncLlamaCloud +) -> None: + with pytest.raises(APIStatusError) as exc_info: + job_id = "split-doesnotexist" + await client.beta.split.get(job_id) + assert exc_info.value.status_code == 404 + assert f"job with ID {job_id} does not exist" in exc_info.value.message