mirror of
https://github.com/run-llama/template-workflow-extract-basic.git
synced 2026-07-19 18:53:47 -04:00
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:
committed by
GitHub
parent
8513a08de0
commit
adda5d6b32
+1
-1
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user