diff --git a/pyproject.toml b/pyproject.toml index 39d4608..d6cc209 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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] diff --git a/src/extraction_review/clients.py b/src/extraction_review/clients.py index 66cd244..be44802 100644 --- a/src/extraction_review/clients.py +++ b/src/extraction_review/clients.py @@ -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) diff --git a/src/extraction_review/config.py b/src/extraction_review/config.py index 22e377c..eabfe46 100644 --- a/src/extraction_review/config.py +++ b/src/extraction_review/config.py @@ -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, diff --git a/src/extraction_review/process_file.py b/src/extraction_review/process_file.py index c51b609..f0f9e36 100644 --- a/src/extraction_review/process_file.py +++ b/src/extraction_review/process_file.py @@ -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)) diff --git a/src/extraction_review/testing_utils/_deterministic.py b/src/extraction_review/testing_utils/_deterministic.py index 8fc1dbf..f4ca1d1 100644 --- a/src/extraction_review/testing_utils/_deterministic.py +++ b/src/extraction_review/testing_utils/_deterministic.py @@ -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) diff --git a/src/extraction_review/testing_utils/classify.py b/src/extraction_review/testing_utils/classify.py index f41e6e2..8b3228f 100644 --- a/src/extraction_review/testing_utils/classify.py +++ b/src/extraction_review/testing_utils/classify.py @@ -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) ) diff --git a/src/extraction_review/testing_utils/extract.py b/src/extraction_review/testing_utils/extract.py index d8d36b8..e16c549 100644 --- a/src/extraction_review/testing_utils/extract.py +++ b/src/extraction_review/testing_utils/extract.py @@ -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, diff --git a/src/extraction_review/testing_utils/files.py b/src/extraction_review/testing_utils/files.py index 91642ce..ed5c960 100644 --- a/src/extraction_review/testing_utils/files.py +++ b/src/extraction_review/testing_utils/files.py @@ -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, diff --git a/src/extraction_review/testing_utils/parse.py b/src/extraction_review/testing_utils/parse.py index dd54148..ae50bde 100644 --- a/src/extraction_review/testing_utils/parse.py +++ b/src/extraction_review/testing_utils/parse.py @@ -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 diff --git a/tests/test_workflow.py b/tests/test_workflow.py index 0eb65b5..d475011 100644 --- a/tests/test_workflow.py +++ b/tests/test_workflow.py @@ -13,18 +13,19 @@ Here are your editing permissions, which you **MUST ALWAYS** follow: """ -import pytest import warnings +import pytest from extraction_review.clients import fake # -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 + # diff --git a/tests/testing_utils/test_agent_data.py b/tests/testing_utils/test_agent_data.py index 57d4e29..1fda332 100644 --- a/tests/testing_utils/test_agent_data.py +++ b/tests/testing_utils/test_agent_data.py @@ -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 diff --git a/tests/testing_utils/test_extract.py b/tests/testing_utils/test_extract.py index 8672e9d..fefe154 100644 --- a/tests/testing_utils/test_extract.py +++ b/tests/testing_utils/test_extract.py @@ -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 diff --git a/tests/testing_utils/test_files.py b/tests/testing_utils/test_files.py index 26f0476..dba1308 100644 --- a/tests/testing_utils/test_files.py +++ b/tests/testing_utils/test_files.py @@ -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 diff --git a/tests/testing_utils/test_parse.py b/tests/testing_utils/test_parse.py index 4e20713..b2fae3f 100644 --- a/tests/testing_utils/test_parse.py +++ b/tests/testing_utils/test_parse.py @@ -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