feat: template with new sdk (#178)

* feat: template with new sdk

* chore: more test updates

* chore: more tests

* chore: update system prompt

* fix: generate_value for nested pydantic models

* chore: template validation

* chore: add ExtractedData

* chore: last tweaks to new sdk template

* feat: add --template flag to bundle-coder.sh (#181)

* feat: add --template flag to bundle-coder.sh

Add a -T/--template flag to select which template to bundle instead of
creating a separate script. This consolidates bundle-coder.sh and
bundle-coder-new-sdk.sh into a single script.

- Default template remains 'extract-basic'
- Use --template extract-basic-new to bundle the new SDK template
- Template name suffix is included in the e2b template name

* Update bundle-coder.sh

---------

Co-authored-by: Claude <noreply@anthropic.com>

* chore: pr suggestions

* chore: package versioning

* chore: move extract-basic-new to extract-basic

* chore: prompt tweaks

* fix: typo

* fix: typo pt2

* chore: pr suggestions

---------

Co-authored-by: Adrian Lyjak <adrianlyjak@gmail.com>
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
Clelia (Astra) Bertelli
2026-01-23 16:12:00 +01:00
committed by GitHub
parent 8513a08de0
commit adda5d6b32
14 changed files with 900 additions and 648 deletions
+1 -1
View File
@@ -5,7 +5,6 @@ description = "Extracts data"
readme = "README.md"
requires-python = ">=3.12"
dependencies = [
"llama-cloud-services>=0.6.69",
"llama-index-workflows>=2.11.3,<3.0.0",
"python-dotenv>=1.1.0",
"jsonref>=1.1.0",
@@ -13,6 +12,7 @@ dependencies = [
"httpx>=0.28.1",
"llama-index-core>=0.14.0",
"respx>=0.22.0,<1",
"llama-cloud>=1.0.0b7",
]
[dependency-groups]
+5 -33
View File
@@ -1,16 +1,9 @@
import os
from typing import Any
import httpx
from llama_cloud_services import LlamaExtract
from llama_cloud_services.beta.agent_data import AsyncAgentDataClient, ExtractedData
from llama_cloud.client import AsyncLlamaCloud
from .testing_utils import FakeLlamaCloudServer
import logging
import os
from extraction_review.config import (
EXTRACTED_DATA_COLLECTION,
)
from llama_cloud import AsyncLlamaCloud
from .testing_utils import FakeLlamaCloudServer
logger = logging.getLogger(__name__)
@@ -31,27 +24,6 @@ else:
fake = None
def get_llama_extract() -> LlamaExtract:
"""Document extractor for parsing and extracting structured fields."""
return LlamaExtract(api_key=api_key, base_url=base_url, project_id=project_id)
def get_data_client() -> AsyncAgentDataClient:
"""Storage for extracted data, enabling human review and corrections."""
return AsyncAgentDataClient(
deployment_name=agent_name,
collection=EXTRACTED_DATA_COLLECTION,
type=ExtractedData[Any],
client=get_llama_cloud_client(),
)
def get_llama_cloud_client() -> AsyncLlamaCloud:
"""Cloud services connection for file storage and processing."""
return AsyncLlamaCloud(
base_url=base_url,
token=api_key,
httpx_client=httpx.AsyncClient(
timeout=60, headers={"Project-Id": project_id} if project_id else None
),
)
return AsyncLlamaCloud(api_key=api_key, base_url=base_url)
+3 -4
View File
@@ -6,8 +6,7 @@ If you need more control, feel free to edit the rest of the application
from __future__ import annotations
from llama_cloud import ExtractConfig
from llama_cloud_services.extract import ExtractMode
from llama_cloud.types.extraction.extract_config_param import ExtractConfigParam
from pydantic import BaseModel, Field
# The name of the collection to use for storing extracted data. This will be qualified by the agent name.
@@ -34,8 +33,8 @@ class ExtractionSchema(BaseModel):
)
EXTRACT_CONFIG = ExtractConfig(
extraction_mode=ExtractMode.PREMIUM,
EXTRACT_CONFIG = ExtractConfigParam(
extraction_mode="PREMIUM",
system_prompt=None,
# advanced. Only compatible with Premium mode.
citation_bbox=True,
+125 -83
View File
@@ -2,26 +2,21 @@ import asyncio
import hashlib
import logging
import os
from pathlib import Path
import tempfile
from typing import Any, Literal, Annotated
from pathlib import Path
from typing import Annotated, Any, Literal
import httpx
from llama_cloud import ExtractRun
from llama_cloud.client import AsyncLlamaCloud
from llama_cloud_services.extract import SourceText, LlamaExtract
from llama_cloud_services.beta.agent_data import (
ExtractedData,
InvalidExtractionData,
AsyncAgentDataClient,
)
from llama_cloud import AsyncLlamaCloud
from llama_cloud.types.beta.extracted_data import ExtractedData, InvalidExtractionData
from llama_cloud.types.file_query_params import Filter
from pydantic import BaseModel
from workflows.resource import Resource
from workflows import Context, Workflow, step
from workflows.events import Event, StartEvent, StopEvent
from workflows.resource import Resource
from .config import EXTRACT_CONFIG, ExtractionSchema
from .clients import get_llama_cloud_client, get_data_client, get_llama_extract
from .clients import agent_name, get_llama_cloud_client, project_id
from .config import EXTRACT_CONFIG, EXTRACTED_DATA_COLLECTION, ExtractionSchema
logger = logging.getLogger(__name__)
@@ -43,6 +38,14 @@ class Status(Event):
message: str
class ExtractJobStartedEvent(Event):
pass
class ExtractJobCompletedEvent(Event):
pass
class ExtractedEvent(Event):
data: ExtractedData
@@ -55,6 +58,8 @@ class ExtractionState(BaseModel):
file_id: str | None = None
file_path: str | None = None
filename: str | None = None
file_hash: str | None = None
extract_job_id: str | None = None
class ProcessFileWorkflow(Workflow):
@@ -82,8 +87,11 @@ class ProcessFileWorkflow(Workflow):
if state.file_id is None:
raise ValueError("File ID is not set")
try:
file_metadata = await llama_cloud_client.files.get_file(id=state.file_id)
file_url = await llama_cloud_client.files.read_file_content(state.file_id)
files = await llama_cloud_client.files.query(
filter=Filter(file_ids=[state.file_id])
)
file_metadata = files.items[0]
file_url = await llama_cloud_client.files.get(file_id=state.file_id)
temp_dir = tempfile.gettempdir()
filename = file_metadata.name
@@ -117,43 +125,82 @@ class ProcessFileWorkflow(Workflow):
self,
event: FileDownloadedEvent,
ctx: Context[ExtractionState],
extractor: Annotated[LlamaExtract, Resource(get_llama_extract)],
) -> ExtractedEvent | ExtractedInvalidEvent:
llama_cloud_client: Annotated[
AsyncLlamaCloud, Resource(get_llama_cloud_client)
],
) -> ExtractJobStartedEvent:
"""Extract structured data fields from the document."""
state = await ctx.store.get_state()
if state.file_path is None or state.filename is None:
raise ValueError("File path or filename is not set")
# track the content of the file, so as to be able to de-duplicate
file_content = Path(state.file_path).read_bytes()
logger.info(f"Extracting data from file {state.filename}")
ctx.write_event_to_stream(
Status(level="info", message=f"Extracting data from file {state.filename}")
)
extract_job = await llama_cloud_client.extraction.run(
config=EXTRACT_CONFIG,
data_schema=ExtractionSchema.model_json_schema(),
file_id=state.file_id,
project_id=project_id,
)
async with ctx.store.edit_state() as st:
st.file_hash = hashlib.sha256(file_content).hexdigest()
st.extract_job_id = extract_job.id
return ExtractJobStartedEvent()
@step()
async def wait_for_extract_job(
self,
event: ExtractJobStartedEvent,
ctx: Context[ExtractionState],
llama_cloud_client: Annotated[
AsyncLlamaCloud, Resource(get_llama_cloud_client)
],
) -> ExtractJobCompletedEvent:
state = await ctx.store.get_state()
if state.extract_job_id is None:
raise ValueError("Job ID cannot be null when waiting for its completion")
await llama_cloud_client.extraction.jobs.wait_for_completion(
state.extract_job_id
)
return ExtractJobCompletedEvent()
@step()
async def get_extraction_job_result(
self,
event: ExtractJobCompletedEvent,
ctx: Context[ExtractionState],
llama_cloud_client: Annotated[
AsyncLlamaCloud, Resource(get_llama_cloud_client)
],
) -> ExtractedEvent | ExtractedInvalidEvent:
state = await ctx.store.get_state()
if state.extract_job_id is None:
raise ValueError("Job ID cannot be null when getting its result")
extracted_result = await llama_cloud_client.extraction.jobs.get_result(
state.extract_job_id
)
extract_run = await llama_cloud_client.extraction.runs.get(
run_id=extracted_result.run_id
)
try:
# track the content of the file, so as to be able to de-duplicate
file_content = Path(state.file_path).read_bytes()
file_hash = hashlib.sha256(file_content).hexdigest()
source_text = SourceText(
file=state.file_path,
filename=state.filename,
logger.info(f"Extracted data: {extracted_result}")
data = ExtractedData.from_extraction_result(
result=extract_run,
schema=ExtractionSchema,
# retain original file name and id, rather than using the extracted duplicate file
file_name=state.filename,
file_id=state.file_id,
file_hash=state.file_hash,
)
logger.info(f"Extracting data from file {state.filename}")
ctx.write_event_to_stream(
Status(
level="info", message=f"Extracting data from file {state.filename}"
)
)
extracted_result: ExtractRun = await extractor.aextract(
data_schema=ExtractionSchema, files=source_text, config=EXTRACT_CONFIG
)
try:
logger.info(f"Extracted data: {extracted_result}")
data = ExtractedData.from_extraction_result(
result=extracted_result,
schema=ExtractionSchema,
# retain original file name and id, rather than using the extracted duplicate file
file_name=state.filename,
file_id=state.file_id,
file_hash=file_hash,
)
return ExtractedEvent(data=data)
except InvalidExtractionData as e:
logger.error(f"Error validating extracted data: {e}", exc_info=True)
return ExtractedInvalidEvent(data=e.invalid_item)
return ExtractedEvent(data=data)
except InvalidExtractionData as e:
logger.error(f"Error validating extracted data: {e}", exc_info=True)
return ExtractedInvalidEvent(data=e.invalid_item)
except Exception as e:
logger.error(
f"Error extracting data from file {state.filename}: {e}",
@@ -172,46 +219,40 @@ class ProcessFileWorkflow(Workflow):
self,
event: ExtractedEvent | ExtractedInvalidEvent,
ctx: Context,
data_client: Annotated[AsyncAgentDataClient, Resource(get_data_client)],
llama_cloud_client: Annotated[
AsyncLlamaCloud, Resource(get_llama_cloud_client)
],
) -> StopEvent:
"""Save extracted data for human review and correction."""
try:
logger.info(f"Recorded extracted data for file {event.data.file_name}")
ctx.write_event_to_stream(
Status(
level="info",
message=f"Recorded extracted data for file {event.data.file_name}",
)
)
# remove past data when reprocessing the same file
if event.data.file_hash:
await data_client.delete(
filter={
"file_hash": {
"eq": event.data.file_hash,
},
data = event.data.model_dump()
# remove past data when reprocessing the same file
if event.data.file_hash is not None:
await llama_cloud_client.beta.agent_data.delete_by_query(
deployment_name=agent_name or "_public",
collection=EXTRACTED_DATA_COLLECTION,
filter={
"file_hash": {
"eq": event.data.file_hash,
},
)
logger.info(
f"Removing past data for file {event.data.file_name} with hash {event.data.file_hash}"
)
# finally, save the new data
item_id = await data_client.create_item(event.data)
return StopEvent(
result=item_id.id,
},
)
except Exception as e:
logger.error(
f"Error recording extracted data for file {event.data.file_name}: {e}",
exc_info=True,
logger.info(
f"Removing past data for file {event.data.file_name} with hash {event.data.file_hash}"
)
# finally, save the new data
item = await llama_cloud_client.beta.agent_data.agent_data(
data=data,
deployment_name=agent_name or "_public",
collection=EXTRACTED_DATA_COLLECTION,
)
logger.info(f"Recorded extracted data for file {event.data.file_name or ''}")
ctx.write_event_to_stream(
Status(
level="info",
message=f"Recorded extracted data for file {event.data.file_name or ''}",
)
ctx.write_event_to_stream(
Status(
level="error",
message=f"Error recording extracted data for file {event.data.file_name}: {e}",
)
)
raise e
)
return StopEvent(result=item.id)
workflow = ProcessFileWorkflow(timeout=None)
@@ -223,8 +264,9 @@ if __name__ == "__main__":
logging.basicConfig(level=logging.INFO)
async def main():
file = await get_llama_cloud_client().files.upload_file(
upload_file=Path("test.pdf").open("rb")
file = await get_llama_cloud_client().files.create(
file=Path("test.pdf").open("rb"),
purpose="extract",
)
await workflow.run(start_event=FileEvent(file_id=file.id))
@@ -3,10 +3,11 @@ from __future__ import annotations
import hashlib
import json
import random
from jsonref import replace_refs
from datetime import datetime, timezone
from typing import Any, Iterable, Mapping, MutableMapping
from jsonref import JsonRef, replace_refs
def hash_chunks(chunks: Iterable[bytes]) -> str:
digest = hashlib.sha256()
@@ -117,6 +118,9 @@ def _generate_value(schema: Any, rng: random.Random, depth: int) -> Any:
)
)
if isinstance(schema, JsonRef):
schema = dict(schema) # type: ignore
if schema is None:
return generate_text_blob(rng.randint(0, 1_000_000), sentences=1)
+13 -12
View File
@@ -4,13 +4,14 @@ from dataclasses import dataclass
from typing import TYPE_CHECKING, Dict, List
import httpx
from llama_cloud.types import (
from llama_cloud.types.classifier import (
ClassifierRule,
ClassifyJob,
ClassifyJobResults,
ClassificationResult,
FileClassification,
StatusEnum,
)
from llama_cloud.types.classifier.job_get_results_response import (
Item,
ItemResult,
JobGetResultsResponse,
)
from ._deterministic import combined_seed, utcnow
@@ -23,7 +24,7 @@ if TYPE_CHECKING:
@dataclass
class ClassificationJobRecord:
job: ClassifyJob
results: ClassifyJobResults
results: JobGetResultsResponse
files: List[StoredFile]
@@ -85,7 +86,7 @@ class FakeClassifyNamespace:
user_id="fake-user",
rules=rules,
parsing_configuration=None,
status=StatusEnum.SUCCESS,
status="SUCCESS",
created_at=utcnow(),
updated_at=utcnow(),
effective_at=utcnow(),
@@ -125,8 +126,8 @@ class FakeClassifyNamespace:
job_id: str,
stored_files: List[StoredFile],
rules: List[ClassifierRule],
) -> ClassifyJobResults:
items: List[FileClassification] = []
) -> JobGetResultsResponse:
items: List[Item] = []
for stored in stored_files:
seed = combined_seed(stored.sha256, job_id)
rule_index = seed % len(rules) if rules else 0
@@ -135,19 +136,19 @@ class FakeClassifyNamespace:
reasoning = (
f"Selected rule '{predicted_type}' using deterministic seed {seed}."
)
classification = FileClassification(
classification = Item(
id=self._server.new_id("classification"),
file_id=stored.file.id,
classify_job_id=job_id,
created_at=utcnow(),
updated_at=utcnow(),
result=ClassificationResult(
result=ItemResult(
type=predicted_type,
confidence=min(confidence, 0.95),
reasoning=reasoning,
),
)
items.append(classification)
return ClassifyJobResults(
return JobGetResultsResponse(
items=items, next_page_token=None, total_size=len(items)
)
+69 -95
View File
@@ -1,28 +1,30 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, cast
import httpx
import copy
from llama_cloud.resources.extraction.runs import AsyncPaginatedExtractRuns
from llama_cloud.types import (
ExtractAgent,
ExtractConfig,
ExtractJob,
ExtractRun,
ExtractState,
File as CloudFile,
PaginatedExtractRunsResponse,
StatusEnum,
)
from llama_cloud.types.extraction.extract_agent import ExtractAgent
from llama_cloud.types.extraction.extract_config import ExtractConfig
from llama_cloud.types.extraction.extract_job import ExtractJob
from llama_cloud.types.extraction.extract_run import ExtractRun
from llama_cloud.types.extraction.extraction_agent_list_response import (
ExtractionAgentListResponse,
)
from llama_cloud.types.extraction.job_get_result_response import JobGetResultResponse
from llama_cloud.types.status_enum import StatusEnum
from ._deterministic import (
combined_seed,
fingerprint_file,
generate_data_from_schema,
hash_schema,
utcnow,
)
from ._deterministic import fingerprint_file
from .files import FakeFilesNamespace, StoredFile
from .matchers import RequestContext, RequestMatcher
@@ -143,30 +145,24 @@ class FakeExtractNamespace:
self._handle_update_agent,
namespace="extract",
)
server.add_route(
"PUT",
"/api/v1/extraction/extraction-agents/{agent_id}",
self._handle_update_agent,
namespace="extract",
)
server.add_route(
"GET",
"/api/v1/extraction/extraction-agents/{agent_id}",
self._handle_get_agent,
namespace="extract",
)
server.add_route(
"GET",
"/api/v1/extraction/extraction-agents/by-name/{name}",
self._handle_get_agent_by_name,
namespace="extract",
)
server.add_route(
"GET",
"/api/v1/extraction/extraction-agents",
self._handle_list_agents,
namespace="extract",
)
server.add_route(
"GET",
"/api/v1/extraction/extraction-agents/default",
self._handle_get_default_agent,
namespace="extract",
)
server.add_route(
"DELETE",
"/api/v1/extraction/extraction-agents/{agent_id}",
@@ -188,12 +184,6 @@ class FakeExtractNamespace:
)
self.routes["agent_job"] = agent_job_route
self.agent_job = agent_job_route
server.add_route(
"POST",
"/api/v1/extraction/jobs/batch",
self._handle_agent_job_batch,
namespace="extract",
)
server.add_route(
"GET",
"/api/v1/extraction/jobs",
@@ -206,6 +196,12 @@ class FakeExtractNamespace:
self._handle_get_job,
namespace="extract",
)
server.add_route(
"GET",
"/api/v1/extraction/jobs/{job_id}/result",
self._handle_get_job_result,
namespace="extract",
)
agent_run_route = server.add_route(
"GET",
"/api/v1/extraction/runs/by-job/{job_id}",
@@ -237,7 +233,7 @@ class FakeExtractNamespace:
# Handlers -------------------------------------------------------
def _handle_stateless_run(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request)
config = ExtractConfig.parse_obj(payload["config"])
config = ExtractConfig.model_validate(payload["config"])
data_schema = payload["data_schema"]
schema_hash = hash_schema(data_schema)
@@ -258,17 +254,17 @@ class FakeExtractNamespace:
)
stub = self._pop_stub(self._run_stubs, context)
job_status = StatusEnum.SUCCESS
run_status = ExtractState.SUCCESS
job_status: StatusEnum = "SUCCESS"
run_status = "SUCCESS"
metadata = {"deterministic": {"value": True}}
error = None
run_data = self._generate_run_data(data_schema, file_info.sha256)
if stub:
if stub.job_status:
job_status = StatusEnum(stub.job_status)
job_status = cast(StatusEnum, stub.job_status)
if stub.status:
run_status = ExtractState(stub.status)
run_status = stub.status
if stub.metadata:
metadata = stub.metadata
if stub.error:
@@ -285,18 +281,18 @@ class FakeExtractNamespace:
data_schema=data_schema,
file_info=file_info,
job_status=job_status,
run_status=run_status,
run_status=cast(StatusEnum, run_status),
metadata=metadata,
data=run_data,
error=error,
project_id=file_info.file.project_id,
)
return self._server.json_response(stored.job.dict())
return self._server.json_response(stored.job.model_dump())
def _handle_create_agent(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request)
name = payload["name"]
config = ExtractConfig.parse_obj(payload["config"])
config = ExtractConfig.model_validate(payload["config"])
data_schema = payload["data_schema"]
agent_id = self._server.new_id("agent")
agent = ExtractAgent(
@@ -313,7 +309,7 @@ class FakeExtractNamespace:
)
self._agents[agent_id] = agent
self._agents_by_name[name] = agent_id
return self._server.json_response(agent.dict())
return self._server.json_response(agent.model_dump())
def _handle_update_agent(self, request: httpx.Request) -> httpx.Response:
agent_id = request.url.path.split("/")[-1]
@@ -325,9 +321,9 @@ class FakeExtractNamespace:
agent = self._agents[agent_id]
config = payload.get("config", agent.config)
data_schema = payload.get("data_schema", agent.data_schema)
updated = agent.copy(
updated = agent.model_copy(
update={
"config": ExtractConfig.parse_obj(config)
"config": ExtractConfig.model_validate(config)
if isinstance(config, dict)
else config,
"data_schema": data_schema,
@@ -335,7 +331,7 @@ class FakeExtractNamespace:
}
)
self._agents[agent_id] = updated
return self._server.json_response(updated.dict())
return self._server.json_response(updated.model_dump())
def _handle_get_agent(self, request: httpx.Request) -> httpx.Response:
agent_id = request.url.path.split("/")[-1]
@@ -344,22 +340,13 @@ class FakeExtractNamespace:
return self._server.json_response(
{"detail": "Agent not found"}, status_code=404
)
return self._server.json_response(agent.dict())
def _handle_get_agent_by_name(self, request: httpx.Request) -> httpx.Response:
name = request.url.path.split("/")[-1]
agent_id = self._agents_by_name.get(name)
if not agent_id:
return self._server.json_response(
{"detail": "Agent not found"}, status_code=404
)
return self._server.json_response(self._agents[agent_id].dict())
return self._server.json_response(agent.model_dump())
def _handle_list_agents(self, request: httpx.Request) -> httpx.Response:
include_default = (
request.url.params.get("include_default", "false").lower() == "true"
)
agents = list(self._agents.values())
agents: ExtractionAgentListResponse = list(self._agents.values())
if include_default and not agents:
default_agent = self._build_ephemeral_agent(
ExtractConfig(),
@@ -367,18 +354,7 @@ class FakeExtractNamespace:
self._server.default_project_id,
)
agents.append(default_agent)
return self._server.json_response([agent.dict() for agent in agents])
def _handle_get_default_agent(self, request: httpx.Request) -> httpx.Response:
if self._agents:
agent = next(iter(self._agents.values()))
else:
agent = self._build_ephemeral_agent(
ExtractConfig(),
{"type": "object", "properties": {}},
self._server.default_project_id,
)
return self._server.json_response(agent.dict())
return self._server.json_response([agent.model_dump() for agent in agents])
def _handle_delete_agent(self, request: httpx.Request) -> httpx.Response:
agent_id = request.url.path.split("/")[-1]
@@ -410,7 +386,7 @@ class FakeExtractNamespace:
schema = payload.get("data_schema_override", agent.data_schema)
config_payload = payload.get("config_override", agent.config)
config = (
ExtractConfig.parse_obj(config_payload)
ExtractConfig.model_validate(config_payload)
if isinstance(config_payload, dict)
else config_payload
)
@@ -418,14 +394,14 @@ class FakeExtractNamespace:
stub = self._pop_agent_stub(
agent_id, RequestContext(request=request, json=payload)
)
job_status = StatusEnum.SUCCESS
run_status = ExtractState.SUCCESS
job_status: StatusEnum = "SUCCESS"
run_status = "SUCCESS"
error = None
if stub:
if stub.job_status:
job_status = StatusEnum(stub.job_status)
job_status = cast(StatusEnum, stub.job_status)
if stub.run_status:
run_status = ExtractState(stub.run_status)
run_status = stub.run_status
if stub.error:
error = stub.error
@@ -435,28 +411,28 @@ class FakeExtractNamespace:
data_schema=schema,
file_info=stored_file,
job_status=job_status,
run_status=run_status,
run_status=cast(StatusEnum, run_status),
metadata={"agent": {"value": agent.id}},
data=self._generate_run_data(schema, stored_file.sha256),
error=error,
project_id=agent.project_id,
)
return self._server.json_response(stored.job.dict())
return self._server.json_response(stored.job.model_dump())
def _handle_agent_job_batch(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request)
file_ids = payload.get("file_ids", [])
jobs = []
for file_id in file_ids:
request_body = payload.copy()
request_body["file_id"] = file_id
fake_request = copy.deepcopy(request)
fake_request._content = self._server.encode_json(request_body)
response = self._handle_agent_job(fake_request)
if response.status_code != 200:
return response
jobs.append(response.json())
return self._server.json_response(jobs)
def _handle_get_job_result(self, request: httpx.Request) -> httpx.Response:
job_id = request.url.path.split("/")[-2]
if job_id not in self._jobs:
return self._server.json_response(
{"detail": f"Job {job_id} not found"}, status_code=404
)
job = self._jobs[job_id]
response = JobGetResultResponse(
data=job.run.data,
extraction_agent_id=job.run.extraction_agent_id,
extraction_metadata=job.run.extraction_metadata or {},
run_id=job.run.id,
)
return self._server.json_response(response.model_dump())
def _handle_list_jobs(self, request: httpx.Request) -> httpx.Response:
agent_id = request.url.params.get("extraction_agent_id")
@@ -464,7 +440,7 @@ class FakeExtractNamespace:
for stored in self._jobs.values():
if agent_id and stored.job.extraction_agent.id != agent_id:
continue
items.append(stored.job.dict())
items.append(stored.job.model_dump())
return self._server.json_response(items)
def _handle_get_job(self, request: httpx.Request) -> httpx.Response:
@@ -474,7 +450,7 @@ class FakeExtractNamespace:
return self._server.json_response(
{"detail": "Job not found"}, status_code=404
)
return self._server.json_response(stored.job.dict())
return self._server.json_response(stored.job.model_dump())
def _handle_get_run_by_job(self, request: httpx.Request) -> httpx.Response:
job_id = request.url.path.split("/")[-1]
@@ -483,7 +459,7 @@ class FakeExtractNamespace:
return self._server.json_response(
{"detail": "Run not found"}, status_code=404
)
return self._server.json_response(stored.run.dict())
return self._server.json_response(stored.run.model_dump())
def _handle_get_run(self, request: httpx.Request) -> httpx.Response:
run_id = request.url.path.split("/")[-1]
@@ -492,7 +468,7 @@ class FakeExtractNamespace:
return self._server.json_response(
{"detail": "Run not found"}, status_code=404
)
return self._server.json_response(run.dict())
return self._server.json_response(run.model_dump())
def _handle_delete_run(self, request: httpx.Request) -> httpx.Response:
run_id = request.url.path.split("/")[-1]
@@ -507,20 +483,18 @@ class FakeExtractNamespace:
def _handle_list_runs(self, request: httpx.Request) -> httpx.Response:
agent_id = request.url.params.get("extraction_agent_id")
skip = int(request.url.params.get("skip", "0"))
limit = int(request.url.params.get("limit", "50"))
filtered = [
stored.run
for stored in self._jobs.values()
if not agent_id or stored.job.extraction_agent.id == agent_id
]
page = filtered[skip : skip + limit]
response = PaginatedExtractRunsResponse(
page = filtered[skip:]
response = AsyncPaginatedExtractRuns[ExtractRun](
items=page,
skip=skip,
limit=limit,
total=len(filtered),
)
return self._server.json_response(response.dict())
return self._server.json_response(response.model_dump())
# Internal helpers -----------------------------------------------
def _extract_file_info(
@@ -609,7 +583,7 @@ class FakeExtractNamespace:
data_schema: Dict[str, Any],
file_info: StoredFile,
job_status: StatusEnum,
run_status: ExtractState,
run_status: StatusEnum,
metadata: Dict[str, Any],
data: Any,
error: Optional[str],
@@ -631,7 +605,7 @@ class FakeExtractNamespace:
job_id=job_id,
file=file_info.file,
extraction_agent_id=agent.id,
status=run_status,
status=cast(Literal["CREATED", "PENDING", "SUCCESS", "ERROR"], run_status),
config=config,
data_schema=data_schema,
data=data,
+56 -98
View File
@@ -9,14 +9,14 @@ from urllib.parse import urlencode
import httpx
import respx
from llama_cloud.types import File as CloudFile
from llama_cloud.types import FileIdPresignedUrl, PresignedUrl
from llama_cloud.types.file_query_response import FileQueryResponse, Item
from llama_cloud.types.presigned_url import PresignedURL
from ._deterministic import (
fingerprint_file,
hash_chunks,
utcnow,
)
from .matchers import RequestContext, RequestMatcher
from .matchers import RequestMatcher
if TYPE_CHECKING:
from .server import FakeLlamaCloudServer
@@ -80,6 +80,20 @@ class FakeFilesNamespace:
def get(self, file_id: str) -> Optional[StoredFile]:
return self._files.get(file_id)
def preload_from_source(self, filename: str, content: bytes) -> str:
file_id = self._server.new_id("file")
name = filename
stored = self._build_file(
file_id=file_id,
name=name,
project_id=self._server.default_project_id,
organization_id=self._server.default_organization_id,
content=content,
external_file_id=None,
)
self._files[file_id] = stored
return file_id
def stub_upload(
self,
matcher: Optional[RequestMatcher],
@@ -94,16 +108,9 @@ class FakeFilesNamespace:
# Route registration ---------------------------------------------
def register(self) -> None:
server = self._server
server.add_route(
"PUT",
"/api/v1/files",
self._handle_generate_presigned_url,
namespace="files",
alias="generate_presigned_url",
)
upload_route = server.add_route(
"POST",
"/api/v1/files",
"/api/v1/beta/files",
self._handle_direct_upload,
namespace="files",
alias="upload",
@@ -111,32 +118,24 @@ class FakeFilesNamespace:
self.routes["upload"] = upload_route
get_route = server.add_route(
"GET",
"/api/v1/files/{file_id}",
self._handle_get_metadata,
"/api/v1/beta/files/{file_id}/content",
self._handle_read_content,
namespace="files",
alias="get",
)
self.routes["get"] = get_route
server.add_route(
"DELETE",
"/api/v1/files/{file_id}",
"/api/v1/beta/files/{file_id}",
self._handle_delete,
namespace="files",
)
server.add_route(
"GET",
"/api/v1/files/{file_id}/content",
self._handle_read_content,
"POST",
"/api/v1/beta/files/query",
self._handle_query,
namespace="files",
alias="read_content",
)
server.add_route(
"PUT",
"/upload/{file_id}",
self._handle_presigned_upload,
namespace="files",
base_urls=[self._upload_base_url],
alias="presigned_upload",
alias="query",
)
server.add_route(
"GET",
@@ -148,32 +147,6 @@ class FakeFilesNamespace:
)
# Handlers -------------------------------------------------------
def _handle_generate_presigned_url(self, request: httpx.Request) -> httpx.Response:
data = self._server.json(request)
now = utcnow()
file_id = self._server.new_id("file")
name = data.get("name") or f"upload-{file_id}.bin"
pending = PendingUpload(
file_id=file_id,
filename=name,
project_id=request.url.params.get(
"project_id", self._server.default_project_id
),
organization_id=request.url.params.get(
"organization_id", self._server.default_organization_id
),
external_file_id=data.get("external_file_id"),
expected_size=data.get("file_size"),
)
self._pending[file_id] = pending
presigned = FileIdPresignedUrl(
file_id=file_id,
url=f"{self._upload_base_url}/upload/{file_id}",
expires_at=now,
form_fields=None,
)
return self._server.json_response(presigned.dict())
def _handle_direct_upload(self, request: httpx.Request) -> httpx.Response:
file_bytes, filename = self._extract_multipart_file(request)
file_id = self._server.new_id("file")
@@ -190,15 +163,7 @@ class FakeFilesNamespace:
external_file_id=request.url.params.get("external_file_id"),
)
self._files[file_id] = stored
return self._server.json_response(stored.file.dict())
def _handle_get_metadata(self, request: httpx.Request) -> httpx.Response:
file_id = request.url.path.split("/")[-1]
if file_id not in self._files:
return self._server.json_response(
{"detail": "File not found"}, status_code=404
)
return self._server.json_response(self._files[file_id].file.dict())
return self._server.json_response(stored.file.model_dump())
def _handle_delete(self, request: httpx.Request) -> httpx.Response:
file_id = request.url.path.split("/")[-1]
@@ -212,47 +177,12 @@ class FakeFilesNamespace:
return self._server.json_response(
{"detail": "File not found"}, status_code=404
)
presigned = PresignedUrl(
presigned = PresignedURL(
url=f"{self._download_base_url}/files/{file_id}?{urlencode({'token': 'fake'})}",
expires_at=utcnow(),
form_fields=None,
)
return self._server.json_response(presigned.dict())
def _handle_presigned_upload(self, request: httpx.Request) -> httpx.Response:
file_id = request.url.path.split("/")[-1]
pending = self._pending.get(file_id)
context = RequestContext(
request=request,
json=None,
file_id=file_id,
filename=pending.filename if pending else None,
file_sha256=hash_chunks([request.content]),
)
for index, (matcher, status, body, once) in enumerate(list(self._upload_stubs)):
if context.matches(matcher):
if once:
self._upload_stubs.pop(index)
return self._server.json_response(body, status_code=status)
if pending is None:
return self._server.json_response(
{"detail": "Unknown file"}, status_code=404
)
stored = self._build_file(
file_id=file_id,
name=pending.filename,
project_id=pending.project_id,
organization_id=pending.organization_id,
content=request.content,
external_file_id=pending.external_file_id,
)
self._files[file_id] = stored
self._pending.pop(file_id, None)
return httpx.Response(204)
return self._server.json_response(presigned.model_dump())
def _handle_presigned_download(self, request: httpx.Request) -> httpx.Response:
file_id = request.url.path.split("/")[-1]
@@ -261,6 +191,34 @@ class FakeFilesNamespace:
return httpx.Response(404, json={"detail": "File not found"})
return httpx.Response(200, content=stored.content)
def _handle_query(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request)
files: list[StoredFile] = []
items: list[Item] = []
if payload.get("filter") is not None:
file_ids = payload["filter"].get("file_ids", [])
for file_id in self._files:
if file_id in file_ids:
files.append(self._files[file_id])
else:
files = list(self._files.values())
for f in files:
item = Item(
id=f.file.id,
name=f.file.name,
project_id=self._server.default_project_id,
expires_at=utcnow(),
external_file_id=f.file.external_file_id,
purpose=f.file.purpose,
last_modified_at=utcnow(),
file_type=f.file.file_type,
)
items.append(item)
response = FileQueryResponse(
items=items, next_page_token=None, total_size=len(items)
)
return self._server.json_response(response.model_dump())
# Internal helpers -----------------------------------------------
def _build_file(
self,
+177 -65
View File
@@ -1,12 +1,27 @@
from __future__ import annotations
import re
from copy import deepcopy
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Dict
import httpx
from llama_cloud.types.parsing_create_response import ParsingCreateResponse
from llama_cloud.types.parsing_get_response import (
Items,
ItemsPage,
ItemsPageStructuredResultPage,
ItemsPageStructuredResultPageItemTextItem,
Job,
Markdown,
MarkdownPage,
MarkdownPageMarkdownResultPage,
ParsingGetResponse,
Text,
TextPage,
)
from ._deterministic import generate_text_blob, hash_schema
from ._deterministic import generate_text_blob, hash_schema, utcnow
if TYPE_CHECKING:
from .server import FakeLlamaCloudServer
@@ -24,103 +39,191 @@ class ParseJobRecord:
class FakeParseNamespace:
def __init__(self, *, server: "FakeLlamaCloudServer") -> None:
self._server = server
self._jobs: Dict[str, ParseJobRecord] = {}
self._jobs: Dict[str, ParsingGetResponse] = {}
self.routes: Dict[str, Any] = {}
self.allowed_expands = ("text", "markdown", "items")
def register(self) -> None:
server = self._server
server.add_route(
"POST",
"/api/parsing/upload",
"/api/v2/parse/upload",
self._handle_upload,
namespace="parse",
)
server.add_route(
"GET",
"/api/parsing/job/{job_id}",
self._handle_job_status,
"/api/v2/parse/{job_id}",
self._handle_job_result,
namespace="parse",
)
server.add_route(
"GET",
"/api/parsing/job/{job_id}/result/{result_type}",
self._handle_job_result,
"POST",
"/api/v2/parse",
self._handle_file_id_source_url,
namespace="parse",
)
def _handle_upload(self, request: httpx.Request) -> httpx.Response:
file_bytes, filename, form_data = self._split_multipart(request)
_, filename, form_data = self._split_multipart(request)
job_id = self._server.new_id("parse-job")
seed_hash = hash_schema({"filename": filename, "form": form_data})
seed = int(seed_hash[:16], 16)
page_text = generate_text_blob(seed, sentences=3)
pages: list[Dict[str, Any]] = [
{
"page": index + 1,
"text": f"{page_text} (page {index + 1})",
"md": f"{page_text} (page {index + 1})",
"images": [],
"charts": [],
"tables": [],
"layout": [],
"items": [],
"status": "SUCCESS",
"links": [],
"width": 8.5,
"height": 11.0,
"parsingMode": "deterministic",
"structuredData": {},
"noStructuredContent": False,
"noTextContent": False,
"isAudioTranscript": False,
"durationInSeconds": None,
"slideSpeakerNotes": None,
}
for index in range(1)
item_pages: list[ItemsPage] = [
ItemsPageStructuredResultPage(
items=[
ItemsPageStructuredResultPageItemTextItem(
md=page_text, value=page_text, bBox=None, type="text"
)
],
page_height=1,
page_number=1,
page_width=1,
success=True,
)
]
result = {
"job_id": job_id,
"status": "SUCCESS",
"file_name": filename,
"is_done": True,
"pages": pages,
"job_metadata": {"job_pages": len(pages)},
"text": "\n\n".join(str(page["text"]) for page in pages),
"markdown": "\n\n".join(str(page["md"]) for page in pages),
"json": {"pages": pages},
}
record = ParseJobRecord(
job_id=job_id,
file_name=filename,
status="SUCCESS",
result=result,
content=file_bytes,
md_pages: list[MarkdownPage] = [
MarkdownPageMarkdownResultPage(
markdown=page_text,
page_number=1,
success=True,
)
]
txt_pages: list[TextPage] = [
TextPage(
text=page_text,
page_number=1,
)
]
record = ParsingGetResponse(
job=Job(
id=job_id,
status="COMPLETED",
project_id=self._server.default_project_id,
created_at=utcnow(),
updated_at=utcnow(),
error_message=None,
),
items=Items(pages=item_pages),
markdown=Markdown(pages=md_pages),
text=Text(pages=txt_pages),
)
self._jobs[job_id] = record
return self._server.json_response({"id": job_id})
response = ParsingCreateResponse(
id=job_id,
project_id=self._server.default_project_id,
status="COMPLETED",
created_at=utcnow(),
updated_at=utcnow(),
error_message=None,
)
return self._server.json_response(response.model_dump())
def _handle_job_status(self, request: httpx.Request) -> httpx.Response:
job_id = request.url.path.split("/")[-1]
job = self._jobs.get(job_id)
if not job:
def _handle_file_id_source_url(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request)
file_id = payload.get("file_id")
source_url = payload.get("source_url")
if file_id is not None:
file = self._server.files.get(file_id)
if file is None:
return self._server.json_response(
{"details": f"File {file_id} not found"},
status_code=404,
)
else:
seed_hash = file.sha256
elif source_url is not None:
response = self._get_file_from_source_url(source_url)
if isinstance(response, int):
return self._server.json_response(
{"details": f"Could not find file associated with {source_url}"},
status_code=response,
)
file_content, filename = response
file_id = self._server.files.preload_from_source(filename, file_content)
seed_hash = self._server.files._files[file_id].sha256
else:
return self._server.json_response(
{"detail": "Job not found"}, status_code=404
{
"details": "At least one between file_id and source_url should be not-null",
},
status_code=400,
)
return self._server.json_response({"id": job_id, "status": job.status})
job_id = self._server.new_id("parse-job")
seed = int(seed_hash[:16], 16)
page_text = generate_text_blob(seed, sentences=3)
item_pages: list[ItemsPage] = [
ItemsPageStructuredResultPage(
items=[
ItemsPageStructuredResultPageItemTextItem(
md=page_text, value=page_text, bBox=None, type="text"
)
],
page_height=1,
page_number=1,
page_width=1,
success=True,
)
]
md_pages: list[MarkdownPage] = [
MarkdownPageMarkdownResultPage(
markdown=page_text,
page_number=1,
success=True,
)
]
txt_pages: list[TextPage] = [
TextPage(
text=page_text,
page_number=1,
)
]
record = ParsingGetResponse(
job=Job(
id=job_id,
status="COMPLETED",
project_id=self._server.default_project_id,
created_at=utcnow(),
updated_at=utcnow(),
error_message=None,
),
items=Items(pages=item_pages),
markdown=Markdown(pages=md_pages),
text=Text(pages=txt_pages),
)
self._jobs[job_id] = record
response = ParsingCreateResponse(
id=job_id,
project_id=self._server.default_project_id,
status="COMPLETED",
created_at=utcnow(),
updated_at=utcnow(),
error_message=None,
)
return self._server.json_response(response.model_dump())
def _handle_job_result(self, request: httpx.Request) -> httpx.Response:
job_id = request.url.path.split("/")[-3]
job = self._jobs.get(job_id)
if not job:
job_id = request.url.path.split("/")[-1]
expandees = request.url.params.get_list("expand")
expandees = (
[e for e in expandees if e in self.allowed_expands]
if len(expandees) > 0
else ["items"]
)
job_response = self._jobs.get(job_id)
if not job_response:
return self._server.json_response(
{"detail": "Result not found"}, status_code=404
)
# Exclude job_id and file_name from result to avoid duplicate argument
# errors in JobResult.__init__ which passes these explicitly and spreads the dict
result = {
k: v for k, v in job.result.items() if k not in ("job_id", "file_name")
}
return self._server.json_response(result)
jb_resp_copy = deepcopy(job_response)
if "markdown" not in expandees:
jb_resp_copy.markdown = None
if "text" not in expandees:
jb_resp_copy.text = None
if "items" not in expandees:
jb_resp_copy.items = None
return self._server.json_response(jb_resp_copy.model_dump())
def _split_multipart(
self, request: httpx.Request
@@ -163,3 +266,12 @@ class FakeParseNamespace:
if not file_bytes:
raise ValueError("File part missing from multipart payload")
return file_bytes, filename, form_data
def _get_file_from_source_url(self, source_url: str) -> tuple[bytes, str] | int:
name = source_url.split("/")[-1]
with httpx.Client() as client:
response = client.get(source_url, follow_redirects=True)
if response.status_code >= 400:
return response.status_code
content = response.content
return content, name
+7 -6
View File
@@ -13,18 +13,19 @@ Here are your editing permissions, which you **MUST ALWAYS** follow:
</guidelines>
"""
import pytest
import warnings
import pytest
from extraction_review.clients import fake
# <edit>
from extraction_review.config import ExtractionSchema, EXTRACTED_DATA_COLLECTION
from extraction_review.process_file import workflow as process_file_workflow
from extraction_review.process_file import FileEvent
from workflows.events import StartEvent
from extraction_review.metadata_workflow import workflow as metadata_workflow
from extraction_review.config import EXTRACTED_DATA_COLLECTION, ExtractionSchema
from extraction_review.metadata_workflow import MetadataResponse
from extraction_review.metadata_workflow import workflow as metadata_workflow
from extraction_review.process_file import FileEvent
from extraction_review.process_file import workflow as process_file_workflow
from workflows.events import StartEvent
# </edit>
+171 -71
View File
@@ -1,12 +1,13 @@
"""Tests for the FakeAgentDataNamespace mock implementation."""
import pytest
from llama_cloud.core.api_error import ApiError
from llama_cloud_services.beta.agent_data import AsyncAgentDataClient
from pydantic import BaseModel, Field
from typing import cast
import pytest
from extraction_review.testing_utils import FakeLlamaCloudServer
from extraction_review.testing_utils._deterministic import hash_schema
from llama_cloud import AsyncLlamaCloud
from llama_cloud._exceptions import APIStatusError
from pydantic import BaseModel, Field
class Receipt(BaseModel):
@@ -22,93 +23,144 @@ def server():
@pytest.fixture
def client(server):
"""Provide an AsyncAgentDataClient configured for the fake server."""
return AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
def client(server) -> AsyncLlamaCloud:
"""Provide an AsyncLlamaCloud client configured for the fake server."""
return AsyncLlamaCloud(api_key="fake-api-key")
@pytest.mark.asyncio
async def test_create_item(server, client):
async def test_create_item(server, client: AsyncLlamaCloud):
"""Verify items can be created and have expected ID format."""
data = Receipt(merchant="Test Inc", total=1000)
item = await client.create_item(data)
item = await client.beta.agent_data.agent_data(
data=data.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
assert item.id == hash_schema(data)[:7]
assert item.data.merchant == data.merchant
assert item.data.total == data.total
assert item.data["merchant"] == data.merchant
assert item.data["total"] == data.total
assert item.collection == "extracted_data"
assert item.deployment_name == "extraction_agent"
@pytest.mark.asyncio
async def test_update_item(server, client):
async def test_update_item(server, client: AsyncLlamaCloud):
"""Verify items can be updated while preserving metadata."""
data = Receipt(merchant="Test Inc", total=1000)
item = await client.create_item(data)
item = await client.beta.agent_data.agent_data(
data=data.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
assert item.id is not None
updated_data = Receipt(merchant="Testing Inc", total=1100)
updated_item = await client.update_item(item_id=item.id, data=updated_data)
updated_item = await client.beta.agent_data.update(
item_id=item.id, data=updated_data.model_dump()
)
assert updated_item.data.merchant == updated_data.merchant
assert updated_item.data.total == updated_data.total
assert updated_item.data["merchant"] == updated_data.merchant
assert updated_item.data["total"] == updated_data.total
assert updated_item.id == item.id
assert updated_item.collection == item.collection
assert updated_item.deployment_name == item.deployment_name
@pytest.mark.asyncio
async def test_search_with_eq_filter(server, client):
async def test_search_with_eq_filter(server, client: AsyncLlamaCloud):
"""Verify search with equality filter returns matching items."""
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
item1 = await client.create_item(data1)
item2 = await client.create_item(data2)
await client.create_item(data3)
item1 = await client.beta.agent_data.agent_data(
data=data1.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
item2 = await client.beta.agent_data.agent_data(
data=data2.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
await client.beta.agent_data.agent_data(
data=data3.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
result = await client.search(filter={"merchant": {"eq": "Test Inc"}})
result = await client.beta.agent_data.search(
deployment_name="extraction_agent",
collection="extracted_data",
filter={"merchant": {"eq": "Test Inc"}},
)
assert result.total == 2
assert result.total_size == 2
assert any(item.id == item1.id for item in result.items)
assert any(item.id == item2.id for item in result.items)
assert all(item.data.merchant == "Test Inc" for item in result.items)
assert all(item.data["merchant"] == "Test Inc" for item in result.items)
@pytest.mark.asyncio
async def test_search_with_lt_filter(server, client):
async def test_search_with_lt_filter(server, client: AsyncLlamaCloud):
"""Verify search with less-than filter returns matching items."""
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
item1 = await client.create_item(data1)
await client.create_item(data2)
item3 = await client.create_item(data3)
item1 = await client.beta.agent_data.agent_data(
data=data1.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
await client.beta.agent_data.agent_data(
data=data2.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
item3 = await client.beta.agent_data.agent_data(
data=data3.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
result = await client.search(filter={"total": {"lt": 1200}})
result = await client.beta.agent_data.search(
deployment_name="extraction_agent",
collection="extracted_data",
filter={"total": {"lt": 1200}},
)
assert result.total == 2
assert result.total_size == 2
assert any(item.id == item1.id for item in result.items)
assert any(item.id == item3.id for item in result.items)
assert all(item.data.total < 1200 for item in result.items)
assert all(cast(int, item.data["total"]) < 1200 for item in result.items)
@pytest.mark.asyncio
async def test_aggregate_with_filter(server, client):
async def test_aggregate_with_filter(server, client: AsyncLlamaCloud):
"""Verify aggregation with filter groups correctly."""
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
await client.create_item(data1)
await client.create_item(data2)
await client.create_item(data3)
await client.beta.agent_data.agent_data(
data=data1.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
await client.beta.agent_data.agent_data(
data=data2.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
await client.beta.agent_data.agent_data(
data=data3.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
result = await client.aggregate(
result = await client.beta.agent_data.aggregate(
deployment_name="extraction_agent",
collection="extracted_data",
filter={"merchant": {"eq": "Test Inc"}},
group_by=["merchant"],
count=True,
@@ -118,21 +170,37 @@ async def test_aggregate_with_filter(server, client):
assert len(result.items) == 1
assert result.items[0].count == 2
assert result.items[0].first_item is not None
assert result.items[0].first_item.merchant == data1.merchant
assert result.items[0].first_item["merchant"] == data1.merchant
assert result.items[0].group_key == {"merchant": "Test Inc"}
@pytest.mark.asyncio
async def test_aggregate_without_filter(server, client):
async def test_aggregate_without_filter(server, client: AsyncLlamaCloud):
"""Verify aggregation without filter groups all items."""
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
await client.create_item(data1)
await client.create_item(data2)
await client.create_item(data3)
result = await client.aggregate(group_by=["merchant"], count=True)
await client.beta.agent_data.agent_data(
data=data1.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
await client.beta.agent_data.agent_data(
data=data2.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
await client.beta.agent_data.agent_data(
data=data3.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
result = await client.beta.agent_data.aggregate(
deployment_name="extraction_agent",
collection="extracted_data",
group_by=["merchant"],
count=True,
)
assert len(result.items) == 2
# First group: Test Inc (2 items)
@@ -144,70 +212,102 @@ async def test_aggregate_without_filter(server, client):
@pytest.mark.asyncio
async def test_get_item(server, client):
async def test_get_item(server, client: AsyncLlamaCloud):
"""Verify items can be retrieved by ID."""
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
item1 = await client.create_item(data1)
item2 = await client.create_item(data2)
item1 = await client.beta.agent_data.agent_data(
data=data1.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
item2 = await client.beta.agent_data.agent_data(
data=data2.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
retrieved = await client.get_item(item1.id)
assert item1.id is not None
retrieved = await client.beta.agent_data.get(item_id=item1.id)
assert retrieved.collection == item1.collection
assert retrieved.deployment_name == item1.deployment_name
assert retrieved.data.merchant == data1.merchant
assert retrieved.data.total == data1.total
assert retrieved.data["merchant"] == data1.merchant
assert retrieved.data["total"] == data1.total
assert item2.id is not None
# Non-existent ID should raise 404
with pytest.raises(ApiError) as exc_info:
await client.get_item(item2.id + "nonexistent")
with pytest.raises(APIStatusError) as exc_info:
await client.beta.agent_data.get(item_id=item2.id + "nonexistent")
assert exc_info.value.status_code == 404
assert exc_info.value.body == {"detail": f"No data with ID: {item2.id}nonexistent"}
@pytest.mark.asyncio
async def test_delete_by_id(server, client):
async def test_delete_by_id(server, client: AsyncLlamaCloud):
"""Verify items can be deleted by ID."""
data = Receipt(merchant="Test Inc", total=1300)
item = await client.create_item(data)
item = await client.beta.agent_data.agent_data(
data=data.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
assert item.id is not None
await client.delete_item(item.id)
await client.beta.agent_data.delete(item.id)
# Item should no longer exist
with pytest.raises(ApiError) as exc_info:
await client.get_item(item.id)
with pytest.raises(APIStatusError) as exc_info:
await client.beta.agent_data.get(item.id)
assert exc_info.value.status_code == 404
# Deleting again should also raise 404
with pytest.raises(ApiError) as exc_info:
await client.delete_item(item.id)
with pytest.raises(APIStatusError) as exc_info:
await client.beta.agent_data.delete(item.id)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_delete_by_query(server, client):
async def test_delete_by_query(server, client: AsyncLlamaCloud):
"""Verify items can be deleted by filter query."""
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
item1 = await client.create_item(data1)
item2 = await client.create_item(data2)
item3 = await client.create_item(data3)
item1 = await client.beta.agent_data.agent_data(
data=data1.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
item2 = await client.beta.agent_data.agent_data(
data=data2.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
item3 = await client.beta.agent_data.agent_data(
data=data3.model_dump(),
deployment_name="extraction_agent",
collection="extracted_data",
)
result = await client.delete(filter={"merchant": {"eq": "Test Inc"}})
result = await client.beta.agent_data.delete_by_query(
deployment_name="extraction_agent",
collection="extracted_data",
filter={"merchant": {"eq": "Test Inc"}},
)
assert result == 2
assert result.deleted_count == 2
# Deleted items should no longer exist
for item in (item1, item2):
with pytest.raises(ApiError) as exc_info:
await client.get_item(item.id)
assert item.id is not None
with pytest.raises(APIStatusError) as exc_info:
await client.beta.agent_data.get(item.id)
assert exc_info.value.status_code == 404
# Non-matching item should still exist
found = await client.get_item(item3.id)
assert item3.id is not None
found = await client.beta.agent_data.get(item3.id)
assert found.id == item3.id
+38 -36
View File
@@ -3,12 +3,9 @@
from pathlib import Path
import pytest
from llama_cloud import ExtractConfig
from llama_cloud.types import ExtractMode
from llama_cloud_services.extract import LlamaExtract
from llama_cloud_services.parse import LlamaParse
from extraction_review.testing_utils import FakeLlamaCloudServer
from llama_cloud import AsyncLlamaCloud
from llama_cloud.types.extraction.extract_config_param import ExtractConfigParam
from pydantic import BaseModel, Field
@@ -38,55 +35,60 @@ def _write_sample_file(tmp_path: Path, name: str, content: str) -> Path:
return target
def test_stateless_extract_is_deterministic(server, tmp_path):
@pytest.mark.asyncio
async def test_stateless_extract_is_deterministic(server, tmp_path):
"""Verify stateless extraction produces deterministic results."""
extractor = LlamaExtract(api_key="unit-test-key", verify=False)
config = ExtractConfig(extraction_mode=ExtractMode.FAST)
client = AsyncLlamaCloud(api_key="unit-test-key")
config = ExtractConfigParam(extraction_mode="FAST")
sample_path = _write_sample_file(
tmp_path, "receipt.txt", "Merchant: Lunar Bistro\nTotal: 123.45"
)
first_run = extractor.extract(Receipt, config, sample_path)
second_run = extractor.extract(Receipt, config, sample_path)
file_obj = await client.files.create(
file=sample_path,
purpose="extract",
external_file_id=str(sample_path),
)
first_run = await client.extraction.extract(
data_schema=Receipt.model_json_schema(),
config=config,
file_id=file_obj.id,
)
second_run = await client.extraction.extract(
data_schema=Receipt.model_json_schema(),
config=config,
file_id=file_obj.id,
)
assert first_run.status.value == "SUCCESS"
assert second_run.data == first_run.data
assert isinstance(first_run.data, dict)
assert "merchant" in first_run.data
assert server.extract.stateless_run.called
def test_agent_flow_uploads_and_processes_files(server, tmp_path):
@pytest.mark.asyncio
async def test_agent_flow_uploads_and_processes_files(server, tmp_path):
"""Verify agent flow correctly uploads files and processes them."""
extractor = LlamaExtract(api_key="unit-test-key", verify=False)
config = ExtractConfig(extraction_mode=ExtractMode.FAST)
agent = extractor.create_agent(
name="unit-test-agent", data_schema=Receipt, config=config
client = AsyncLlamaCloud(api_key="unit-test-key")
config = ExtractConfigParam(extraction_mode="FAST")
agent = await client.extraction.extraction_agents.create(
name="unit-test-agent", data_schema=Receipt.model_json_schema(), config=config
)
sample_path = _write_sample_file(
tmp_path, "contract.pdf", "Agreement between parties."
)
run = agent.extract(sample_path)
file_obj = await client.files.create(
file=sample_path,
purpose="extract",
external_file_id=str(sample_path),
)
run = await client.extraction.jobs.extract(
extraction_agent_id=agent.id,
file_id=file_obj.id,
)
assert run.status.value == "SUCCESS"
assert isinstance(run.data, dict)
assert "merchant" in run.data
uploaded_bytes = server.files.read(run.file.id)
assert uploaded_bytes.startswith(b"Agreement")
assert server.extract.agent_job.called
assert server.extract.agent_run.called
def test_parse_load_data_returns_documents(server, tmp_path):
"""Verify LlamaParse returns documents with expected content."""
parser = LlamaParse(
api_key="unit-test-key", base_url=FakeLlamaCloudServer.DEFAULT_BASE_URL
)
sample_path = _write_sample_file(
tmp_path, "report.pdf", "Executive summary of quarterly goals."
)
documents = parser.load_data(sample_path)
assert documents
assert "(page 1)" in documents[0].text
+79 -18
View File
@@ -1,9 +1,12 @@
"""Tests for the FakeFilesNamespace mock implementation."""
import pytest
import httpx
from urllib.parse import urlencode
import httpx
import pytest
from extraction_review.testing_utils import FakeLlamaCloudServer
from llama_cloud import APIStatusError, AsyncLlamaCloud
from llama_cloud.types.file_query_params import Filter
@pytest.fixture
@@ -13,7 +16,8 @@ def server():
yield srv
def test_preload_and_read(server, tmp_path):
@pytest.mark.asyncio
async def test_preload_and_download_as_presigned_url(server, tmp_path):
"""Verify files can be preloaded and read back."""
test_file = tmp_path / "test_file.txt"
test_file.write_bytes(b"test content here")
@@ -23,34 +27,91 @@ def test_preload_and_read(server, tmp_path):
content = server.files.read(file_id)
assert content == b"test content here"
response = httpx.get(f"{server.DEFAULT_BASE_URL}/api/v1/files/{file_id}")
assert response.status_code == 200
metadata = response.json()
assert metadata["id"] == file_id
assert metadata["name"] == "test_file.txt"
client = AsyncLlamaCloud(api_key="fake-api-key")
presigned_url = await client.files.get(
file_id=file_id,
)
assert (
presigned_url.url
== f"{server._download_base_url}/files/{file_id}?{urlencode({'token': 'fake'})}"
)
response = httpx.get(presigned_url.url)
assert response.content == b"test content here"
def test_not_found_returns_404(server):
@pytest.mark.asyncio
async def test_not_found_returns_404(server):
"""Verify non-existent file returns 404."""
response = httpx.get(f"{server.DEFAULT_BASE_URL}/api/v1/files/nonexistent-file-id")
assert response.status_code == 404
client = AsyncLlamaCloud(api_key="fake-api-key")
with pytest.raises(APIStatusError) as exc_info:
await client.files.get(
"does-not-exist",
)
assert exc_info.value.status_code == 404
def test_delete_file(server, tmp_path):
@pytest.mark.asyncio
async def test_delete_file(server, tmp_path):
"""Verify files can be deleted."""
test_file = tmp_path / "to_delete.txt"
test_file.write_bytes(b"delete me")
file_id = server.files.preload(path=test_file)
client = AsyncLlamaCloud(api_key="fake-api-key")
# File should exist
response = httpx.get(f"{server.DEFAULT_BASE_URL}/api/v1/files/{file_id}")
assert response.status_code == 200
response = await client.files.get(file_id)
assert file_id in response.url
# Delete the file
delete_response = httpx.delete(f"{server.DEFAULT_BASE_URL}/api/v1/files/{file_id}")
assert delete_response.status_code == 200
await client.files.delete(
file_id,
)
# File should no longer exist
response = httpx.get(f"{server.DEFAULT_BASE_URL}/api/v1/files/{file_id}")
assert response.status_code == 404
with pytest.raises(APIStatusError) as exc_info:
await client.files.get(
file_id,
)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_files_native_upload(server, tmp_path):
"""Verify that the client can natively upload the files without having to pass through server.preload"""
test_file = tmp_path / "test_file.txt"
test_file.write_bytes(b"test content here")
client = AsyncLlamaCloud(api_key="fake-api-key")
file_obj = await client.files.create(
file=test_file,
purpose="parse",
external_file_id=str(test_file),
)
assert isinstance(file_obj.id, str)
assert file_obj.file_type == "application/octet-stream"
assert file_obj.id.startswith("file_")
@pytest.mark.asyncio
async def test_files_query_by_id(server, tmp_path):
"""Test that you can upload and query files selecting them by file ID"""
test_file_1 = tmp_path / "test_file1.txt"
test_file_1.write_bytes(b"test content here 1")
test_file_2 = tmp_path / "test_file2.txt"
test_file_2.write_bytes(b"test content here 2")
client = AsyncLlamaCloud(api_key="fake-api-key")
file_obj_1 = await client.files.create(
file=test_file_1,
purpose="parse",
external_file_id=str(test_file_1),
)
await client.files.create(
file=test_file_2,
purpose="parse",
external_file_id=str(test_file_1),
)
response = await client.files.query(filter=Filter(file_ids=[file_obj_1.id]))
assert len(response.items) == 1
assert response.total_size == 1
assert response.items[0].id == file_obj_1.id
+151 -125
View File
@@ -1,153 +1,179 @@
"""Tests for the FakeParseNamespace mock implementation."""
import pytest
import httpx
import respx
from extraction_review.testing_utils import FakeLlamaCloudServer
from llama_cloud import APIStatusError, AsyncLlamaCloud
@pytest.fixture
def server():
"""Provide a server with parse namespace enabled."""
with FakeLlamaCloudServer(namespaces=["parse"]) as srv:
with FakeLlamaCloudServer() as srv:
yield srv
def _make_multipart_body(
filename: str, content: bytes = b"fake pdf content"
) -> tuple[bytes, str]:
"""Create a multipart form body with the given filename."""
boundary = "----TestBoundary123"
body = (
(
f"------{boundary}\r\n"
f'Content-Disposition: form-data; name="file"; filename="{filename}"\r\n'
f"Content-Type: application/pdf\r\n"
f"\r\n"
).encode()
+ content
+ f"\r\n------{boundary}--\r\n".encode()
@pytest.fixture()
def client() -> AsyncLlamaCloud:
return AsyncLlamaCloud(api_key="fake-api-key")
@pytest.fixture()
def data() -> tuple[str, bytes, str]:
with open("tests/files/test.pdf", "rb") as f:
content = f.read()
return ("tests/files/test.pdf", content, "application/pdf")
@pytest.mark.asyncio
async def test_parse_with_upload_file(
server: FakeLlamaCloudServer, client: AsyncLlamaCloud, data: tuple[str, bytes, str]
) -> None:
job_create = await client.parsing.create(
tier="fast",
version="latest",
upload_file=data,
)
return body, boundary
def test_job_result_excludes_duplicate_fields(server):
"""Verify job result doesn't include job_id/file_name to avoid duplicate arg errors.
The llama_cloud_services library's JobResult.__init__ passes job_id and
file_name as explicit args AND spreads the job_result dict. If both are
present, it causes "multiple values for keyword argument" errors.
"""
body, boundary = _make_multipart_body("test.pdf")
upload_response = httpx.post(
f"{server.DEFAULT_BASE_URL}/api/parsing/upload",
content=body,
headers={"Content-Type": f"multipart/form-data; boundary=----{boundary}"},
assert job_create.error_message is None
assert job_create.status == "COMPLETED"
assert job_create.project_id == server.default_project_id
job_response = await client.parsing.get(
job_id=job_create.id, expand=["text", "markdown", "items"]
)
assert upload_response.status_code == 200
job_id = upload_response.json()["id"]
assert job_response.job.id == job_create.id
assert job_response.job.status == job_create.status
assert job_response.job.project_id == job_create.project_id
assert job_response.items is not None
assert job_response.markdown is not None
assert job_response.text is not None
result_response = httpx.get(
f"{server.DEFAULT_BASE_URL}/api/parsing/job/{job_id}/result/json"
@pytest.mark.asyncio
async def test_parse_with_different_expand(
server: FakeLlamaCloudServer, client: AsyncLlamaCloud, data: tuple[str, bytes, str]
) -> None:
job_create = await client.parsing.create(
tier="fast",
version="latest",
upload_file=data,
)
assert result_response.status_code == 200
result = result_response.json()
# These fields should NOT be in the result to avoid duplicate argument errors
assert "job_id" not in result, "job_id should be excluded from result"
assert "file_name" not in result, "file_name should be excluded from result"
# But other expected fields should still be present
assert "status" in result
assert "is_done" in result
assert "pages" in result
job_response = await client.parsing.get(job_id=job_create.id, expand=["text"])
assert job_response.items is None
assert job_response.markdown is None
assert job_response.text is not None
job_response = await client.parsing.get(job_id=job_create.id, expand=["markdown"])
assert job_response.items is None
assert job_response.markdown is not None
assert job_response.text is None
job_response = await client.parsing.get(job_id=job_create.id, expand=["items"])
assert job_response.items is not None
assert job_response.markdown is None
assert job_response.text is None
# no expands -> defaul to items
job_response = await client.parsing.get(job_id=job_create.id)
assert job_response.items is not None
assert job_response.markdown is None
assert job_response.text is None
def test_filename_with_double_quotes(server):
"""Verify filename is correctly parsed from double-quoted Content-Disposition."""
body, boundary = _make_multipart_body("my_document.pdf")
response = httpx.post(
f"{server.DEFAULT_BASE_URL}/api/parsing/upload",
content=body,
headers={"Content-Type": f"multipart/form-data; boundary=----{boundary}"},
@pytest.mark.asyncio
async def test_parse_with_file_id(
server: FakeLlamaCloudServer, client: AsyncLlamaCloud, data: tuple[str, bytes, str]
) -> None:
file_name, _, _ = data
file_obj = await client.files.create(
file=file_name,
purpose="parse",
external_file_id=file_name,
)
assert response.status_code == 200
job_id = response.json()["id"]
job = server.parse._jobs[job_id]
assert job.file_name == "my_document.pdf"
def test_filename_with_single_quotes(server):
"""Verify filename is correctly parsed from single-quoted Content-Disposition."""
boundary = "----TestBoundary789"
body = (
f"------{boundary}\r\n"
f"Content-Disposition: form-data; name=\"file\"; filename='single_quoted.pdf'\r\n"
f"Content-Type: application/pdf\r\n"
f"\r\n"
f"fake pdf content\r\n"
f"------{boundary}--\r\n"
).encode()
response = httpx.post(
f"{server.DEFAULT_BASE_URL}/api/parsing/upload",
content=body,
headers={"Content-Type": f"multipart/form-data; boundary=----{boundary}"},
job_create = await client.parsing.create(
tier="fast",
version="latest",
file_id=file_obj.id,
)
assert response.status_code == 200
job_id = response.json()["id"]
job = server.parse._jobs[job_id]
assert job.file_name == "single_quoted.pdf"
def test_filename_does_not_capture_subsequent_headers(server):
"""Verify filename parsing stops at header boundary, not capturing Content-Type."""
body, boundary = _make_multipart_body("test.pdf")
response = httpx.post(
f"{server.DEFAULT_BASE_URL}/api/parsing/upload",
content=body,
headers={"Content-Type": f"multipart/form-data; boundary=----{boundary}"},
assert job_create.error_message is None
assert job_create.status == "COMPLETED"
assert job_create.project_id == server.default_project_id
job_response = await client.parsing.get(
job_id=job_create.id, expand=["text", "markdown", "items"]
)
assert response.status_code == 200
job_id = response.json()["id"]
job = server.parse._jobs[job_id]
# Should be just "test.pdf", not "test.pdf\r\nContent-Type: application/pdf"
assert job.file_name == "test.pdf"
assert "Content-Type" not in job.file_name
assert job_response.job.id == job_create.id
assert job_response.job.status == job_create.status
assert job_response.job.project_id == job_create.project_id
assert job_response.items is not None
assert job_response.markdown is not None
assert job_response.text is not None
def test_job_status_endpoint(server):
"""Verify job status endpoint returns correct status."""
body, boundary = _make_multipart_body("status_test.pdf")
@pytest.mark.asyncio
async def test_parse_with_file_id_file_not_found(
server: FakeLlamaCloudServer,
client: AsyncLlamaCloud,
) -> None:
with pytest.raises(APIStatusError) as exc_info:
await client.parsing.create(
tier="fast",
version="latest",
file_id="does-not-exist",
)
assert exc_info.value.status_code == 404
upload_response = httpx.post(
f"{server.DEFAULT_BASE_URL}/api/parsing/upload",
content=body,
headers={"Content-Type": f"multipart/form-data; boundary=----{boundary}"},
@pytest.mark.asyncio
async def test_parse_without_fileid_or_sourceurl(
server: FakeLlamaCloudServer,
client: AsyncLlamaCloud,
) -> None:
with pytest.raises(APIStatusError) as exc_info:
await client.parsing.create(
tier="fast",
version="latest",
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
@respx.mock(assert_all_mocked=False)
async def test_parse_with_source_url(
server: FakeLlamaCloudServer,
client: AsyncLlamaCloud,
) -> None:
job_create = await client.parsing.create(
tier="fast",
version="latest",
source_url="https://pdfobject.com/pdf/sample.pdf",
)
job_id = upload_response.json()["id"]
status_response = httpx.get(f"{server.DEFAULT_BASE_URL}/api/parsing/job/{job_id}")
assert status_response.status_code == 200
status = status_response.json()
assert status["id"] == job_id
assert status["status"] == "SUCCESS"
def test_job_not_found_returns_404(server):
"""Verify non-existent job returns 404."""
response = httpx.get(
f"{server.DEFAULT_BASE_URL}/api/parsing/job/nonexistent-job-id"
assert job_create.error_message is None
assert job_create.status == "COMPLETED"
assert job_create.project_id == server.default_project_id
job_response = await client.parsing.get(
job_id=job_create.id, expand=["text", "markdown", "items"]
)
assert response.status_code == 404
assert job_response.job.id == job_create.id
assert job_response.job.status == job_create.status
assert job_response.job.project_id == job_create.project_id
assert job_response.items is not None
assert job_response.markdown is not None
assert job_response.text is not None
result_response = httpx.get(
f"{server.DEFAULT_BASE_URL}/api/parsing/job/nonexistent-job-id/result/json"
@pytest.mark.asyncio
async def test_parse_e2e(
server: FakeLlamaCloudServer, client: AsyncLlamaCloud, data: tuple[str, bytes, str]
) -> None:
file_name, _, _ = data
file_obj = await client.files.create(
file=file_name,
purpose="parse",
external_file_id=file_name,
)
assert result_response.status_code == 404
result = await client.parsing.parse(
file_id=file_obj.id,
expand=["markdown"],
tier="agentic",
version="latest",
)
assert result.markdown is not None
assert len(result.markdown.pages) == 1
assert hasattr(result.markdown.pages[0], "markdown")
assert isinstance(result.markdown.pages[0].markdown, str) # type: ignore