chore: add split to mock server (#199)

This commit is contained in:
Clelia (Astra) Bertelli
2026-01-27 14:17:04 +01:00
committed by GitHub
parent 9b970b7c83
commit 6f4e661fb4
5 changed files with 289 additions and 3 deletions
+7 -1
View File
@@ -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):
@@ -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
+15 -2
View File
@@ -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
@@ -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
+111
View File
@@ -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