mirror of
https://github.com/run-llama/template-workflow-extract-basic.git
synced 2026-07-19 18:53:47 -04:00
chore: add split to mock server (#199)
This commit is contained in:
committed by
GitHub
parent
9b970b7c83
commit
6f4e661fb4
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user