mirror of
https://github.com/run-llama/llama_cloud_services.git
synced 2026-07-21 03:55:22 -04:00
Compare commits
53 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 9ce2044995 | |||
| 90d1608a71 | |||
| 2448a42b90 | |||
| c75a900174 | |||
| 2fb7adfe0e | |||
| dc82270724 | |||
| d880a48dd0 | |||
| 7567e8b45e | |||
| 0d59a90151 | |||
| 98ad550b1a | |||
| b58f43ce9f | |||
| acf6adcd91 | |||
| daf6576c3c | |||
| 8caa4defa6 | |||
| 26918b8de4 | |||
| 6fb5ebe2f9 | |||
| c0aa67995b | |||
| 9f841f8328 | |||
| 99c75eece9 | |||
| 57d2586ee3 | |||
| 4280a43ec8 | |||
| 7f1082bbb2 | |||
| 57cfc45804 | |||
| 30e8913875 | |||
| 0ce6d4d7a4 | |||
| 584ba8d48e | |||
| 925805ee11 | |||
| 76fb73c971 | |||
| 6d19ea9ac0 | |||
| 90431090e9 | |||
| 6dff35b204 | |||
| e634c7978d | |||
| 7a9e99bba2 | |||
| efcdd4405b | |||
| bf3614690f | |||
| 7463e00da3 | |||
| cbe9de0c57 | |||
| a023507d42 | |||
| e48f544ddc | |||
| 4aa7ad5642 | |||
| c39cdbcd01 | |||
| 71eaa8bcc6 | |||
| 1e1cbdfc79 | |||
| cc8af4a43a | |||
| 43fbd48ab8 | |||
| 5ec66e9452 | |||
| 211521c82e | |||
| 4ddaab1efb | |||
| 53e5ce2e83 | |||
| 9f4bd1cb64 | |||
| 456863752b | |||
| c2dc34bbd6 | |||
| fcabb04baf |
@@ -0,0 +1,11 @@
|
||||
# Please see the documentation for all configuration options:
|
||||
# https://docs.github.com/github/administering-a-repository/configuration-options-for-dependency-updates
|
||||
# and
|
||||
# https://docs.github.com/code-security/dependabot/dependabot-version-updates/configuration-options-for-the-dependabot.yml-file
|
||||
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "github-actions"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
@@ -21,9 +21,9 @@ jobs:
|
||||
os: [ubuntu-latest, windows-latest]
|
||||
python-version: ["3.9"]
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v4
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install Poetry
|
||||
|
||||
@@ -1,14 +1,3 @@
|
||||
# For most projects, this workflow file will not need changing; you simply need
|
||||
# to commit it to your repository.
|
||||
#
|
||||
# You may wish to alter this file to override the set of languages analyzed,
|
||||
# or to provide custom queries or build logic.
|
||||
#
|
||||
# ******** NOTE ********
|
||||
# We have attempted to detect the languages in your repository. Please check
|
||||
# the `language` matrix defined below to confirm you have the correct set of
|
||||
# supported CodeQL languages.
|
||||
#
|
||||
name: "CodeQL"
|
||||
|
||||
on:
|
||||
@@ -28,54 +17,25 @@ jobs:
|
||||
# - https://gh.io/supported-runners-and-hardware-resources
|
||||
# - https://gh.io/using-larger-runners
|
||||
# Consider using larger runners for possible analysis time improvements.
|
||||
runs-on: ${{ (matrix.language == 'swift' && 'macos-latest') || 'ubuntu-latest' }}
|
||||
timeout-minutes: ${{ (matrix.language == 'swift' && 120) || 360 }}
|
||||
runs-on: "ubuntu-latest"
|
||||
timeout-minutes: 360
|
||||
permissions:
|
||||
actions: read
|
||||
contents: read
|
||||
security-events: write
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
language: ["python"]
|
||||
# CodeQL supports [ 'cpp', 'csharp', 'go', 'java', 'javascript', 'python', 'ruby', 'swift' ]
|
||||
# Use only 'java' to analyze code written in Java, Kotlin or both
|
||||
# Use only 'javascript' to analyze code written in JavaScript, TypeScript or both
|
||||
# Learn more about CodeQL language support at https://aka.ms/codeql-docs/language-support
|
||||
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v3
|
||||
uses: actions/checkout@v4
|
||||
|
||||
# Initializes the CodeQL tools for scanning.
|
||||
- name: Initialize CodeQL
|
||||
uses: github/codeql-action/init@v2
|
||||
uses: github/codeql-action/init@v3
|
||||
with:
|
||||
languages: ${{ matrix.language }}
|
||||
# If you wish to specify custom queries, you can do so here or in a config file.
|
||||
# By default, queries listed here will override any specified in a config file.
|
||||
# Prefix the list here with "+" to use these queries and those in the config file.
|
||||
|
||||
# For more details on CodeQL's query packs, refer to: https://docs.github.com/en/code-security/code-scanning/automatically-scanning-your-code-for-vulnerabilities-and-errors/configuring-code-scanning#using-queries-in-ql-packs
|
||||
# queries: security-extended,security-and-quality
|
||||
|
||||
# Autobuild attempts to build any compiled languages (C/C++, C#, Go, Java, or Swift).
|
||||
# If this step fails, then you should remove it and run the build manually (see below)
|
||||
- name: Autobuild
|
||||
uses: github/codeql-action/autobuild@v2
|
||||
|
||||
# ℹ️ Command-line programs to run using the OS shell.
|
||||
# 📚 See https://docs.github.com/en/actions/using-workflows/workflow-syntax-for-github-actions#jobsjob_idstepsrun
|
||||
|
||||
# If the Autobuild fails above, remove it and uncomment the following three lines.
|
||||
# modify them (or add more) to build your code if your project, please refer to the EXAMPLE below for guidance.
|
||||
|
||||
# - run: |
|
||||
# echo "Run, Build Application using script"
|
||||
# ./location_of_script_within_repo/buildscript.sh
|
||||
languages: python
|
||||
dependency-caching: true
|
||||
|
||||
- name: Perform CodeQL Analysis
|
||||
uses: github/codeql-action/analyze@v2
|
||||
uses: github/codeql-action/analyze@v3
|
||||
with:
|
||||
category: "/language:${{matrix.language}}"
|
||||
category: "/language:python"
|
||||
|
||||
@@ -18,11 +18,11 @@ jobs:
|
||||
matrix:
|
||||
python-version: ["3.9"]
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: ${{ github.event_name == 'pull_request' && 2 || 0 }}
|
||||
- name: Set up python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v4
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install Poetry
|
||||
|
||||
@@ -18,9 +18,9 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up python ${{ env.PYTHON_VERSION }}
|
||||
uses: actions/setup-python@v4
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ env.PYTHON_VERSION }}
|
||||
|
||||
@@ -39,10 +39,18 @@ jobs:
|
||||
pypi_token: ${{ secrets.LLAMA_PARSE_PYPI_TOKEN }}
|
||||
poetry_install_options: "--without dev"
|
||||
|
||||
- name: Wait for PyPI to update
|
||||
run: |
|
||||
sleep 120
|
||||
|
||||
- name: Update llama-parse lock file
|
||||
run: |
|
||||
cd llama_parse && poetry lock
|
||||
|
||||
- name: Build and publish llama-parse
|
||||
uses: JRubics/poetry-publish@v2.1
|
||||
with:
|
||||
working_directory: "llama_parse"
|
||||
package_directory: "./llama_parse"
|
||||
pypi_token: ${{ secrets.LLAMA_PARSE_PYPI_TOKEN }}
|
||||
poetry_install_options: "--without dev"
|
||||
|
||||
|
||||
@@ -19,11 +19,11 @@ jobs:
|
||||
matrix:
|
||||
python-version: ["3.9", "3.10", "3.11", "3.12"]
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- name: Set up python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v4
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- name: Install Poetry
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
After Width: | Height: | Size: 3.3 MiB |
@@ -0,0 +1 @@
|
||||
sec_form_4_dump.json
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
After Width: | Height: | Size: 202 KiB |
@@ -0,0 +1,440 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Extract Data from Financial Reports - with Citations and Reasoning\n",
|
||||
"\n",
|
||||
"Given complex files like financial reports, contracts, invoices etc, Llama Extract allows you to make use of an LLM to extract the information relevant to you, in a structured format.\n",
|
||||
"\n",
|
||||
"In this example, we'll be using [LlamaExtract](https://docs.cloud.llamaindex.ai/llamaextract/getting_started?utm_campaign=extract&utm_medium=recipe) to extract structured data from an SEC filing (specifically, the filing by Nvidia for fiscal year 2025).\n",
|
||||
"\n",
|
||||
"On top of simple data extraction, we'll ask our extraction agent to provide citations and reasoning for each extracted field. This allows us to:\n",
|
||||
"- Confirm the accuracy of the extracted field\n",
|
||||
"- Understand the reasoning behind why the LLM extracted a given piece of information\n",
|
||||
"- This last point allows us an opportunity to adjust the system prompt or field descriptions and improve on results where needed.\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"The example we go through below is also replicable within Llama Cloud as well, where you will also be able to pick between a number of pre-defined schemas, instead of building your own."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!pip install llama-cloud-services"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Connect to Llama Cloud\n",
|
||||
"\n",
|
||||
"To get started, make sure you provide your [Llama Cloud](https://cloud.llamaindex.ai?utm_campaign=extract&utm_medium=recipe) API key."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Enter your Llama Cloud API Key: ··········\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"from getpass import getpass\n",
|
||||
"\n",
|
||||
"if \"LLAMA_CLOUD_API_KEY\" not in os.environ:\n",
|
||||
" os.environ[\"LLAMA_CLOUD_API_KEY\"] = getpass(\"Enter your Llama Cloud API Key: \")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Extract Data with Llama Extract Agent"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"No project_id provided, fetching default project.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from llama_cloud_services import LlamaExtract\n",
|
||||
"\n",
|
||||
"# Optionally, provide your project id, if not, it will use the 'Default' project\n",
|
||||
"llama_extract = LlamaExtract()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Provide Your Custom Schema\n",
|
||||
"\n",
|
||||
"When using LlamaExtract via the API, you provide your own schema that describes what you want extracted from files and data provided to your agent. Here, we are essentially building an SEC filings extraction agent."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from pydantic import BaseModel, Field\n",
|
||||
"from enum import Enum\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"class FilingType(str, Enum):\n",
|
||||
" ten_k = \"10 K\"\n",
|
||||
" ten_q = \"10-Q\"\n",
|
||||
" ten_ka = \"10-K/A\"\n",
|
||||
" ten_qa = \"10-Q/A\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"class FinancialReport(BaseModel):\n",
|
||||
" company_name: str = Field(description=\"The name of the company\")\n",
|
||||
" description: str = Field(\n",
|
||||
" description=\"Short description of the filing and what it contains\"\n",
|
||||
" )\n",
|
||||
" filing_type: FilingType = Field(description=\"Type of SEC filing\")\n",
|
||||
" filing_date: str = Field(description=\"Date when filing was submitted to SEC\")\n",
|
||||
" fiscal_year: int = Field(description=\"Fiscal year\")\n",
|
||||
" unit: str = Field(\n",
|
||||
" description=\"Unit of financial figures (thousands, millions, etc.)\"\n",
|
||||
" )\n",
|
||||
" revenue: int = Field(description=\"Total revenue for period\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Set Up Citations and Reasoning\n",
|
||||
"\n",
|
||||
"Optionally, we can set the `ExtractConfig` to extract citations for each field the agent extracts. These cications will cite the specific pages and sections of the file from which a given field was extractedd.\n",
|
||||
"\n",
|
||||
"By setting `use_reasoning` to True, we als ask the agent to do an additional reasoning step, explaining why a given field was extracted."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from llama_cloud.types import ExtractConfig, ExtractMode\n",
|
||||
"\n",
|
||||
"config = ExtractConfig(\n",
|
||||
" use_reasoning=True, cite_sources=True, extraction_mode=ExtractMode.MULTIMODAL\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/usr/local/lib/python3.11/dist-packages/llama_cloud_services/extract/extract.py:127: ExperimentalWarning: `use_reasoning` is an experimental feature. Results will be available in the `extraction_metadata` field for the extraction run.\n",
|
||||
" warnings.warn(\n",
|
||||
"/usr/local/lib/python3.11/dist-packages/llama_cloud_services/extract/extract.py:133: ExperimentalWarning: `cite_sources` is an experimental feature. This may greatly increase the size of the response, and slow down the extraction. Results will be available in the `extraction_metadata` field for the extraction run.\n",
|
||||
" warnings.warn(\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"agent = llama_extract.create_agent(\n",
|
||||
" name=\"filing-parser\", data_schema=FinancialReport, config=config\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Demo Time - Download a PDF and Extract Data with Citations"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"PDF downloaded successfully.\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"import requests\n",
|
||||
"\n",
|
||||
"url = \"https://raw.githubusercontent.com/run-llama/llama_cloud_services/refs/heads/main/examples/extract/data/sec_filings/nvda_10k.pdf\"\n",
|
||||
"\n",
|
||||
"response = requests.get(url)\n",
|
||||
"\n",
|
||||
"if response.status_code == 200:\n",
|
||||
" with open(\"/content/nvda_10k.pdf\", \"wb\") as f:\n",
|
||||
" f.write(response.content)\n",
|
||||
" print(\"PDF downloaded successfully.\")\n",
|
||||
"else:\n",
|
||||
" print(f\"Failed to download. Status code: {response.status_code}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Uploading files: 100%|██████████| 1/1 [00:00<00:00, 1.83it/s]\n",
|
||||
"Creating extraction jobs: 100%|██████████| 1/1 [00:00<00:00, 4.38it/s]\n",
|
||||
"Extracting files: 100%|██████████| 1/1 [02:03<00:00, 123.40s/it]\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"filing_info = agent.extract(\"/content/nvda_10k.pdf\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"{'company_name': 'NVIDIA Corporation',\n",
|
||||
" 'description': \"The filing provides a detailed overview of NVIDIA's business as a full-stack computing infrastructure company, discusses various technologies including digital avatars and autonomous vehicles, outlines numerous risk factors affecting operations such as supply chain issues and geopolitical tensions, and describes employee stock purchase plans and related compliance requirements.\",\n",
|
||||
" 'filing_type': '10 K',\n",
|
||||
" 'filing_date': 'February 26, 2025',\n",
|
||||
" 'fiscal_year': 2025,\n",
|
||||
" 'unit': 'millions',\n",
|
||||
" 'revenue': 130497}"
|
||||
]
|
||||
},
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"filing_info.data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Inspect Citations and Reasoning"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"{'field_metadata': {'company_name': {'reasoning': 'VERBATIM EXTRACTION',\n",
|
||||
" 'citation': [{'page': 1, 'matching_text': 'NVIDIA CORPORATION'},\n",
|
||||
" {'page': 2, 'matching_text': 'NVIDIA Corporation'},\n",
|
||||
" {'page': 3,\n",
|
||||
" 'matching_text': 'All references to \"NVIDIA,\" \"we,\" \"us,\" \"our,\" or the \"Company\" mean NVIDIA Corporation and its subsidiaries.'},\n",
|
||||
" {'page': 35,\n",
|
||||
" 'matching_text': 'Comparison of 5 Year Cumulative Total Return* Among NVIDIA Corporation'},\n",
|
||||
" {'page': 49,\n",
|
||||
" 'matching_text': 'To the Board of Directors and Shareholders of NVIDIA Corporation'},\n",
|
||||
" {'page': 90, 'matching_text': 'NVIDIA Corporation'},\n",
|
||||
" {'page': 119,\n",
|
||||
" 'matching_text': '*\"Company\"* means NVIDIA Corporation, a Delaware corporation.'},\n",
|
||||
" {'page': 126,\n",
|
||||
" 'matching_text': 'Annual Report on Form 10-K of NVIDIA Corporation'}]},\n",
|
||||
" 'filing_type': {'reasoning': \"VERBATIM EXTRACTION from multiple sources confirming the filing type as '10 K'.\",\n",
|
||||
" 'citation': [{'page': 1, 'matching_text': 'FORM 10-K'},\n",
|
||||
" {'page': 2, 'matching_text': 'Item 16. | Form 10-K Summary'},\n",
|
||||
" {'page': 3,\n",
|
||||
" 'matching_text': 'This Annual Report on Form 10-K contains forward-looking statements...'},\n",
|
||||
" {'page': 13, 'matching_text': 'this Annual Report on Form 10-K'},\n",
|
||||
" {'page': 15, 'matching_text': 'this Annual Report on Form 10-K'},\n",
|
||||
" {'page': 32,\n",
|
||||
" 'matching_text': 'Annual Report on Form 10-K, which information is hereby incorporated by reference.'},\n",
|
||||
" {'page': 36, 'matching_text': 'this Annual Report on Form 10-K'},\n",
|
||||
" {'page': 43,\n",
|
||||
" 'matching_text': 'Annual Report on Form 10-K for additional information'},\n",
|
||||
" {'page': 45, 'matching_text': 'Annual Report on Form 10-K'},\n",
|
||||
" {'page': 46, 'matching_text': 'this Annual Report on Form 10-K'},\n",
|
||||
" {'page': 62, 'matching_text': 'Annual Report on Form 10-K'},\n",
|
||||
" {'page': 83,\n",
|
||||
" 'matching_text': 'Restated Certificate of Incorporation | 10-K'},\n",
|
||||
" {'page': 84, 'matching_text': 'Item 16. Form 10-K Summary'},\n",
|
||||
" {'page': 126, 'matching_text': 'which appears in this Form 10-K'},\n",
|
||||
" {'page': 127, 'matching_text': 'Annual Report on Form 10-K'},\n",
|
||||
" {'page': 128, 'matching_text': 'Annual Report on Form 10-K'},\n",
|
||||
" {'page': 129, 'matching_text': \"The Company's Annual Report on Form 10-K\"},\n",
|
||||
" {'page': 130,\n",
|
||||
" 'matching_text': \"The Company's Annual Report on Form 10-K for the year ended January 26, 2025\"}]},\n",
|
||||
" 'fiscal_year': {'reasoning': 'The fiscal year ended January 26, 2025, indicates the fiscal year is 2025. Additionally, multiple references throughout the text confirm the fiscal year 2025 in various contexts.',\n",
|
||||
" 'citation': [{'page': 1,\n",
|
||||
" 'matching_text': 'For the fiscal year ended January 26, 2025'},\n",
|
||||
" {'page': 6,\n",
|
||||
" 'matching_text': 'In fiscal year 2025, we launched the NVIDIA Blackwell architecture'},\n",
|
||||
" {'page': 12, 'matching_text': 'fiscal year 2025'},\n",
|
||||
" {'page': 17,\n",
|
||||
" 'matching_text': 'our gross margins in the second quarter of fiscal year 2025 were negatively impacted'},\n",
|
||||
" {'page': 20,\n",
|
||||
" 'matching_text': 'we generated 53% of our revenue in fiscal year 2025 from sales outside the United States.'},\n",
|
||||
" {'page': 23,\n",
|
||||
" 'matching_text': 'For fiscal year 2025, an indirect customer which primarily purchases our products through system integrators...'},\n",
|
||||
" {'page': 33,\n",
|
||||
" 'matching_text': 'In fiscal year 2025, we repurchased 310 million shares of our common stock for $34.0 billion.'},\n",
|
||||
" {'page': 37,\n",
|
||||
" 'matching_text': 'Our Data Center revenue in China grew in fiscal year 2025.'},\n",
|
||||
" {'page': 44,\n",
|
||||
" 'matching_text': 'Cash provided by operating activities increased in fiscal year 2025 compared to fiscal year 2024'},\n",
|
||||
" {'page': 57,\n",
|
||||
" 'matching_text': 'Fiscal years 2025, 2024 and 2023 were all 52-week years.'},\n",
|
||||
" {'page': 65,\n",
|
||||
" 'matching_text': 'Beginning in the second quarter of fiscal year 2025'},\n",
|
||||
" {'page': 69, 'matching_text': 'In the fourth quarter of fiscal year 2025'},\n",
|
||||
" {'page': 78,\n",
|
||||
" 'matching_text': 'Depreciation and amortization expense attributable to our Compute and Networking segment for fiscal years 2025'},\n",
|
||||
" {'page': 129, 'matching_text': 'for the year ended January 26, 2025'}]},\n",
|
||||
" 'description': {'reasoning': 'The extracted data combines multiple descriptions from the source text, ensuring no duplication while maintaining the order and context of the information. Each section of the filing is summarized to reflect the key points without losing the essence of the original text.',\n",
|
||||
" 'citation': [{'page': 4,\n",
|
||||
" 'matching_text': 'NVIDIA is now a full-stack computing infrastructure company with data-center-scale offerings that are reshaping industry.'},\n",
|
||||
" {'page': 8,\n",
|
||||
" 'matching_text': 'a suite of technologies that help developers bring digital avatars to life with generative Al...autonomous vehicles, or AV, and electric vehicles, or EV, is revolutionizing the transportation industry...Our worldwide sales and marketing strategy is key to achieving our objective of providing markets with our high-performance and efficient computing platforms and software.'},\n",
|
||||
" {'page': 14, 'matching_text': 'Risk Factors Summary'},\n",
|
||||
" {'page': 16,\n",
|
||||
" 'matching_text': 'Risks Related to Demand, Supply, and Manufacturing\\n\\nLong manufacturing lead times and uncertain supply and component availability...'},\n",
|
||||
" {'page': 18,\n",
|
||||
" 'matching_text': 'cryptocurrency mining, on demand for our products. Volatility in the cryptocurrency market, including new compute technologies...'},\n",
|
||||
" {'page': 21,\n",
|
||||
" 'matching_text': 'supply-chain attacks or other business disruptions. We cannot guarantee that third parties and infrastructure in our supply chain...'},\n",
|
||||
" {'page': 22,\n",
|
||||
" 'matching_text': 'We are monitoring the impact of the geopolitical conflict in and around Israel on our operations... Climate change may have a long-term impact on our business.'},\n",
|
||||
" {'page': 25,\n",
|
||||
" 'matching_text': 'We are subject to complex laws, rules, regulations, and political and other actions, including restrictions on the export of our products, which may adversely impact our business.'},\n",
|
||||
" {'page': 28,\n",
|
||||
" 'matching_text': 'Our competitive position has been harmed by the existing export controls, and our competitive position and future results may be further harmed'},\n",
|
||||
" {'page': 29,\n",
|
||||
" 'matching_text': 'restrictions imposed by the Chinese government on the duration of gaming activities and access to games may adversely affect our Gaming revenue'},\n",
|
||||
" {'page': 29,\n",
|
||||
" 'matching_text': 'our business depends on our ability to receive consistent and reliable supply from our overseas partners, especially in Taiwan and South Korea'},\n",
|
||||
" {'page': 29,\n",
|
||||
" 'matching_text': 'Increased scrutiny from shareholders, regulators and others regarding our corporate sustainability practices could result in additional costs'},\n",
|
||||
" {'page': 29,\n",
|
||||
" 'matching_text': 'Concerns relating to the responsible use of new and evolving technologies, such as Al, in our products and services may result in reputational or financial harm'},\n",
|
||||
" {'page': 31,\n",
|
||||
" 'matching_text': 'Data protection laws around the world are quickly changing and may be interpreted and applied in an increasingly stringent fashion...'}]},\n",
|
||||
" 'filing_date': {'reasoning': 'The filing date is consistently mentioned as February 26, 2025 across multiple entries, making it the most reliable date for the filing.',\n",
|
||||
" 'citation': [{'page': 51, 'matching_text': 'February 26, 2025'},\n",
|
||||
" {'page': 86, 'matching_text': 'on February 26, 2025.'},\n",
|
||||
" {'page': 87, 'matching_text': 'February 26, 2025'},\n",
|
||||
" {'page': 126, 'matching_text': 'our report dated February 26, 2025'},\n",
|
||||
" {'page': 127, 'matching_text': 'Date: February 26, 2025'},\n",
|
||||
" {'page': 128, 'matching_text': 'Date: February 26, 2025'},\n",
|
||||
" {'page': 129, 'matching_text': 'Date: February 26, 2025'},\n",
|
||||
" {'page': 130, 'matching_text': 'Date: February 26, 2025'}]},\n",
|
||||
" 'unit': {'reasoning': \"The unit of financial figures is explicitly mentioned multiple times in the text as 'millions', including in table headers and notes. This is confirmed by various citations from pages 38, 42, 43, 52, 53, 54, 56, 65, 71, 72, 73, 75, 77, 79, 80, and 82.\",\n",
|
||||
" 'citation': [{'page': 38,\n",
|
||||
" 'matching_text': '($ in millions, except per share data)'},\n",
|
||||
" {'page': 42, 'matching_text': '($ in millions)'},\n",
|
||||
" {'page': 43, 'matching_text': '($ in millions)'},\n",
|
||||
" {'page': 52, 'matching_text': '(In millions, except per share data)'},\n",
|
||||
" {'page': 53,\n",
|
||||
" 'matching_text': 'Consolidated Statements of Comprehensive Income (In millions)'},\n",
|
||||
" {'page': 54,\n",
|
||||
" 'matching_text': 'Consolidated Balance Sheets (In millions, except par value)'},\n",
|
||||
" {'page': 55, 'matching_text': '(In millions, except per share data)'},\n",
|
||||
" {'page': 56,\n",
|
||||
" 'matching_text': 'Consolidated Statements of Cash Flows (In millions)'},\n",
|
||||
" {'page': 65,\n",
|
||||
" 'matching_text': 'Year Ended<br/>Jan 26, 2025<br/>(In millions, except per share data)'},\n",
|
||||
" {'page': 71, 'matching_text': '(In millions) | (In millions)'},\n",
|
||||
" {'page': 72, 'matching_text': '(In millions)'}]},\n",
|
||||
" 'revenue': {'reasoning': 'The total revenue for fiscal year 2025 is extracted from multiple sources within the text, all confirming the same figure of $130,497 million. The revenue recognized for fiscal year 2025 is also noted as $4,607 million, which is a separate figure. However, the primary focus is on the total revenue figure, which is consistently cited.',\n",
|
||||
" 'citation': [{'page': 38,\n",
|
||||
" 'matching_text': 'Revenue for fiscal year 2025 was $130.5 billion'},\n",
|
||||
" {'page': 41,\n",
|
||||
" 'matching_text': 'Total | $ 130,497 | $ | 60,922'},\n",
|
||||
" {'page': 52, 'matching_text': 'Revenue | $ 130,497'},\n",
|
||||
" {'page': 78,\n",
|
||||
" 'matching_text': 'Revenue | $ 116,193 | $ 14,304 | $ - | $ 130,497'},\n",
|
||||
" {'page': 79, 'matching_text': 'Total revenue | $ 130,497'},\n",
|
||||
" {'page': 80, 'matching_text': 'Total revenue | $ 130,497'}]}},\n",
|
||||
" 'usage': {'num_pages_extracted': 130,\n",
|
||||
" 'num_document_tokens': 105932,\n",
|
||||
" 'num_output_tokens': 31306}}"
|
||||
]
|
||||
},
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"filing_info.extraction_metadata"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## What's Next?\n",
|
||||
"\n",
|
||||
"In this example, we built an Extraction Agent that is capable of citing it's sources from the document it's extracting data from, and reasoning about its reponse. To further customize and improve on the results, you can also try to customize the `system_prompt` in the `ExtractConfig`.\n",
|
||||
"\n",
|
||||
"#### Learn More\n",
|
||||
"\n",
|
||||
"- [LlamaExtract Documentation](https://docs.cloud.llamaindex.ai/llamaextract/getting_started)\n",
|
||||
"- [Example Notebooks](https://github.com/run-llama/llama_cloud_services/tree/main/examples/extract)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,19 @@
|
||||
from .schema import (
|
||||
TypedAgentData,
|
||||
ExtractedData,
|
||||
TypedAgentDataItems,
|
||||
StatusType,
|
||||
ExtractedT,
|
||||
AgentDataT,
|
||||
)
|
||||
from .client import AsyncAgentDataClient
|
||||
|
||||
__all__ = [
|
||||
"TypedAgentData",
|
||||
"AsyncAgentDataClient",
|
||||
"ExtractedData",
|
||||
"TypedAgentDataItems",
|
||||
"StatusType",
|
||||
"ExtractedT",
|
||||
"AgentDataT",
|
||||
]
|
||||
@@ -0,0 +1,267 @@
|
||||
import os
|
||||
from typing import Dict, Generic, List, Optional, Type
|
||||
|
||||
from llama_cloud import FilterOperation
|
||||
from llama_cloud.client import AsyncLlamaCloud
|
||||
from tenacity import (
|
||||
WrappedFn,
|
||||
retry,
|
||||
stop_after_attempt,
|
||||
wait_exponential,
|
||||
retry_if_exception_type,
|
||||
)
|
||||
import httpx
|
||||
|
||||
from .schema import (
|
||||
AgentDataT,
|
||||
TypedAgentData,
|
||||
TypedAgentDataItems,
|
||||
TypedAggregateGroup,
|
||||
TypedAggregateGroupItems,
|
||||
)
|
||||
|
||||
|
||||
def agent_data_retry(func: WrappedFn) -> WrappedFn:
|
||||
"""
|
||||
Decorator that adds automatic retry logic to agent data API calls.
|
||||
|
||||
Applies exponential backoff retry strategy for common network-related exceptions:
|
||||
- Up to 3 retry attempts
|
||||
- Exponential wait time between 0.5s and 10s
|
||||
- Retries on timeout, connection, and HTTP status errors
|
||||
|
||||
This ensures resilient API communication in distributed environments where
|
||||
temporary network issues or service unavailability may occur.
|
||||
"""
|
||||
return retry(
|
||||
stop=stop_after_attempt(3),
|
||||
wait=wait_exponential(min=0.5, max=10),
|
||||
retry=retry_if_exception_type(
|
||||
(httpx.TimeoutException, httpx.ConnectError, httpx.HTTPStatusError)
|
||||
),
|
||||
)(func)
|
||||
|
||||
|
||||
def get_default_agent_id() -> Optional[str]:
|
||||
"""
|
||||
Retrieve the default agent ID from environment variables.
|
||||
|
||||
Returns:
|
||||
The value of LLAMA_DEPLOY_DEPLOYMENT_NAME environment variable,
|
||||
or None if not set
|
||||
|
||||
Note:
|
||||
This provides a convenient way to configure agent ID globally
|
||||
via environment variables instead of passing it explicitly
|
||||
to each client instance.
|
||||
"""
|
||||
return os.getenv("LLAMA_DEPLOY_DEPLOYMENT_NAME")
|
||||
|
||||
|
||||
class AsyncAgentDataClient(Generic[AgentDataT]):
|
||||
"""
|
||||
Async client for managing agent-generated structured data with type safety.
|
||||
|
||||
This client provides a high-level interface for CRUD operations, searching, and
|
||||
aggregation of structured data created by agents. It enforces type safety by
|
||||
validating all data against a specified Pydantic model type.
|
||||
|
||||
The client is generic over AgentDataT, which must be a Pydantic BaseModel that
|
||||
defines the structure of your agent's data output.
|
||||
|
||||
Example:
|
||||
```python
|
||||
from pydantic import BaseModel
|
||||
from llama_cloud.client import AsyncLlamaCloud
|
||||
from llama_cloud_services.beta.agent_data import AsyncAgentDataClient
|
||||
|
||||
class ExtractedPerson(BaseModel):
|
||||
name: str
|
||||
age: int
|
||||
email: str
|
||||
|
||||
# Initialize client
|
||||
llama_client = AsyncLlamaCloud(token="your-api-key")
|
||||
agent_client = AsyncAgentDataClient(
|
||||
client=llama_client,
|
||||
type=ExtractedPerson,
|
||||
collection_name="extracted_people",
|
||||
agent_url_id="person-extraction-agent"
|
||||
)
|
||||
|
||||
# Create data
|
||||
person = ExtractedPerson(name="John Doe", age=30, email="john@example.com")
|
||||
result = await agent_client.create_agent_data(person)
|
||||
|
||||
# Search data
|
||||
results = await agent_client.search_agent_data(
|
||||
filter={"age": FilterOperation(gt=25)},
|
||||
order_by="data.name",
|
||||
page_size=20
|
||||
)
|
||||
```
|
||||
|
||||
Type Parameters:
|
||||
AgentDataT: Pydantic BaseModel type that defines the structure of agent data
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: AsyncLlamaCloud,
|
||||
type: Type[AgentDataT],
|
||||
collection_name: str = "default",
|
||||
agent_url_id: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Initialize the AsyncAgentDataClient.
|
||||
|
||||
Args:
|
||||
client: AsyncLlamaCloud client instance for API communication
|
||||
type: Pydantic BaseModel class that defines the data structure.
|
||||
All agent data will be validated against this type.
|
||||
collection_name: Named collection within the agent for organizing data.
|
||||
Defaults to "default". Collections allow logical separation of
|
||||
different data types or workflows within the same agent.
|
||||
agent_url_id: Unique identifier for the agent. This normally appears in the
|
||||
url of an agent within the llama cloud platform. If not provided,
|
||||
will attempt to use the LLAMA_DEPLOY_DEPLOYMENT_NAME environment
|
||||
variable. Data can only be added to an already existing agent in the
|
||||
platform.
|
||||
|
||||
Raises:
|
||||
ValueError: If agent_url_id is not provided and the
|
||||
LLAMA_DEPLOY_DEPLOYMENT_NAME environment variable is not set
|
||||
|
||||
Note:
|
||||
The client automatically applies retry logic to all API calls with
|
||||
exponential backoff for timeout, connection, and HTTP status errors.
|
||||
"""
|
||||
|
||||
self.agent_url_id = agent_url_id or get_default_agent_id()
|
||||
if not self.agent_url_id:
|
||||
raise ValueError(
|
||||
"Agent ID is required, or set the LLAMA_DEPLOY_DEPLOYMENT_NAME environment variable"
|
||||
)
|
||||
|
||||
self.collection_name = collection_name
|
||||
self.client = client
|
||||
self.type = type
|
||||
|
||||
@agent_data_retry
|
||||
async def get_agent_data(self, item_id: str) -> TypedAgentData[AgentDataT]:
|
||||
raw_data = await self.client.beta.get_agent_data(
|
||||
item_id=item_id,
|
||||
)
|
||||
return TypedAgentData.from_raw(raw_data, validator=self.type)
|
||||
|
||||
@agent_data_retry
|
||||
async def create_agent_data(self, data: AgentDataT) -> TypedAgentData[AgentDataT]:
|
||||
raw_data = await self.client.beta.create_agent_data(
|
||||
agent_slug=self.agent_url_id,
|
||||
collection=self.collection_name,
|
||||
data=data.model_dump(),
|
||||
)
|
||||
return TypedAgentData.from_raw(raw_data, validator=self.type)
|
||||
|
||||
@agent_data_retry
|
||||
async def update_agent_data(
|
||||
self, item_id: str, data: AgentDataT
|
||||
) -> TypedAgentData[AgentDataT]:
|
||||
raw_data = await self.client.beta.update_agent_data(
|
||||
item_id=item_id,
|
||||
data=data.model_dump(),
|
||||
)
|
||||
return TypedAgentData.from_raw(raw_data, validator=self.type)
|
||||
|
||||
@agent_data_retry
|
||||
async def delete_agent_data(self, item_id: str) -> None:
|
||||
await self.client.beta.delete_agent_data(item_id=item_id)
|
||||
|
||||
@agent_data_retry
|
||||
async def search_agent_data(
|
||||
self,
|
||||
filter: Optional[Dict[str, Optional[FilterOperation]]] = None,
|
||||
order_by: Optional[str] = None,
|
||||
offset: Optional[int] = None,
|
||||
page_size: Optional[int] = None,
|
||||
include_total: bool = False,
|
||||
) -> TypedAgentDataItems[AgentDataT]:
|
||||
"""
|
||||
Search agent data with filtering, sorting, and pagination.
|
||||
Args:
|
||||
filter: Filter conditions to apply to the search. Dict mapping field names to FilterOperation objects. Filters only by data fields
|
||||
Examples:
|
||||
- {"age": FilterOperation(gt=18)} - age greater than 18
|
||||
- {"status": FilterOperation(eq="active")} - status equals "active"
|
||||
- {"tags": FilterOperation(includes=["python", "ml"])} - tags include "python" or "ml"
|
||||
- {"created_at": FilterOperation(gte="2024-01-01")} - created after date
|
||||
- {"score": FilterOperation(lt=100, gte=50)} - score between 50 and 100
|
||||
order_by: Comma delimited list of fields to sort results by. Can order by standard agent fields like created_at, or by data fields. Data fields must be prefixed with "data.". If ordering desceding, use a " desc" suffix.
|
||||
Examples:
|
||||
- "data.name desc, created_at" - sort by name in descending order, and then by creation date
|
||||
page_size: Maximum number of items to return per page. Defaults to 10.
|
||||
offset: Number of items to skip from the beginning. Defaults to 0.
|
||||
include_total: Whether to include the total count in the response. Defaults to False to improve performance. It's recommended to only request on the first page.
|
||||
"""
|
||||
raw = await self.client.beta.search_agent_data_api_v_1_beta_agent_data_search_post(
|
||||
agent_slug=self.agent_url_id,
|
||||
collection=self.collection_name,
|
||||
filter=filter,
|
||||
order_by=order_by,
|
||||
offset=offset,
|
||||
page_size=page_size,
|
||||
include_total=include_total,
|
||||
)
|
||||
return TypedAgentDataItems(
|
||||
items=[
|
||||
TypedAgentData.from_raw(item, validator=self.type) for item in raw.items
|
||||
],
|
||||
has_more=raw.next_page_token is not None,
|
||||
total=raw.total_size,
|
||||
)
|
||||
|
||||
@agent_data_retry
|
||||
async def aggregate_agent_data(
|
||||
self,
|
||||
filter: Optional[Dict[str, Optional[FilterOperation]]] = None,
|
||||
group_by: Optional[List[str]] = None,
|
||||
count: Optional[bool] = None,
|
||||
first: Optional[bool] = None,
|
||||
order_by: Optional[str] = None,
|
||||
offset: Optional[int] = None,
|
||||
page_size: Optional[int] = None,
|
||||
) -> TypedAggregateGroupItems[AgentDataT]:
|
||||
"""
|
||||
Aggregate agent data into groups according to the group_by fields.
|
||||
Args:
|
||||
filter: Filter conditions to apply to the search. Dict mapping field names to FilterOperation objects. Filters only by data fields
|
||||
See search_agent_data for more details on filtering.
|
||||
group_by: List of fields to group by. Groups strictly by equality. Can only group by data fields.
|
||||
Examples:
|
||||
- ["name"] - group by name
|
||||
- ["name", "age"] - group by name and age
|
||||
count: Whether to include the count of items in each group.
|
||||
first: Whether to include the first item in each group.
|
||||
order_by: Comma delimited list of fields to sort results by. See search_agent_data for more details on ordering.
|
||||
offset: Number of groups to skip from the beginning. Defaults to 0.
|
||||
page_size: Maximum number of groups to return per page.
|
||||
"""
|
||||
raw = await self.client.beta.aggregate_agent_data_api_v_1_beta_agent_data_aggregate_post(
|
||||
agent_slug=self.agent_url_id,
|
||||
collection=self.collection_name,
|
||||
page_size=page_size,
|
||||
filter=filter,
|
||||
order_by=order_by,
|
||||
group_by=group_by,
|
||||
count=count,
|
||||
first=first,
|
||||
offset=offset,
|
||||
)
|
||||
return TypedAggregateGroupItems(
|
||||
items=[
|
||||
TypedAggregateGroup.from_raw(item, validator=self.type)
|
||||
for item in raw.items
|
||||
],
|
||||
has_more=raw.next_page_token is not None,
|
||||
total=raw.total_size,
|
||||
)
|
||||
@@ -0,0 +1,357 @@
|
||||
"""
|
||||
Agent Data API Schema Definitions
|
||||
|
||||
This module provides typed wrappers around the raw LlamaCloud agent data API,
|
||||
enabling type-safe interactions with agent-generated structured data.
|
||||
|
||||
The agent data API serves as a persistent storage system for structured data
|
||||
produced by LlamaCloud agents (particularly extraction agents). It provides
|
||||
CRUD operations, search capabilities, filtering, and aggregation functionality
|
||||
for managing agent-generated data at scale.
|
||||
|
||||
Key Concepts:
|
||||
- Agent Slug: Unique identifier for an agent instance
|
||||
- Collection: Named grouping of data within an agent (defaults to "default"). Data within a collection should be of the same type.
|
||||
- Agent Data: Individual structured data records with metadata and timestamps
|
||||
|
||||
Example Usage:
|
||||
```python
|
||||
from pydantic import BaseModel
|
||||
|
||||
class Person(BaseModel):
|
||||
name: str
|
||||
age: int
|
||||
|
||||
client = AsyncAgentDataClient(
|
||||
client=async_llama_cloud,
|
||||
type=Person,
|
||||
collection="people",
|
||||
agent_url_id="my-extraction-agent-xyz"
|
||||
)
|
||||
|
||||
# Create typed data
|
||||
person = Person(name="John", age=30)
|
||||
result = await client.create_agent_data(person)
|
||||
print(result.data.name) # Type-safe access
|
||||
```
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from llama_cloud.types.agent_data import AgentData
|
||||
from llama_cloud.types.aggregate_group import AggregateGroup
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import (
|
||||
Generic,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Dict,
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
Any,
|
||||
)
|
||||
|
||||
|
||||
# Type variable for user-defined data models
|
||||
AgentDataT = TypeVar("AgentDataT", bound=BaseModel)
|
||||
|
||||
# Type variable for extracted data (can be dict or Pydantic model)
|
||||
ExtractedT = TypeVar("ExtractedT", bound=Union[BaseModel, dict])
|
||||
|
||||
# Status types for extracted data workflow
|
||||
StatusType = Union[Literal["error", "accepted", "rejected", "in_review"], str]
|
||||
|
||||
|
||||
class TypedAgentData(BaseModel, Generic[AgentDataT]):
|
||||
"""
|
||||
Type-safe wrapper for agent data records.
|
||||
|
||||
This class represents a single data record stored in the agent data API,
|
||||
combining the structured data payload with metadata about when and where
|
||||
it was created.
|
||||
|
||||
Attributes:
|
||||
id: Unique identifier for this data record
|
||||
agent_url_id: Identifier of the agent that created this data
|
||||
collection: Named collection within the agent (used for organization)
|
||||
data: The actual structured data payload (typed as AgentDataT)
|
||||
created_at: Timestamp when the record was first created
|
||||
updated_at: Timestamp when the record was last modified
|
||||
|
||||
Example:
|
||||
```python
|
||||
# Access typed data
|
||||
person_data: TypedAgentData[Person] = await client.get_agent_data(id)
|
||||
print(person_data.data.name) # Type-safe access to Person fields
|
||||
print(person_data.created_at) # Access metadata
|
||||
```
|
||||
"""
|
||||
|
||||
id: Optional[str] = Field(description="Unique identifier for this data record")
|
||||
agent_url_id: str = Field(
|
||||
description="Identifier of the agent that created this data"
|
||||
)
|
||||
collection: Optional[str] = Field(
|
||||
description="Named collection within the agent for data organization"
|
||||
)
|
||||
data: AgentDataT = Field(description="The structured data payload")
|
||||
created_at: Optional[datetime] = Field(description="When this record was created")
|
||||
updated_at: Optional[datetime] = Field(
|
||||
description="When this record was last modified"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_raw(
|
||||
cls, raw_data: AgentData, validator: Type[AgentDataT]
|
||||
) -> "TypedAgentData[AgentDataT]":
|
||||
"""
|
||||
Convert raw API response to typed agent data.
|
||||
|
||||
Args:
|
||||
raw_data: Raw agent data from the API
|
||||
validator: Pydantic model class to validate the data field
|
||||
|
||||
Returns:
|
||||
TypedAgentData instance with validated data
|
||||
"""
|
||||
data: AgentDataT = validator.model_validate(raw_data.data)
|
||||
|
||||
return cls(
|
||||
id=raw_data.id,
|
||||
agent_url_id=raw_data.agent_slug,
|
||||
collection=raw_data.collection,
|
||||
data=data,
|
||||
created_at=raw_data.created_at,
|
||||
updated_at=raw_data.updated_at,
|
||||
)
|
||||
|
||||
|
||||
class TypedAgentDataItems(BaseModel, Generic[AgentDataT]):
|
||||
"""
|
||||
Paginated collection of agent data records.
|
||||
|
||||
This class represents a page of search results from the agent data API,
|
||||
providing both the data records and pagination metadata.
|
||||
|
||||
Attributes:
|
||||
items: List of agent data records in this page
|
||||
total: Total number of records matching the query (only present if requested)
|
||||
has_more: Whether there are more records available beyond this page
|
||||
|
||||
Example:
|
||||
```python
|
||||
# Search with pagination
|
||||
results = await client.search_agent_data(
|
||||
page_size=10,
|
||||
include_total=True
|
||||
)
|
||||
|
||||
for item in results.items:
|
||||
print(item.data.name)
|
||||
|
||||
if results.has_more:
|
||||
# Load next page
|
||||
next_page = await client.search_agent_data(
|
||||
page_size=10,
|
||||
offset=10
|
||||
)
|
||||
```
|
||||
"""
|
||||
|
||||
items: List[TypedAgentData[AgentDataT]] = Field(
|
||||
description="List of agent data records in this page"
|
||||
)
|
||||
total: Optional[int] = Field(
|
||||
description="Total number of records matching the query (only present if requested)"
|
||||
)
|
||||
has_more: bool = Field(
|
||||
description="Whether there are more records available beyond this page"
|
||||
)
|
||||
|
||||
|
||||
class ExtractedData(BaseModel, Generic[ExtractedT]):
|
||||
"""
|
||||
Wrapper for extracted data with workflow status tracking.
|
||||
|
||||
This class is designed for extraction workflows where data goes through
|
||||
review and approval stages. It maintains both the original extracted data
|
||||
and the current state after any modifications.
|
||||
|
||||
Attributes:
|
||||
original_data: The data as originally extracted from the source
|
||||
data: The current state of the data (may differ from original after edits)
|
||||
status: Current workflow status (in_review, accepted, rejected, error)
|
||||
confidence: Confidence scores for individual fields (if available)
|
||||
|
||||
Status Workflow:
|
||||
- "in_review": Initial state, awaiting human review
|
||||
- "accepted": Data approved and ready for use
|
||||
- "rejected": Data rejected, needs re-extraction or manual fix
|
||||
- "error": Processing error occurred
|
||||
|
||||
Example:
|
||||
```python
|
||||
# Create extracted data for review
|
||||
extracted = ExtractedData.create(
|
||||
extracted_data=person_data,
|
||||
status="in_review",
|
||||
confidence={"name": 0.95, "age": 0.87}
|
||||
)
|
||||
|
||||
# Later, after review
|
||||
if extracted.status == "accepted":
|
||||
# Use the data
|
||||
process_person(extracted.data)
|
||||
```
|
||||
"""
|
||||
|
||||
original_data: ExtractedT = Field(
|
||||
description="The original data that was extracted from the document"
|
||||
)
|
||||
data: ExtractedT = Field(
|
||||
description="The latest state of the data. Will differ if data has been updated"
|
||||
)
|
||||
status: Union[Literal["error", "accepted", "rejected", "in_review"], str] = Field(
|
||||
description="The status of the extracted data"
|
||||
)
|
||||
confidence: Dict[str, Union[float, Dict]] = Field(
|
||||
default_factory=dict,
|
||||
description="Confidence scores, if any, for each primitive field in the original_data data",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
extracted_data: ExtractedT,
|
||||
status: StatusType = "in_review",
|
||||
confidence: Optional[Dict[str, Union[float, Dict]]] = None,
|
||||
) -> "ExtractedData[ExtractedT]":
|
||||
"""
|
||||
Create a new ExtractedData instance with sensible defaults.
|
||||
|
||||
Args:
|
||||
extracted_data: The extracted data payload
|
||||
status: Initial workflow status
|
||||
confidence: Optional confidence scores for fields
|
||||
|
||||
Returns:
|
||||
New ExtractedData instance ready for storage
|
||||
"""
|
||||
return cls(
|
||||
original_data=extracted_data,
|
||||
data=extracted_data,
|
||||
status=status,
|
||||
confidence=confidence or {},
|
||||
)
|
||||
|
||||
|
||||
class TypedAggregateGroup(BaseModel, Generic[AgentDataT]):
|
||||
"""
|
||||
Represents a group of agent data records aggregated by common field values.
|
||||
|
||||
This class is used for grouping and analyzing agent data based on shared
|
||||
characteristics. It's particularly useful for generating summaries and
|
||||
statistics across large datasets.
|
||||
|
||||
Attributes:
|
||||
group_key: The field values that define this group
|
||||
count: Number of records in this group (if count aggregation was requested)
|
||||
first_item: Representative data record from this group (if requested)
|
||||
|
||||
Example:
|
||||
```python
|
||||
# Group by age range
|
||||
groups = await client.aggregate_agent_data(
|
||||
group_by=["age_range"],
|
||||
count=True,
|
||||
first=True
|
||||
)
|
||||
|
||||
for group in groups.items:
|
||||
print(f"Age range {group.group_key['age_range']}: {group.count} people")
|
||||
if group.first_item:
|
||||
print(f"Example: {group.first_item.name}")
|
||||
```
|
||||
"""
|
||||
|
||||
group_key: Dict[str, Any] = Field(
|
||||
description="The field values that define this group"
|
||||
)
|
||||
count: Optional[int] = Field(
|
||||
description="Number of records in this group (if count aggregation was requested)"
|
||||
)
|
||||
first_item: Optional[AgentDataT] = Field(
|
||||
description="Representative data record from this group (if requested)"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_raw(
|
||||
cls, raw_data: AggregateGroup, validator: Type[AgentDataT]
|
||||
) -> "TypedAggregateGroup[AgentDataT]":
|
||||
"""
|
||||
Convert raw API response to typed aggregate group.
|
||||
|
||||
Args:
|
||||
raw_data: Raw aggregate group from the API
|
||||
validator: Pydantic model class to validate the first_item field
|
||||
|
||||
Returns:
|
||||
TypedAggregateGroup instance with validated first_item
|
||||
"""
|
||||
first_item: Optional[AgentDataT] = raw_data.first_item
|
||||
if first_item is not None:
|
||||
first_item = validator.model_validate(first_item)
|
||||
|
||||
return cls(
|
||||
group_key=raw_data.group_key,
|
||||
count=raw_data.count,
|
||||
first_item=first_item,
|
||||
)
|
||||
|
||||
|
||||
class TypedAggregateGroupItems(BaseModel, Generic[AgentDataT]):
|
||||
"""
|
||||
Paginated collection of aggregate groups.
|
||||
|
||||
This class represents a page of aggregation results from the agent data API,
|
||||
providing both the grouped data and pagination metadata.
|
||||
|
||||
Attributes:
|
||||
items: List of aggregate groups in this page
|
||||
total: Total number of groups matching the query (only present if requested)
|
||||
has_more: Whether there are more groups available beyond this page
|
||||
|
||||
Example:
|
||||
```python
|
||||
# Get first page of groups
|
||||
results = await client.aggregate_agent_data(
|
||||
group_by=["department"],
|
||||
count=True,
|
||||
page_size=20
|
||||
)
|
||||
|
||||
for group in results.items:
|
||||
dept = group.group_key["department"]
|
||||
print(f"{dept}: {group.count} employees")
|
||||
|
||||
# Load more if needed
|
||||
if results.has_more:
|
||||
next_page = await client.aggregate_agent_data(
|
||||
group_by=["department"],
|
||||
count=True,
|
||||
page_size=20,
|
||||
offset=20
|
||||
)
|
||||
```
|
||||
"""
|
||||
|
||||
items: List[TypedAggregateGroup[AgentDataT]] = Field(
|
||||
description="List of aggregate groups in this page"
|
||||
)
|
||||
total: Optional[int] = Field(
|
||||
description="Total number of groups matching the query (only present if requested)"
|
||||
)
|
||||
has_more: bool = Field(
|
||||
description="Whether there are more groups available beyond this page"
|
||||
)
|
||||
@@ -1,7 +1,17 @@
|
||||
from llama_cloud_services.extract.extract import (
|
||||
LlamaExtract,
|
||||
ExtractConfig,
|
||||
ExtractionAgent,
|
||||
SourceText,
|
||||
ExtractTarget,
|
||||
ExtractMode,
|
||||
)
|
||||
|
||||
__all__ = ["LlamaExtract", "ExtractionAgent", "SourceText"]
|
||||
__all__ = [
|
||||
"LlamaExtract",
|
||||
"ExtractionAgent",
|
||||
"SourceText",
|
||||
"ExtractConfig",
|
||||
"ExtractTarget",
|
||||
"ExtractMode",
|
||||
]
|
||||
|
||||
@@ -8,25 +8,32 @@ import secrets
|
||||
import warnings
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from tenacity import (
|
||||
retry_if_exception,
|
||||
stop_after_attempt,
|
||||
wait_exponential_jitter,
|
||||
AsyncRetrying,
|
||||
)
|
||||
from llama_cloud import (
|
||||
ExtractAgent as CloudExtractAgent,
|
||||
ExtractAgentCreate,
|
||||
ExtractConfig,
|
||||
ExtractJob,
|
||||
ExtractJobCreate,
|
||||
ExtractRun,
|
||||
ExtractSchemaValidateRequest,
|
||||
ExtractAgentUpdate,
|
||||
File,
|
||||
ExtractMode,
|
||||
StatusEnum,
|
||||
Project,
|
||||
ExtractTarget,
|
||||
LlamaExtractSettings,
|
||||
PaginatedExtractRunsResponse,
|
||||
)
|
||||
from llama_cloud.client import AsyncLlamaCloud
|
||||
from llama_cloud_services.extract.utils import JSONObjectType, augment_async_errors
|
||||
from llama_cloud.core.api_error import ApiError
|
||||
from llama_cloud_services.extract.utils import (
|
||||
JSONObjectType,
|
||||
augment_async_errors,
|
||||
ExperimentalWarning,
|
||||
)
|
||||
from llama_index.core.schema import BaseComponent
|
||||
from llama_index.core.async_utils import run_jobs
|
||||
from llama_index.core.bridge.pydantic import Field, PrivateAttr
|
||||
@@ -44,6 +51,17 @@ DEFAULT_EXTRACT_CONFIG = ExtractConfig(
|
||||
)
|
||||
|
||||
|
||||
def _is_retryable_error(exception: BaseException) -> bool:
|
||||
"""Check if an exception is retryable."""
|
||||
if isinstance(exception, ApiError):
|
||||
return exception.status_code in (502, 503, 504, 425, 408)
|
||||
elif isinstance(
|
||||
exception, (httpx.HTTPStatusError, httpx.RequestError, httpx.TimeoutException)
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class SourceText:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -118,6 +136,22 @@ def run_in_thread(
|
||||
return thread_pool.submit(run_coro).result()
|
||||
|
||||
|
||||
def _extraction_config_warning(config: ExtractConfig) -> None:
|
||||
if config.use_reasoning:
|
||||
warnings.warn(
|
||||
"`use_reasoning` is an experimental feature. Results will be available in "
|
||||
"the `extraction_metadata` field for the extraction run.",
|
||||
ExperimentalWarning,
|
||||
)
|
||||
if config.cite_sources:
|
||||
warnings.warn(
|
||||
"`cite_sources` is an experimental feature. This may greatly increase the "
|
||||
"size of the response, and slow down the extraction. Results will be "
|
||||
"available in the `extraction_metadata` field for the extraction run.",
|
||||
ExperimentalWarning,
|
||||
)
|
||||
|
||||
|
||||
class ExtractionAgent:
|
||||
"""Class representing a single extraction agent with methods for extraction operations."""
|
||||
|
||||
@@ -178,7 +212,7 @@ class ExtractionAgent:
|
||||
)
|
||||
validated_schema = self._run_in_thread(
|
||||
self._client.llama_extract.validate_extraction_schema(
|
||||
request=ExtractSchemaValidateRequest(data_schema=processed_schema)
|
||||
data_schema=processed_schema
|
||||
)
|
||||
)
|
||||
self._data_schema = validated_schema.data_schema
|
||||
@@ -189,6 +223,7 @@ class ExtractionAgent:
|
||||
|
||||
@config.setter
|
||||
def config(self, config: ExtractConfig) -> None:
|
||||
_extraction_config_warning(config)
|
||||
self._config = config
|
||||
|
||||
def _run_in_thread(self, coro: Coroutine[Any, Any, T]) -> T:
|
||||
@@ -211,9 +246,8 @@ class ExtractionAgent:
|
||||
ValueError: If filename is not provided for bytes input or for file-like objects
|
||||
without a name attribute.
|
||||
"""
|
||||
file_contents: Optional[Union[BufferedIOBase, BytesIO]] = None
|
||||
try:
|
||||
file_contents: Union[BufferedIOBase, BytesIO]
|
||||
|
||||
if file_input.text_content is not None:
|
||||
# Handle direct text content
|
||||
file_contents = BytesIO(file_input.text_content.encode("utf-8"))
|
||||
@@ -240,7 +274,7 @@ class ExtractionAgent:
|
||||
project_id=self._project_id, upload_file=file_contents
|
||||
)
|
||||
finally:
|
||||
if isinstance(file_contents, BufferedReader):
|
||||
if file_contents is not None and isinstance(file_contents, BufferedReader):
|
||||
file_contents.close()
|
||||
|
||||
async def _upload_file(self, file_input: FileInput) -> File:
|
||||
@@ -268,35 +302,60 @@ class ExtractionAgent:
|
||||
|
||||
return await self.upload_file(source_text)
|
||||
|
||||
async def _get_job_with_retry(self, job_id: str) -> ExtractJob:
|
||||
"""Get job with retry logic for transient errors."""
|
||||
async for attempt in AsyncRetrying(
|
||||
retry=retry_if_exception(_is_retryable_error),
|
||||
stop=stop_after_attempt(5),
|
||||
wait=wait_exponential_jitter(initial=1, max=60, jitter=5),
|
||||
reraise=True,
|
||||
):
|
||||
with attempt:
|
||||
return await self._client.llama_extract.get_job(job_id=job_id)
|
||||
|
||||
async def _get_run_with_retry(self, job_id: str) -> ExtractRun:
|
||||
"""Get extraction run with retry logic for transient errors."""
|
||||
async for attempt in AsyncRetrying(
|
||||
retry=retry_if_exception(_is_retryable_error),
|
||||
stop=stop_after_attempt(3),
|
||||
wait=wait_exponential_jitter(initial=1, max=20, jitter=3),
|
||||
reraise=True,
|
||||
):
|
||||
with attempt:
|
||||
return await self._client.llama_extract.get_run_by_job_id(job_id=job_id)
|
||||
|
||||
async def _wait_for_job_result(self, job_id: str) -> Optional[ExtractRun]:
|
||||
"""Wait for and return the results of an extraction job."""
|
||||
start = time.perf_counter()
|
||||
tries = 0
|
||||
|
||||
while True:
|
||||
await asyncio.sleep(self.check_interval)
|
||||
tries += 1
|
||||
job = await self._client.llama_extract.get_job(
|
||||
job_id=job_id,
|
||||
)
|
||||
|
||||
if job.status == StatusEnum.SUCCESS:
|
||||
return await self._client.llama_extract.get_run_by_job_id(
|
||||
job_id=job_id,
|
||||
)
|
||||
elif job.status == StatusEnum.PENDING:
|
||||
end = time.perf_counter()
|
||||
if end - start > self.max_timeout:
|
||||
raise Exception(f"Timeout while extracting the file: {job_id}")
|
||||
if self._verbose and tries % 10 == 0:
|
||||
print(".", end="", flush=True)
|
||||
continue
|
||||
else:
|
||||
warnings.warn(
|
||||
f"Failure in job: {job_id}, status: {job.status}, error: {job.error}"
|
||||
)
|
||||
return await self._client.llama_extract.get_run_by_job_id(
|
||||
job_id=job_id,
|
||||
)
|
||||
try:
|
||||
job = await self._get_job_with_retry(job_id)
|
||||
|
||||
if job.status == StatusEnum.SUCCESS:
|
||||
return await self._get_run_with_retry(job_id)
|
||||
elif job.status == StatusEnum.PENDING:
|
||||
end = time.perf_counter()
|
||||
if end - start > self.max_timeout:
|
||||
raise Exception(f"Timeout while extracting the file: {job_id}")
|
||||
if self._verbose and tries % 10 == 0:
|
||||
print(".", end="", flush=True)
|
||||
continue
|
||||
else:
|
||||
warnings.warn(
|
||||
f"Failure in job: {job_id}, status: {job.status}, error: {job.error}"
|
||||
)
|
||||
return await self._get_run_with_retry(job_id)
|
||||
|
||||
except Exception as e:
|
||||
# If we get a non-retryable error or all retries are exhausted, re-raise
|
||||
if self._verbose:
|
||||
print(f"\nError in job polling for {job_id}: {e}")
|
||||
raise e
|
||||
|
||||
def save(self) -> None:
|
||||
"""Persist the extraction agent's schema and config to the database.
|
||||
@@ -307,10 +366,8 @@ class ExtractionAgent:
|
||||
self._agent = self._run_in_thread(
|
||||
self._client.llama_extract.update_extraction_agent(
|
||||
extraction_agent_id=self.id,
|
||||
request=ExtractAgentUpdate(
|
||||
data_schema=self.data_schema,
|
||||
config=self.config,
|
||||
),
|
||||
data_schema=self.data_schema,
|
||||
config=self.config,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -602,7 +659,7 @@ class LlamaExtract(BaseComponent):
|
||||
httpx_timeout=httpx_timeout,
|
||||
verbose=verbose,
|
||||
)
|
||||
self._httpx_client = httpx.AsyncClient(verify=verify, timeout=httpx_timeout)
|
||||
self._httpx_client = httpx.AsyncClient(verify=verify, timeout=httpx_timeout) # type: ignore
|
||||
self.verify = verify
|
||||
self.httpx_timeout = httpx_timeout
|
||||
|
||||
@@ -614,21 +671,8 @@ class LlamaExtract(BaseComponent):
|
||||
self._thread_pool = ThreadPoolExecutor(
|
||||
max_workers=min(10, (os.cpu_count() or 1) + 4)
|
||||
)
|
||||
# Fetch default project id if not provided
|
||||
if not project_id:
|
||||
project_id = os.getenv("LLAMA_CLOUD_PROJECT_ID", None)
|
||||
if not project_id:
|
||||
print("No project_id provided, fetching default project.")
|
||||
projects: List[Project] = self._run_in_thread(
|
||||
self._async_client.projects.list_projects()
|
||||
)
|
||||
default_project = [p for p in projects if p.is_default]
|
||||
if not default_project:
|
||||
raise ValueError(
|
||||
"No default project found. Please provide a project_id."
|
||||
)
|
||||
project_id = default_project[0].id
|
||||
|
||||
self._project_id = project_id
|
||||
self._organization_id = organization_id
|
||||
|
||||
@@ -659,11 +703,7 @@ class LlamaExtract(BaseComponent):
|
||||
ExtractionAgent: The created extraction agent
|
||||
"""
|
||||
if config is not None:
|
||||
if config.extraction_mode == ExtractMode.ACCURATE:
|
||||
warnings.warn(
|
||||
"ACCURATE extraction mode is deprecated. Using BALANCED instead."
|
||||
)
|
||||
config.extraction_mode = ExtractMode.BALANCED
|
||||
_extraction_config_warning(config)
|
||||
else:
|
||||
config = DEFAULT_EXTRACT_CONFIG
|
||||
|
||||
@@ -680,11 +720,9 @@ class LlamaExtract(BaseComponent):
|
||||
self._async_client.llama_extract.create_extraction_agent(
|
||||
project_id=self._project_id,
|
||||
organization_id=self._organization_id,
|
||||
request=ExtractAgentCreate(
|
||||
name=name,
|
||||
data_schema=data_schema,
|
||||
config=config,
|
||||
),
|
||||
name=name,
|
||||
data_schema=data_schema,
|
||||
config=config,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -698,6 +736,8 @@ class LlamaExtract(BaseComponent):
|
||||
num_workers=self.num_workers,
|
||||
show_progress=self.show_progress,
|
||||
verbose=self.verbose,
|
||||
verify=self.verify,
|
||||
httpx_timeout=self.httpx_timeout,
|
||||
)
|
||||
|
||||
def get_agent(
|
||||
|
||||
@@ -32,3 +32,9 @@ def augment_async_errors() -> Generator[None, None, None]:
|
||||
|
||||
JSONType = Union[Dict[str, Any], List[Any], str, int, float, bool, None]
|
||||
JSONObjectType = Dict[str, JSONType]
|
||||
|
||||
|
||||
class ExperimentalWarning(Warning):
|
||||
"""Warning for experimental features."""
|
||||
|
||||
pass
|
||||
|
||||
@@ -1,3 +1,8 @@
|
||||
from llama_cloud_services.parse.base import LlamaParse, ResultType
|
||||
from llama_cloud_services.parse.base import (
|
||||
LlamaParse,
|
||||
ResultType,
|
||||
ParsingMode,
|
||||
FailedPageMode,
|
||||
)
|
||||
|
||||
__all__ = ["LlamaParse", "ResultType"]
|
||||
__all__ = ["LlamaParse", "ResultType", "ParsingMode", "FailedPageMode"]
|
||||
|
||||
@@ -2,30 +2,41 @@ import asyncio
|
||||
import mimetypes
|
||||
import os
|
||||
import time
|
||||
import warnings
|
||||
from contextlib import asynccontextmanager
|
||||
from copy import deepcopy
|
||||
from enum import Enum
|
||||
from io import BufferedIOBase
|
||||
from pathlib import Path, PurePath, PurePosixPath
|
||||
from typing import Any, AsyncGenerator, Dict, List, Optional, Union
|
||||
from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple, Union
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
from fsspec import AbstractFileSystem
|
||||
from llama_index.core.async_utils import asyncio_run, run_jobs
|
||||
from llama_index.core.bridge.pydantic import Field, PrivateAttr, field_validator
|
||||
from llama_index.core.bridge.pydantic import (
|
||||
Field,
|
||||
PrivateAttr,
|
||||
field_validator,
|
||||
model_validator,
|
||||
)
|
||||
from llama_index.core.constants import DEFAULT_BASE_URL
|
||||
from llama_index.core.readers.base import BasePydanticReader
|
||||
from llama_index.core.readers.file.base import get_default_fs
|
||||
from llama_index.core.schema import Document
|
||||
|
||||
from llama_cloud_services.utils import check_extra_params
|
||||
from llama_cloud_services.parse.types import JobResult
|
||||
from llama_cloud_services.parse.utils import (
|
||||
SUPPORTED_FILE_TYPES,
|
||||
ResultType,
|
||||
ParsingMode,
|
||||
FailedPageMode,
|
||||
expand_target_pages,
|
||||
nest_asyncio_err,
|
||||
nest_asyncio_msg,
|
||||
make_api_request,
|
||||
partition_pages,
|
||||
)
|
||||
|
||||
# can put in a path to the file or the file bytes itself
|
||||
@@ -55,6 +66,36 @@ def build_url(
|
||||
return base_url
|
||||
|
||||
|
||||
class JobFailedException(Exception):
|
||||
"""Parse job failed exception."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
job_id: str,
|
||||
status: str,
|
||||
error_code: Optional[str] = None,
|
||||
error_message: Optional[str] = None,
|
||||
):
|
||||
exception_str = (
|
||||
f"Job ID: {job_id} failed with status: {status}, "
|
||||
f'Error code: {error_code or "No error code found"}, '
|
||||
f'Error message: {error_message or "No error message found"}'
|
||||
)
|
||||
super().__init__(exception_str)
|
||||
self.job_id = job_id
|
||||
self.status = status
|
||||
self.error_code = error_code
|
||||
self.error_message = error_message
|
||||
|
||||
@classmethod
|
||||
def from_result(cls, result_json: Dict[str, Any]) -> "JobFailedException":
|
||||
job_id = result_json["id"]
|
||||
status = result_json["status"]
|
||||
error_code = result_json.get("error_code")
|
||||
error_message = result_json.get("error_message")
|
||||
return cls(job_id, status, error_code=error_code, error_message=error_message)
|
||||
|
||||
|
||||
class BackoffPattern(str, Enum):
|
||||
"""Backoff pattern for polling."""
|
||||
|
||||
@@ -113,7 +154,7 @@ class LlamaParse(BasePydanticReader):
|
||||
num_workers: int = Field(
|
||||
default=4,
|
||||
gt=0,
|
||||
lt=10,
|
||||
lt=20,
|
||||
description="The number of workers to use sending API requests for parsing.",
|
||||
)
|
||||
result_type: ResultType = Field(
|
||||
@@ -143,6 +184,10 @@ class LlamaParse(BasePydanticReader):
|
||||
default=False,
|
||||
description="If set to true, the parser will automatically select the best mode to extract text from documents based on the rules provide. Will use the 'accurate' default mode by default and will upgrade page that match the rule to Premium mode.",
|
||||
)
|
||||
auto_mode_configuration_json: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A JSON string containing the configuration for the auto mode. If set, the parser will use the provided configuration for the auto mode.",
|
||||
)
|
||||
auto_mode_trigger_on_image_in_page: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="If auto_mode is set to true, the parser will upgrade the page that contain an image to Premium mode.",
|
||||
@@ -228,6 +273,10 @@ class LlamaParse(BasePydanticReader):
|
||||
default=False,
|
||||
description="Whether to guess the sheet names of the xlsx file.",
|
||||
)
|
||||
high_res_ocr: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="If set to true, the parser will use high resolution OCR to extract text from images. This will increase the accuracy of the parsing job, but reduce the speed.",
|
||||
)
|
||||
html_make_all_elements_visible: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="If set to true, when parsing HTML the parser will consider all elements display not element as display block.",
|
||||
@@ -291,6 +340,14 @@ class LlamaParse(BasePydanticReader):
|
||||
default=False,
|
||||
description="If set to true, the parser will output tables as HTML in the markdown.",
|
||||
)
|
||||
outlined_table_extraction: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="If set to true, the parser will use a dedicated approach to extract tables with outlined cells. This is useful for documents with spreadsheet-like tables where cells are outlined with borders. This could lead to false positives, so use with caution.",
|
||||
)
|
||||
page_error_tolerance: Optional[float] = Field(
|
||||
default=None,
|
||||
description="The error tolerance for the number of pages with error in a doc (percentage express as 0-1). If we fail to parse a greater percentage of pages than the tolerance value we fail the job.",
|
||||
)
|
||||
page_prefix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A templated prefix to add to the beginning of each page. If it contain `{page_number}`, it will be replaced by the page number.",
|
||||
@@ -303,7 +360,7 @@ class LlamaParse(BasePydanticReader):
|
||||
default=None,
|
||||
description="A templated suffix to add to the beginning of each page. If it contain `{page_number}`, it will be replaced by the page number.",
|
||||
)
|
||||
parse_mode: Optional[str] = Field(
|
||||
parse_mode: Optional[Union[ParsingMode, str]] = Field(
|
||||
default=None,
|
||||
description="The parsing mode to use, see ParsingMode enum for possible values ",
|
||||
)
|
||||
@@ -311,10 +368,26 @@ class LlamaParse(BasePydanticReader):
|
||||
default=False,
|
||||
description="Use our best parser mode if set to True.",
|
||||
)
|
||||
preset: Optional[str] = Field(
|
||||
default=None,
|
||||
description="The preset to use for the parser. If set, the parser will use the preset configuration. See LlamaParse documentation for available presets. Preset override most other parameters.",
|
||||
)
|
||||
preserve_layout_alignment_across_pages: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="Preserve grid alignment across page in text mode.",
|
||||
)
|
||||
replace_failed_page_mode: Optional[FailedPageMode] = Field(
|
||||
default=None,
|
||||
description="The mode to use to replace the failed page, see FailedPageMode enum for possible value. If set, the parser will replace the failed page with the specified mode. If not set, the default mode (raw_text) will be used.",
|
||||
)
|
||||
replace_failed_page_with_error_message_prefix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A prefix to add before error message in failed pages. If not set, no prefix will be used.",
|
||||
)
|
||||
replace_failed_page_with_error_message_suffix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A suffix to add after error message in failed pages. If not set, no suffix will be used.",
|
||||
)
|
||||
skip_diagonal_text: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="If set to true, the parser will ignore diagonal text (when the text rotation in degrees modulo 90 is not 0).",
|
||||
@@ -384,10 +457,42 @@ class LlamaParse(BasePydanticReader):
|
||||
default=None,
|
||||
description="The model name for the vendor multimodal API.",
|
||||
)
|
||||
model: Optional[str] = Field(
|
||||
default=None,
|
||||
description="The document model name to be used with `parse_with_agent`.",
|
||||
)
|
||||
webhook_url: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A URL that needs to be called at the end of the parsing job.",
|
||||
)
|
||||
partition_pages: Optional[int] = Field(
|
||||
default=None,
|
||||
description="If set, documents will automatically be partitioned into segments containing the specified number of pages at most. Parsing will be split into separate jobs for each partition segment. Can be used in combination with targetPages and maxPages.",
|
||||
)
|
||||
hide_headers: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="Whether to hide page header in output markdown.",
|
||||
)
|
||||
hide_footers: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="Whether to hide page footers in output markdown.",
|
||||
)
|
||||
page_header_suffix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A suffix to add to the page header in the output markdown.",
|
||||
)
|
||||
page_header_prefix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A prefix to add to the page header in the output markdown.",
|
||||
)
|
||||
page_footer_suffix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A suffix to add to the page footer in the output markdown.",
|
||||
)
|
||||
page_footer_prefix: Optional[str] = Field(
|
||||
default=None,
|
||||
description="A prefix to add to the page footer in the output markdown.",
|
||||
)
|
||||
|
||||
# Deprecated
|
||||
bounding_box: Optional[str] = Field(
|
||||
@@ -427,6 +532,21 @@ class LlamaParse(BasePydanticReader):
|
||||
description="Whether to use the vendor multimodal API.",
|
||||
)
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def warn_extra_params(cls, data: Dict[str, Any]) -> Dict[str, Any]:
|
||||
extra_params, suggestions = check_extra_params(cls, data)
|
||||
if extra_params:
|
||||
suggestions = [f"\n - {suggestion}" for suggestion in suggestions]
|
||||
suggestions_str = "".join(suggestions)
|
||||
warnings.warn(
|
||||
"The following parameters are unused: "
|
||||
+ ", ".join(extra_params)
|
||||
+ f".\n{suggestions_str}",
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
@field_validator("api_key", mode="before", check_fields=True)
|
||||
@classmethod
|
||||
def validate_api_key(cls, v: str) -> str:
|
||||
@@ -514,6 +634,7 @@ class LlamaParse(BasePydanticReader):
|
||||
file_input: FileInput,
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
partition_target_pages: Optional[str] = None,
|
||||
) -> str:
|
||||
files = None
|
||||
file_handle = None
|
||||
@@ -564,6 +685,9 @@ class LlamaParse(BasePydanticReader):
|
||||
if self.auto_mode:
|
||||
data["auto_mode"] = self.auto_mode
|
||||
|
||||
if self.auto_mode_configuration_json is not None:
|
||||
data["auto_mode_configuration_json"] = self.auto_mode_configuration_json
|
||||
|
||||
if self.auto_mode_trigger_on_image_in_page:
|
||||
data[
|
||||
"auto_mode_trigger_on_image_in_page"
|
||||
@@ -661,6 +785,9 @@ class LlamaParse(BasePydanticReader):
|
||||
if self.html_make_all_elements_visible:
|
||||
data["html_make_all_elements_visible"] = self.html_make_all_elements_visible
|
||||
|
||||
if self.high_res_ocr:
|
||||
data["high_res_ocr"] = self.high_res_ocr
|
||||
|
||||
if self.html_remove_fixed_elements:
|
||||
data["html_remove_fixed_elements"] = self.html_remove_fixed_elements
|
||||
|
||||
@@ -723,9 +850,33 @@ class LlamaParse(BasePydanticReader):
|
||||
if self.output_tables_as_HTML:
|
||||
data["output_tables_as_HTML"] = self.output_tables_as_HTML
|
||||
|
||||
if self.outlined_table_extraction:
|
||||
data["outlined_table_extraction"] = self.outlined_table_extraction
|
||||
|
||||
if self.page_error_tolerance is not None:
|
||||
data["page_error_tolerance"] = self.page_error_tolerance
|
||||
|
||||
if self.page_prefix is not None:
|
||||
data["page_prefix"] = self.page_prefix
|
||||
|
||||
if self.hide_headers:
|
||||
data["hide_headers"] = self.hide_headers
|
||||
|
||||
if self.hide_footers:
|
||||
data["hide_footers"] = self.hide_footers
|
||||
|
||||
if self.page_header_suffix is not None:
|
||||
data["page_header_suffix"] = self.page_header_suffix
|
||||
|
||||
if self.page_header_prefix is not None:
|
||||
data["page_header_prefix"] = self.page_header_prefix
|
||||
|
||||
if self.page_footer_suffix is not None:
|
||||
data["page_footer_suffix"] = self.page_footer_suffix
|
||||
|
||||
if self.page_footer_prefix is not None:
|
||||
data["page_footer_prefix"] = self.page_footer_prefix
|
||||
|
||||
# only send page separator to server if it is not None
|
||||
# as if a null, "" string is sent the server will then ignore the page separator instead of using the default
|
||||
if self.page_separator is not None:
|
||||
@@ -751,6 +902,22 @@ class LlamaParse(BasePydanticReader):
|
||||
"preserve_layout_alignment_across_pages"
|
||||
] = self.preserve_layout_alignment_across_pages
|
||||
|
||||
if self.preset is not None:
|
||||
data["preset"] = self.preset
|
||||
|
||||
if self.replace_failed_page_mode is not None:
|
||||
data["replace_failed_page_mode"] = self.replace_failed_page_mode.value
|
||||
|
||||
if self.replace_failed_page_with_error_message_prefix is not None:
|
||||
data[
|
||||
"replace_failed_page_with_error_message_prefix"
|
||||
] = self.replace_failed_page_with_error_message_prefix
|
||||
|
||||
if self.replace_failed_page_with_error_message_suffix is not None:
|
||||
data[
|
||||
"replace_failed_page_with_error_message_suffix"
|
||||
] = self.replace_failed_page_with_error_message_suffix
|
||||
|
||||
if self.skip_diagonal_text:
|
||||
data["skip_diagonal_text"] = self.skip_diagonal_text
|
||||
|
||||
@@ -786,7 +953,9 @@ class LlamaParse(BasePydanticReader):
|
||||
if self.take_screenshot:
|
||||
data["take_screenshot"] = self.take_screenshot
|
||||
|
||||
if self.target_pages is not None:
|
||||
if partition_target_pages is not None:
|
||||
data["target_pages"] = partition_target_pages
|
||||
elif self.target_pages is not None:
|
||||
data["target_pages"] = self.target_pages
|
||||
if self.user_prompt is not None:
|
||||
data["user_prompt"] = self.user_prompt
|
||||
@@ -799,6 +968,9 @@ class LlamaParse(BasePydanticReader):
|
||||
if self.vendor_multimodal_model_name is not None:
|
||||
data["vendor_multimodal_model_name"] = self.vendor_multimodal_model_name
|
||||
|
||||
if self.model is not None:
|
||||
data["model"] = self.model
|
||||
|
||||
if self.webhook_url is not None:
|
||||
data["webhook_url"] = self.webhook_url
|
||||
|
||||
@@ -880,15 +1052,7 @@ class LlamaParse(BasePydanticReader):
|
||||
print(".", end="", flush=True)
|
||||
current_interval = self._calculate_backoff(current_interval)
|
||||
else:
|
||||
error_code = result_json.get("error_code", "No error code found")
|
||||
error_message = result_json.get(
|
||||
"error_message", "No error message found"
|
||||
)
|
||||
exception_str = (
|
||||
f"Job ID: {job_id} failed with status: {status}, "
|
||||
f"Error code: {error_code}, Error message: {error_message}"
|
||||
)
|
||||
raise Exception(exception_str)
|
||||
raise JobFailedException.from_result(result_json)
|
||||
except (
|
||||
httpx.ConnectError,
|
||||
httpx.ReadError,
|
||||
@@ -897,6 +1061,7 @@ class LlamaParse(BasePydanticReader):
|
||||
httpx.ReadTimeout,
|
||||
httpx.WriteTimeout,
|
||||
httpx.HTTPStatusError,
|
||||
httpx.RemoteProtocolError,
|
||||
) as err:
|
||||
error_count += 1
|
||||
end = time.time()
|
||||
@@ -911,26 +1076,151 @@ class LlamaParse(BasePydanticReader):
|
||||
)
|
||||
current_interval = self._calculate_backoff(current_interval)
|
||||
|
||||
async def _parse_one(
|
||||
self,
|
||||
file_path: FileInput,
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
result_type: Optional[str] = None,
|
||||
num_workers: Optional[int] = None,
|
||||
) -> List[Tuple[str, Dict[str, Any]]]:
|
||||
if self.partition_pages is None:
|
||||
job_results = [
|
||||
await self._parse_one_unpartitioned(
|
||||
file_path,
|
||||
extra_info=extra_info,
|
||||
fs=fs,
|
||||
result_type=result_type,
|
||||
)
|
||||
]
|
||||
else:
|
||||
job_results = await self._parse_one_partitioned(
|
||||
file_path,
|
||||
extra_info,
|
||||
fs=fs,
|
||||
result_type=result_type,
|
||||
num_workers=num_workers,
|
||||
)
|
||||
return job_results
|
||||
|
||||
async def _parse_one_unpartitioned(
|
||||
self,
|
||||
file_path: FileInput,
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
result_type: Optional[str] = None,
|
||||
**create_kwargs: Any,
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""Create one parse job and wait for the result."""
|
||||
job_id = await self._create_job(
|
||||
file_path, extra_info=extra_info, fs=fs, **create_kwargs
|
||||
)
|
||||
if self.verbose:
|
||||
print("Started parsing the file under job_id %s" % job_id)
|
||||
result = await self._get_job_result(
|
||||
job_id, result_type or self.result_type.value, verbose=self.verbose
|
||||
)
|
||||
return job_id, result
|
||||
|
||||
async def _parse_one_partitioned(
|
||||
self,
|
||||
file_path: FileInput,
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
result_type: Optional[str] = None,
|
||||
num_workers: Optional[int] = None,
|
||||
) -> List[Tuple[str, Dict[str, Any]]]:
|
||||
"""Partition a file and run separate parse jobs per partition segment."""
|
||||
assert self.partition_pages is not None
|
||||
|
||||
num_workers = num_workers or self.num_workers
|
||||
if num_workers < 1:
|
||||
raise ValueError("Invalid number of workers")
|
||||
if self.target_pages is not None:
|
||||
jobs = [
|
||||
self._parse_one_unpartitioned(
|
||||
file_path,
|
||||
extra_info=extra_info,
|
||||
fs=fs,
|
||||
result_type=result_type,
|
||||
partition_target_pages=target_pages,
|
||||
)
|
||||
for target_pages in partition_pages(
|
||||
expand_target_pages(self.target_pages),
|
||||
self.partition_pages,
|
||||
max_pages=self.max_pages,
|
||||
)
|
||||
]
|
||||
return await run_jobs(
|
||||
jobs,
|
||||
workers=num_workers,
|
||||
desc="Getting job results",
|
||||
show_progress=self.show_progress,
|
||||
)
|
||||
|
||||
total = 0
|
||||
results: List[Tuple[str, Dict[str, Any]]] = []
|
||||
while self.max_pages is None or total < self.max_pages:
|
||||
if (
|
||||
self.max_pages is not None
|
||||
and total + self.partition_pages >= self.max_pages
|
||||
):
|
||||
size = self.max_pages - total
|
||||
else:
|
||||
size = self.partition_pages
|
||||
if not size:
|
||||
break
|
||||
try:
|
||||
# Fetch JSON result type first to get accurate pagination data
|
||||
# and then fetch the user's desired result type if needed
|
||||
job_id, json_result = await self._parse_one_unpartitioned(
|
||||
file_path,
|
||||
extra_info=extra_info,
|
||||
fs=fs,
|
||||
result_type=ResultType.JSON.value,
|
||||
partition_target_pages=f"{total}-{total + size - 1}",
|
||||
)
|
||||
result_type = result_type or self.result_type.value
|
||||
if result_type == ResultType.JSON.value:
|
||||
job_result = json_result
|
||||
else:
|
||||
job_result = await self._get_job_result(
|
||||
job_id, result_type, verbose=self.verbose
|
||||
)
|
||||
except JobFailedException as e:
|
||||
if results and e.error_code == "NO_DATA_FOUND_IN_FILE":
|
||||
# Expected when we try to read past the end of the file
|
||||
return results
|
||||
raise
|
||||
results.append((job_id, job_result))
|
||||
if len(json_result["pages"]) < size:
|
||||
break
|
||||
total += size
|
||||
return results
|
||||
|
||||
async def _aload_data(
|
||||
self,
|
||||
file_path: FileInput,
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
verbose: bool = False,
|
||||
num_workers: Optional[int] = None,
|
||||
) -> List[Document]:
|
||||
"""Load data from the input path."""
|
||||
try:
|
||||
job_id = await self._create_job(file_path, extra_info=extra_info, fs=fs)
|
||||
if verbose:
|
||||
print("Started parsing the file under job_id %s" % job_id)
|
||||
|
||||
result = await self._get_job_result(
|
||||
job_id, self.result_type.value, verbose=verbose
|
||||
)
|
||||
|
||||
results = [
|
||||
job_result
|
||||
for _, job_result in await self._parse_one(
|
||||
file_path, extra_info, fs=fs, num_workers=num_workers
|
||||
)
|
||||
]
|
||||
# Flatten the resulting doc if it was partitioned
|
||||
separator = self.page_separator or _DEFAULT_SEPARATOR
|
||||
docs = [
|
||||
Document(
|
||||
text=result[self.result_type.value],
|
||||
text=separator.join(
|
||||
result[self.result_type.value] for result in results
|
||||
),
|
||||
metadata=extra_info or {},
|
||||
)
|
||||
]
|
||||
@@ -953,7 +1243,11 @@ class LlamaParse(BasePydanticReader):
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
) -> List[Document]:
|
||||
"""Load data from the input path."""
|
||||
"""Load data from the input path.
|
||||
|
||||
File(s) which were partitioned before parsing will be loaded as a single
|
||||
re-assembled Document.
|
||||
"""
|
||||
if isinstance(file_path, (str, PurePosixPath, Path, bytes, BufferedIOBase)):
|
||||
return await self._aload_data(
|
||||
file_path, extra_info=extra_info, fs=fs, verbose=self.verbose
|
||||
@@ -965,6 +1259,7 @@ class LlamaParse(BasePydanticReader):
|
||||
extra_info=extra_info,
|
||||
fs=fs,
|
||||
verbose=self.verbose and not self.show_progress,
|
||||
num_workers=1,
|
||||
)
|
||||
for f in file_path
|
||||
]
|
||||
@@ -1003,6 +1298,34 @@ class LlamaParse(BasePydanticReader):
|
||||
else:
|
||||
raise e
|
||||
|
||||
async def _aparse_one(
|
||||
self,
|
||||
file_path: FileInput,
|
||||
file_name: str,
|
||||
extra_info: Optional[dict] = None,
|
||||
fs: Optional[AbstractFileSystem] = None,
|
||||
num_workers: Optional[int] = None,
|
||||
) -> List[JobResult]:
|
||||
job_results = await self._parse_one(
|
||||
file_path,
|
||||
extra_info,
|
||||
fs=fs,
|
||||
result_type=ResultType.JSON.value,
|
||||
num_workers=num_workers,
|
||||
)
|
||||
return [
|
||||
JobResult(
|
||||
job_id=job_id,
|
||||
file_name=file_name,
|
||||
job_result=job_result,
|
||||
api_key=self.api_key,
|
||||
base_url=self.base_url,
|
||||
client=self.aclient,
|
||||
page_separator=self.page_separator or _DEFAULT_SEPARATOR,
|
||||
)
|
||||
for job_id, job_result in job_results
|
||||
]
|
||||
|
||||
async def aparse(
|
||||
self,
|
||||
file_path: Union[List[FileInput], FileInput],
|
||||
@@ -1021,14 +1344,10 @@ class LlamaParse(BasePydanticReader):
|
||||
fs: Optional filesystem to use for reading files.
|
||||
|
||||
Returns:
|
||||
JobResult object or list of JobResult objects if multiple files were provided
|
||||
JobResult object or list of JobResult objects if either multiple files were provided or file(s) were partitioned before parsing.
|
||||
"""
|
||||
|
||||
if isinstance(file_path, (str, PurePosixPath, Path, bytes, BufferedIOBase)):
|
||||
job_id = await self._create_job(file_path, extra_info=extra_info, fs=fs)
|
||||
if self.verbose:
|
||||
print("Started parsing the file under job_id %s" % job_id)
|
||||
|
||||
if isinstance(file_path, (bytes, BufferedIOBase)):
|
||||
if not extra_info or "file_name" not in extra_info:
|
||||
raise ValueError(
|
||||
@@ -1037,29 +1356,12 @@ class LlamaParse(BasePydanticReader):
|
||||
file_name = extra_info["file_name"]
|
||||
else:
|
||||
file_name = str(file_path)
|
||||
|
||||
job_result = await self._get_job_result(
|
||||
job_id, ResultType.JSON.value, verbose=self.verbose
|
||||
)
|
||||
return JobResult(
|
||||
job_id=job_id,
|
||||
file_name=file_name,
|
||||
job_result=job_result,
|
||||
api_key=self.api_key,
|
||||
base_url=self.base_url,
|
||||
client=self.aclient,
|
||||
page_separator=self.page_separator or _DEFAULT_SEPARATOR,
|
||||
result = await self._aparse_one(
|
||||
file_path, file_name, extra_info=extra_info, fs=fs
|
||||
)
|
||||
return result[0] if len(result) == 1 else result
|
||||
|
||||
elif isinstance(file_path, list):
|
||||
jobs = [
|
||||
self._create_job(
|
||||
f,
|
||||
extra_info=extra_info,
|
||||
fs=fs,
|
||||
)
|
||||
for f in file_path
|
||||
]
|
||||
file_names = []
|
||||
for f in file_path:
|
||||
if isinstance(f, (bytes, BufferedIOBase)):
|
||||
@@ -1071,40 +1373,24 @@ class LlamaParse(BasePydanticReader):
|
||||
else:
|
||||
file_names.append(str(f))
|
||||
|
||||
job_results = []
|
||||
try:
|
||||
job_ids = await run_jobs(
|
||||
jobs,
|
||||
workers=self.num_workers,
|
||||
desc="Creating parsing jobs",
|
||||
show_progress=self.show_progress,
|
||||
)
|
||||
|
||||
job_results = await run_jobs(
|
||||
for result in await run_jobs(
|
||||
[
|
||||
self._get_job_result(
|
||||
job_id, ResultType.JSON.value, verbose=self.verbose
|
||||
self._aparse_one(
|
||||
f,
|
||||
file_names[i],
|
||||
extra_info=extra_info,
|
||||
fs=fs,
|
||||
num_workers=1,
|
||||
)
|
||||
for job_id in job_ids
|
||||
for i, f in enumerate(file_path)
|
||||
],
|
||||
workers=self.num_workers,
|
||||
desc="Getting job results",
|
||||
show_progress=self.show_progress,
|
||||
)
|
||||
|
||||
# Create JobResults just using the job_ids and job_results
|
||||
job_results = [
|
||||
JobResult(
|
||||
job_id=job_id,
|
||||
file_name=file_names[i],
|
||||
job_result=job_results[i],
|
||||
api_key=self.api_key,
|
||||
base_url=self.base_url,
|
||||
client=self.aclient,
|
||||
page_separator=self.page_separator or _DEFAULT_SEPARATOR,
|
||||
)
|
||||
for i, job_id in enumerate(job_ids)
|
||||
]
|
||||
|
||||
):
|
||||
job_results.extend(result)
|
||||
return job_results
|
||||
|
||||
except RuntimeError as e:
|
||||
@@ -1146,20 +1432,27 @@ class LlamaParse(BasePydanticReader):
|
||||
raise e
|
||||
|
||||
async def _aget_json(
|
||||
self, file_path: FileInput, extra_info: Optional[dict] = None
|
||||
self,
|
||||
file_path: FileInput,
|
||||
extra_info: Optional[dict] = None,
|
||||
num_workers: Optional[int] = None,
|
||||
) -> List[dict]:
|
||||
"""Load data from the input path."""
|
||||
try:
|
||||
job_id = await self._create_job(file_path, extra_info=extra_info)
|
||||
if self.verbose:
|
||||
print("Started parsing the file under job_id %s" % job_id)
|
||||
result = await self._get_job_result(job_id, "json")
|
||||
result["job_id"] = job_id
|
||||
job_results = await self._parse_one(
|
||||
file_path,
|
||||
extra_info=extra_info,
|
||||
result_type=ResultType.JSON.value,
|
||||
num_workers=num_workers,
|
||||
)
|
||||
|
||||
if not isinstance(file_path, (bytes, BufferedIOBase)):
|
||||
result["file_path"] = str(file_path)
|
||||
|
||||
return [result]
|
||||
results = []
|
||||
for job_id, job_result in job_results:
|
||||
job_result["job_id"] = job_id
|
||||
if not isinstance(file_path, (bytes, BufferedIOBase)):
|
||||
job_result["file_path"] = str(file_path)
|
||||
results.append(job_result)
|
||||
return results
|
||||
except Exception as e:
|
||||
file_repr = file_path if isinstance(file_path, str) else "<bytes/buffer>"
|
||||
print(f"Error while parsing the file '{file_repr}':", e)
|
||||
@@ -1174,7 +1467,7 @@ class LlamaParse(BasePydanticReader):
|
||||
extra_info: Optional[dict] = None,
|
||||
) -> List[dict]:
|
||||
"""Load data from the input path."""
|
||||
if isinstance(file_path, (str, Path)):
|
||||
if isinstance(file_path, (str, PurePosixPath, Path, bytes, BufferedIOBase)):
|
||||
return await self._aget_json(file_path, extra_info=extra_info)
|
||||
elif isinstance(file_path, list):
|
||||
jobs = [self._aget_json(f, extra_info=extra_info) for f in file_path]
|
||||
@@ -1195,7 +1488,7 @@ class LlamaParse(BasePydanticReader):
|
||||
raise e
|
||||
else:
|
||||
raise ValueError(
|
||||
"The input file_path must be a string or a list of strings."
|
||||
"The input file_path must be a string, Path, bytes, BufferedIOBase, or a list of these types."
|
||||
)
|
||||
|
||||
def get_json_result(
|
||||
|
||||
@@ -14,14 +14,14 @@ PAGE_REGEX = r"page[-_](\d+)\.jpg$"
|
||||
class JobMetadata(BaseModel):
|
||||
"""Metadata about the job."""
|
||||
|
||||
job_credits_usage: int = Field(
|
||||
default_factory=dict, description="The credits usage for the job."
|
||||
job_pages: int = Field(default=0, description="The number of pages in the job.")
|
||||
job_auto_mode_triggered_pages: Optional[int] = Field(
|
||||
default=None,
|
||||
description="The number of pages that triggered auto mode (thus increasing the cost).",
|
||||
)
|
||||
job_pages: int = Field(description="The number of pages in the job.")
|
||||
job_auto_mode_triggered_pages: int = Field(
|
||||
description="The number of pages that triggered auto mode (thus increasing the cost)."
|
||||
job_is_cache_hit: bool = Field(
|
||||
default=False, description="Whether the job was a cache hit."
|
||||
)
|
||||
job_is_cache_hit: bool = Field(description="Whether the job was a cache hit.")
|
||||
|
||||
|
||||
class BBox(BaseModel):
|
||||
@@ -46,22 +46,34 @@ class PageItem(BaseModel):
|
||||
md: Optional[str] = Field(
|
||||
default=None, description="The markdown-formatted content of the item."
|
||||
)
|
||||
rows: Optional[List[List[str]]] = Field(
|
||||
rows: Optional[List[List[Any]]] = Field(
|
||||
default=None, description="The rows of the item."
|
||||
)
|
||||
bBox: BBox = Field(description="The bounding box of the item.")
|
||||
bBox: Optional[BBox] = Field(
|
||||
default=None, description="The bounding box of the item."
|
||||
)
|
||||
|
||||
|
||||
class ImageItem(BaseModel):
|
||||
"""An image in a page."""
|
||||
|
||||
name: str = Field(description="The name of the image.")
|
||||
height: float = Field(description="The height of the image.")
|
||||
width: float = Field(description="The width of the image.")
|
||||
x: float = Field(description="The x-coordinate of the image.")
|
||||
y: float = Field(description="The y-coordinate of the image.")
|
||||
original_width: int = Field(description="The original width of the image.")
|
||||
original_height: int = Field(description="The original height of the image.")
|
||||
height: Optional[float] = Field(
|
||||
default=None, description="The height of the image."
|
||||
)
|
||||
width: Optional[float] = Field(default=None, description="The width of the image.")
|
||||
x: Optional[float] = Field(
|
||||
default=None, description="The x-coordinate of the image."
|
||||
)
|
||||
y: Optional[float] = Field(
|
||||
default=None, description="The y-coordinate of the image."
|
||||
)
|
||||
original_width: Optional[int] = Field(
|
||||
default=None, description="The original width of the image."
|
||||
)
|
||||
original_height: Optional[int] = Field(
|
||||
default=None, description="The original height of the image."
|
||||
)
|
||||
type: Optional[str] = Field(default=None, description="The type of the image.")
|
||||
|
||||
|
||||
@@ -71,7 +83,9 @@ class LayoutItem(BaseModel):
|
||||
image: str = Field(description="The name of the image containing the layout item")
|
||||
confidence: float = Field(description="The confidence of the layout item.")
|
||||
label: str = Field(description="The label of the layout item.")
|
||||
bbox: BBox = Field(description="The bounding box of the layout item.")
|
||||
bbox: Optional[BBox] = Field(
|
||||
default=None, description="The bounding box of the layout item."
|
||||
)
|
||||
isLikelyNoise: bool = Field(description="Whether the layout item is likely noise.")
|
||||
|
||||
|
||||
@@ -79,18 +93,24 @@ class ChartItem(BaseModel):
|
||||
"""A chart in a page."""
|
||||
|
||||
name: str = Field(description="The name of the chart.")
|
||||
x: float = Field(description="The x-coordinate of the chart.")
|
||||
y: float = Field(description="The y-coordinate of the chart.")
|
||||
width: float = Field(description="The width of the chart.")
|
||||
height: float = Field(description="The height of the chart.")
|
||||
x: Optional[float] = Field(
|
||||
default=None, description="The x-coordinate of the chart."
|
||||
)
|
||||
y: Optional[float] = Field(
|
||||
default=None, description="The y-coordinate of the chart."
|
||||
)
|
||||
width: Optional[float] = Field(default=None, description="The width of the chart.")
|
||||
height: Optional[float] = Field(
|
||||
default=None, description="The height of the chart."
|
||||
)
|
||||
|
||||
|
||||
class Page(BaseModel):
|
||||
"""A page of the document."""
|
||||
|
||||
page: int = Field(description="The page number.")
|
||||
text: str = Field(description="The text of the page.")
|
||||
md: str = Field(description="The markdown of the page.")
|
||||
text: Optional[str] = Field(default=None, description="The text of the page.")
|
||||
md: Optional[str] = Field(default=None, description="The markdown of the page.")
|
||||
images: List[ImageItem] = Field(
|
||||
default_factory=list,
|
||||
description="The names of the image IDs in the page, including both objects and page screenshots.",
|
||||
@@ -107,32 +127,43 @@ class Page(BaseModel):
|
||||
items: List[PageItem] = Field(
|
||||
default_factory=list, description="The items in the page."
|
||||
)
|
||||
status: str = Field(description="The status of the page.")
|
||||
status: Optional[str] = Field(default=None, description="The status of the page.")
|
||||
links: List[SerializeAsAny[Any]] = Field(
|
||||
default_factory=list, description="The links in the page."
|
||||
)
|
||||
width: float = Field(description="The width of the page.")
|
||||
height: float = Field(description="The height of the page.")
|
||||
triggeredAutoMode: bool = Field(
|
||||
description="Whether the page triggered auto mode (thus increasing the cost)."
|
||||
width: Optional[float] = Field(default=None, description="The width of the page.")
|
||||
height: Optional[float] = Field(default=None, description="The height of the page.")
|
||||
triggeredAutoMode: Optional[bool] = Field(
|
||||
default=False,
|
||||
description="Whether the page triggered auto mode (thus increasing the cost).",
|
||||
)
|
||||
parsingMode: str = Field(
|
||||
default="", description="The parsing mode used for the page."
|
||||
)
|
||||
parsingMode: str = Field(description="The parsing mode used for the page.")
|
||||
structuredData: Optional[Dict[str, Any]] = Field(
|
||||
description="The structured data of the page."
|
||||
default=None, description="The structured data of the page."
|
||||
)
|
||||
noStructuredContent: bool = Field(
|
||||
description="Whether the page has no structured data."
|
||||
default=True, description="Whether the page has no structured data."
|
||||
)
|
||||
noTextContent: bool = Field(
|
||||
default=False, description="Whether the page has no text content."
|
||||
)
|
||||
noTextContent: bool = Field(description="Whether the page has no text content.")
|
||||
|
||||
|
||||
class JobResult(BaseModel):
|
||||
"""The raw JSON result from the LlamaParse API."""
|
||||
|
||||
pages: List[Page] = Field(description="The pages of the document.")
|
||||
job_metadata: JobMetadata = Field(description="The metadata of the job.")
|
||||
file_name: str = Field(description="The path to the file that was parsed.")
|
||||
job_id: str = Field(description="The ID of the job.")
|
||||
pages: List[Page] = Field(
|
||||
default_factory=list, description="The pages of the document."
|
||||
)
|
||||
job_metadata: JobMetadata = Field(
|
||||
default_factory=JobMetadata, description="The metadata of the job."
|
||||
)
|
||||
file_name: str = Field(
|
||||
default="", description="The path to the file that was parsed."
|
||||
)
|
||||
job_id: str = Field(default="", description="The ID of the job.")
|
||||
is_done: bool = Field(default=False, description="Whether the job is done.")
|
||||
error: Optional[str] = Field(
|
||||
default=None, description="The error message if the job failed."
|
||||
@@ -185,7 +216,9 @@ class JobResult(BaseModel):
|
||||
for page in self.pages
|
||||
]
|
||||
else:
|
||||
text = self._page_separator.join([page.text for page in self.pages])
|
||||
text = self._page_separator.join(
|
||||
[page.text if page.text is not None else "" for page in self.pages]
|
||||
)
|
||||
return [Document(text=text, metadata={"file_name": self.file_name})]
|
||||
|
||||
async def aget_text_documents(self, split_by_page: bool = False) -> List[Document]:
|
||||
@@ -230,7 +263,9 @@ class JobResult(BaseModel):
|
||||
else:
|
||||
return [
|
||||
Document(
|
||||
text=self._page_separator.join([page.md for page in self.pages]),
|
||||
text=self._page_separator.join(
|
||||
[page.md if page.md is not None else "" for page in self.pages]
|
||||
),
|
||||
metadata={"file_name": self.file_name},
|
||||
)
|
||||
]
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import httpx
|
||||
import itertools
|
||||
import logging
|
||||
from enum import Enum
|
||||
from tenacity import (
|
||||
@@ -8,7 +9,7 @@ from tenacity import (
|
||||
retry_if_exception,
|
||||
before_sleep_log,
|
||||
)
|
||||
from typing import Any
|
||||
from typing import Any, Iterable, Iterator, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -34,6 +35,17 @@ class ParsingMode(str, Enum):
|
||||
parse_page_with_lvm = "parse_page_with_lvm"
|
||||
parse_page_with_agent = "parse_page_with_agent"
|
||||
parse_document_with_llm = "parse_document_with_llm"
|
||||
parse_document_with_agent = "parse_document_with_agent"
|
||||
|
||||
|
||||
class FailedPageMode(str, Enum):
|
||||
"""
|
||||
Enum for representing the different available page error handling modes
|
||||
"""
|
||||
|
||||
raw_text = "raw_text"
|
||||
blank_page = "blank_page"
|
||||
error_message = "error_message"
|
||||
|
||||
|
||||
class Language(str, Enum):
|
||||
@@ -286,3 +298,56 @@ async def make_api_request(
|
||||
return response
|
||||
|
||||
return await _make_request(url, **httpx_kwargs)
|
||||
|
||||
|
||||
def expand_target_pages(target_pages: str) -> Iterator[int]:
|
||||
"""Yield all values in target_pages."""
|
||||
for target in target_pages.strip().split(","):
|
||||
if "-" in target:
|
||||
try:
|
||||
start, end = map(int, target.strip().split("-"))
|
||||
if start > end:
|
||||
raise ValueError
|
||||
yield from range(start, end + 1)
|
||||
except ValueError as e:
|
||||
raise ValueError(f"Invalid page range: {target}") from e
|
||||
else:
|
||||
try:
|
||||
yield int(target)
|
||||
except ValueError as e:
|
||||
raise ValueError(f"Invalid page number: {target}") from e
|
||||
|
||||
|
||||
def partition_pages(
|
||||
pages: Iterable[int], size: int, max_pages: Optional[int] = None
|
||||
) -> Iterator[str]:
|
||||
"""Yield partitioned target_pages segments."""
|
||||
if size < 1:
|
||||
raise ValueError(f"Invalid partition segment size: {size}")
|
||||
if max_pages is not None and max_pages < 1:
|
||||
raise ValueError("Max pages must be > 0")
|
||||
it = iter(pages)
|
||||
total = 0
|
||||
while max_pages is None or total < max_pages:
|
||||
segment = tuple(itertools.islice(it, size))
|
||||
if segment:
|
||||
targets = []
|
||||
for _k, g in itertools.groupby(enumerate(segment), lambda x: x[0] - x[1]):
|
||||
group = [item[1] for item in g]
|
||||
if len(group) > 1:
|
||||
start, end = group[0], group[-1]
|
||||
group_size = end - start + 1
|
||||
if max_pages is not None and total + group_size > max_pages:
|
||||
end -= total + group_size - max_pages
|
||||
group_size = end - start + 1
|
||||
if group_size > 1:
|
||||
targets.append(f"{start}-{end}")
|
||||
else:
|
||||
targets.append(str(start))
|
||||
total += group_size
|
||||
else:
|
||||
targets.append(str(group[0]))
|
||||
total += 1
|
||||
yield ",".join(targets)
|
||||
else:
|
||||
return
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
import difflib
|
||||
from pydantic import BaseModel
|
||||
from typing import Any, Dict, List, Tuple, Type
|
||||
|
||||
|
||||
def check_extra_params(
|
||||
model_cls: Type[BaseModel], data: Dict[str, Any]
|
||||
) -> Tuple[List[str], List[str]]:
|
||||
# check if one of the parameters is unused, and warn the user
|
||||
model_attributes = set(model_cls.model_fields.keys())
|
||||
extra_params = [param for param in data.keys() if param not in model_attributes]
|
||||
|
||||
suggestions: List[str] = []
|
||||
if extra_params:
|
||||
# for each unused parameter, check if it is similar to a valid parameter and suggest a typo correction, else suggest to check the documentation / update the package
|
||||
for param in extra_params:
|
||||
similar_params = difflib.get_close_matches(
|
||||
param, model_attributes, n=1, cutoff=0.8
|
||||
)
|
||||
if similar_params:
|
||||
suggestions.append(
|
||||
f"'{param}' is not a valid parameter. Did you mean '{similar_params[0]}' instead of '{param}'?"
|
||||
)
|
||||
else:
|
||||
suggestions.append(
|
||||
f"'{param}' is not a valid parameter. Please check the documentation or update the package."
|
||||
)
|
||||
|
||||
return extra_params, suggestions
|
||||
@@ -1,3 +1,8 @@
|
||||
from llama_cloud_services.parse import LlamaParse, ResultType
|
||||
from llama_cloud_services.parse import (
|
||||
LlamaParse,
|
||||
ResultType,
|
||||
ParsingMode,
|
||||
FailedPageMode,
|
||||
)
|
||||
|
||||
__all__ = ["LlamaParse", "ResultType"]
|
||||
__all__ = ["LlamaParse", "ResultType", "ParsingMode", "FailedPageMode"]
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from llama_cloud_services.parse.base import (
|
||||
LlamaParse,
|
||||
ResultType,
|
||||
ParsingMode,
|
||||
FailedPageMode,
|
||||
FileInput,
|
||||
_DEFAULT_SEPARATOR,
|
||||
JOB_RESULT_URL,
|
||||
@@ -12,6 +14,8 @@ __all__ = [
|
||||
"LlamaParse",
|
||||
"ResultType",
|
||||
"FileInput",
|
||||
"ParsingMode",
|
||||
"FailedPageMode",
|
||||
"_DEFAULT_SEPARATOR",
|
||||
"JOB_RESULT_URL",
|
||||
"JOB_STATUS_ROUTE",
|
||||
|
||||
@@ -2,10 +2,14 @@ from llama_cloud_services.parse.utils import (
|
||||
SUPPORTED_FILE_TYPES,
|
||||
Language,
|
||||
ResultType,
|
||||
ParsingMode,
|
||||
FailedPageMode,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SUPPORTED_FILE_TYPES",
|
||||
"Language",
|
||||
"ResultType",
|
||||
"ParsingMode",
|
||||
"FailedPageMode",
|
||||
]
|
||||
|
||||
Generated
+3088
File diff suppressed because it is too large
Load Diff
@@ -4,7 +4,7 @@ build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.poetry]
|
||||
name = "llama-parse"
|
||||
version = "0.6.13"
|
||||
version = "0.6.44"
|
||||
description = "Parse files into RAG-Optimized formats."
|
||||
authors = ["Logan Markewich <logan@llamaindex.ai>"]
|
||||
license = "MIT"
|
||||
@@ -13,7 +13,7 @@ packages = [{include = "llama_parse"}]
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.9,<4.0"
|
||||
llama-cloud-services = ">=0.6.13"
|
||||
llama-cloud-services = ">=0.6.44"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
pytest = "^8.0.0"
|
||||
|
||||
Generated
+849
-652
File diff suppressed because it is too large
Load Diff
+5
-4
@@ -8,7 +8,7 @@ python_version = "3.10"
|
||||
|
||||
[tool.poetry]
|
||||
name = "llama-cloud-services"
|
||||
version = "0.6.13"
|
||||
version = "0.6.44"
|
||||
description = "Tailored SDK clients for LlamaCloud services."
|
||||
authors = ["Logan Markewich <logan@runllama.ai>"]
|
||||
license = "MIT"
|
||||
@@ -17,13 +17,14 @@ packages = [{include = "llama_cloud_services"}]
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.9,<4.0"
|
||||
llama-index-core = ">=0.11.0"
|
||||
llama-cloud = "^0.1.18"
|
||||
pydantic = "!=2.10"
|
||||
llama-index-core = ">=0.12.0"
|
||||
llama-cloud = "==0.1.33"
|
||||
pydantic = ">=2.8,!=2.10"
|
||||
click = "^8.1.7"
|
||||
python-dotenv = "^1.0.1"
|
||||
eval-type-backport = {python = "<3.10", version = "^0.2.0"}
|
||||
platformdirs = "^4.3.7"
|
||||
tenacity = ">=8.5.0, <10.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
pytest = "^8.0.0"
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
import os
|
||||
import httpx
|
||||
import pytest
|
||||
import uuid
|
||||
from pydantic import BaseModel
|
||||
from dotenv import load_dotenv
|
||||
from pathlib import Path
|
||||
|
||||
from llama_cloud.client import AsyncLlamaCloud
|
||||
from llama_cloud_services.beta.agent_data import AsyncAgentDataClient
|
||||
|
||||
|
||||
class TrailingSlashHttpxClient(httpx.AsyncClient):
|
||||
"""Custom httpx client that ensures all URLs have trailing slashes"""
|
||||
|
||||
async def request(self, method, url, **kwargs):
|
||||
# Convert URL to string and ensure trailing slash
|
||||
url_str = str(url)
|
||||
if not url_str.endswith("/") and "?" not in url_str:
|
||||
url_str += "/"
|
||||
self.headers["Authorization"] = f"Bearer {LLAMA_CLOUD_API_KEY}"
|
||||
kwargs.pop("headers", None)
|
||||
return await super().request(method, url_str, headers=self.headers, **kwargs)
|
||||
|
||||
|
||||
# Load environment variables
|
||||
def load_test_dotenv():
|
||||
dotenv_path = Path(__file__).parent.parent.parent.parent / ".env.dev"
|
||||
load_dotenv(dotenv_path, override=True)
|
||||
|
||||
|
||||
load_test_dotenv()
|
||||
|
||||
# Get configuration from environment
|
||||
LLAMA_CLOUD_API_KEY = os.getenv("LLAMA_CLOUD_API_KEY")
|
||||
LLAMA_CLOUD_BASE_URL = os.getenv("LLAMA_CLOUD_BASE_URL")
|
||||
LLAMA_DEPLOY_DEPLOYMENT_NAME = os.getenv("LLAMA_DEPLOY_DEPLOYMENT_NAME")
|
||||
|
||||
|
||||
class TestData(BaseModel):
|
||||
"""Simple test data model for agent data testing"""
|
||||
|
||||
name: str
|
||||
test_id: str
|
||||
value: int
|
||||
|
||||
|
||||
# Skip all tests if API key is not set
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not LLAMA_CLOUD_API_KEY or not LLAMA_DEPLOY_DEPLOYMENT_NAME,
|
||||
reason="LLAMA_CLOUD_API_KEY or LLAMA_DEPLOY_DEPLOYMENT_NAME not set",
|
||||
)
|
||||
async def test_agent_data_crud_operations():
|
||||
"""Test basic CRUD operations for agent data with automatic cleanup"""
|
||||
# Create unique test identifier to avoid conflicts
|
||||
test_id = str(uuid.uuid4())
|
||||
|
||||
# Set up client
|
||||
client = AsyncLlamaCloud(
|
||||
token=LLAMA_CLOUD_API_KEY,
|
||||
base_url=LLAMA_CLOUD_BASE_URL,
|
||||
httpx_client=TrailingSlashHttpxClient(timeout=60, follow_redirects=True),
|
||||
)
|
||||
|
||||
# Create agent data client with unique collection name
|
||||
agent_data_client = AsyncAgentDataClient(
|
||||
client=client,
|
||||
type=TestData,
|
||||
collection_name=f"test-collection-{test_id[:8]}",
|
||||
agent_url_id=LLAMA_DEPLOY_DEPLOYMENT_NAME,
|
||||
)
|
||||
|
||||
# Create test data
|
||||
test_data = TestData(name="test-item", test_id=test_id, value=42)
|
||||
|
||||
created_item = None
|
||||
try:
|
||||
# Test CREATE
|
||||
created_item = await agent_data_client.create_agent_data(test_data)
|
||||
assert created_item.data.name == "test-item"
|
||||
assert created_item.data.test_id == test_id
|
||||
assert created_item.data.value == 42
|
||||
assert created_item.id is not None
|
||||
|
||||
# Test READ
|
||||
retrieved_item = await agent_data_client.get_agent_data(created_item.id)
|
||||
assert retrieved_item.id == created_item.id
|
||||
assert retrieved_item.data.name == "test-item"
|
||||
assert retrieved_item.data.test_id == test_id
|
||||
assert retrieved_item.data.value == 42
|
||||
|
||||
# Test SEARCH
|
||||
search_results = await agent_data_client.search_agent_data(
|
||||
filter={"test_id": {"eq": test_id}}, page_size=10, include_total=True
|
||||
)
|
||||
assert len(search_results.items) == 1
|
||||
assert search_results.items[0].data.test_id == test_id
|
||||
assert search_results.total == 1
|
||||
|
||||
# Test AGGREGATE
|
||||
aggregate_results = await agent_data_client.aggregate_agent_data(
|
||||
group_by=["test_id"], count=True
|
||||
)
|
||||
assert len(aggregate_results.items) == 1
|
||||
assert aggregate_results.items[0].group_key["test_id"] == test_id
|
||||
assert aggregate_results.items[0].count == 1
|
||||
|
||||
# Test UPDATE
|
||||
updated_data = TestData(name="updated-item", test_id=test_id, value=84)
|
||||
updated_item = await agent_data_client.update_agent_data(
|
||||
created_item.id, updated_data
|
||||
)
|
||||
assert updated_item.data.name == "updated-item"
|
||||
assert updated_item.data.value == 84
|
||||
assert updated_item.id == created_item.id
|
||||
|
||||
# Verify update persisted
|
||||
verified_item = await agent_data_client.get_agent_data(created_item.id)
|
||||
assert verified_item.data.name == "updated-item"
|
||||
assert verified_item.data.value == 84
|
||||
|
||||
finally:
|
||||
# Clean up test data
|
||||
if created_item is not None:
|
||||
try:
|
||||
await agent_data_client.delete_agent_data(created_item.id)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to cleanup test data {created_item.id}: {e}")
|
||||
@@ -0,0 +1,109 @@
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict
|
||||
|
||||
import pytest
|
||||
from llama_cloud.types.agent_data import AgentData
|
||||
from llama_cloud.types.aggregate_group import AggregateGroup
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from llama_cloud_services.beta.agent_data.schema import (
|
||||
ExtractedData,
|
||||
TypedAgentData,
|
||||
TypedAggregateGroup,
|
||||
)
|
||||
|
||||
|
||||
# Test data models
|
||||
class Person(BaseModel):
|
||||
name: str
|
||||
age: int
|
||||
email: str
|
||||
|
||||
|
||||
class Company(BaseModel):
|
||||
name: str
|
||||
industry: str
|
||||
employees: int
|
||||
|
||||
|
||||
def test_typed_agent_data_from_raw():
|
||||
"""Test TypedAgentData.from_raw class method."""
|
||||
raw_data = AgentData(
|
||||
id="456",
|
||||
agent_slug="extraction-agent",
|
||||
collection="employees",
|
||||
data={"name": "Jane Smith", "age": 25, "email": "jane@company.com"},
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
typed_data = TypedAgentData.from_raw(raw_data, Person)
|
||||
|
||||
assert typed_data.id == "456"
|
||||
assert typed_data.agent_url_id == "extraction-agent"
|
||||
assert typed_data.collection == "employees"
|
||||
assert typed_data.data.name == "Jane Smith"
|
||||
assert typed_data.data.age == 25
|
||||
assert typed_data.data.email == "jane@company.com"
|
||||
|
||||
|
||||
def test_typed_agent_data_from_raw_validation_error():
|
||||
"""Test TypedAgentData.from_raw with invalid data."""
|
||||
raw_data = AgentData(
|
||||
id="789",
|
||||
agent_slug="test-agent",
|
||||
collection="people",
|
||||
data={"name": "Invalid Person", "age": "not_a_number"}, # Invalid age
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
TypedAgentData.from_raw(raw_data, Person)
|
||||
|
||||
|
||||
def test_extracted_data_create_method():
|
||||
"""Test ExtractedData.create class method."""
|
||||
person = Person(name="Created Person", age=35, email="created@example.com")
|
||||
|
||||
# Test with defaults
|
||||
extracted = ExtractedData.create(person)
|
||||
assert extracted.original_data == person
|
||||
assert extracted.data == person
|
||||
assert extracted.status == "in_review"
|
||||
assert extracted.confidence == {}
|
||||
|
||||
# Test with custom values
|
||||
extracted_custom = ExtractedData.create(
|
||||
person, status="accepted", confidence={"name": 0.99}
|
||||
)
|
||||
assert extracted_custom.status == "accepted"
|
||||
assert extracted_custom.confidence["name"] == 0.99
|
||||
|
||||
|
||||
def test_extracted_data_with_dict():
|
||||
"""Test ExtractedData with dict data instead of Pydantic model."""
|
||||
data_dict = {"name": "Dict Person", "age": 45, "email": "dict@example.com"}
|
||||
|
||||
extracted = ExtractedData[Dict[str, Any]](
|
||||
original_data=data_dict, data=data_dict, status="accepted", confidence={}
|
||||
)
|
||||
|
||||
assert extracted.original_data["name"] == "Dict Person"
|
||||
assert extracted.data["age"] == 45
|
||||
|
||||
|
||||
def test_typed_aggregate_group_from_raw():
|
||||
"""Test TypedAggregateGroup.from_raw class method."""
|
||||
raw_group = AggregateGroup(
|
||||
group_key={"industry": "Technology"},
|
||||
count=25,
|
||||
first_item={"name": "Tech Corp", "industry": "Technology", "employees": 500},
|
||||
)
|
||||
|
||||
typed_group = TypedAggregateGroup.from_raw(raw_group, Company)
|
||||
|
||||
assert typed_group.group_key["industry"] == "Technology"
|
||||
assert typed_group.count == 25
|
||||
assert typed_group.first_item.name == "Tech Corp"
|
||||
assert typed_group.first_item.employees == 500
|
||||
@@ -0,0 +1,41 @@
|
||||
import os
|
||||
from typing import List
|
||||
from llama_cloud_services.extract import LlamaExtract
|
||||
|
||||
# Global storage for agents to cleanup
|
||||
_TEST_AGENTS_TO_CLEANUP: List[str] = []
|
||||
|
||||
|
||||
def pytest_sessionfinish(session, exitstatus):
|
||||
"""Hook that runs after all tests complete - cleanup agents here"""
|
||||
print(
|
||||
f"pytest_sessionfinish hook called! Agents to cleanup: {_TEST_AGENTS_TO_CLEANUP}"
|
||||
)
|
||||
|
||||
if _TEST_AGENTS_TO_CLEANUP:
|
||||
print("Creating cleanup client...")
|
||||
# Create a fresh client just for cleanup
|
||||
cleanup_client = LlamaExtract(
|
||||
api_key=os.getenv("LLAMA_CLOUD_API_KEY"),
|
||||
base_url=os.getenv("LLAMA_CLOUD_BASE_URL"),
|
||||
project_id=os.getenv("LLAMA_CLOUD_PROJECT_ID"),
|
||||
verbose=True,
|
||||
)
|
||||
|
||||
for agent_id in _TEST_AGENTS_TO_CLEANUP:
|
||||
try:
|
||||
print(f"Deleting agent {agent_id}...")
|
||||
cleanup_client.delete_agent(agent_id)
|
||||
print(f"Cleaned up agent {agent_id}")
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to delete agent {agent_id}: {e}")
|
||||
|
||||
_TEST_AGENTS_TO_CLEANUP.clear()
|
||||
print("Agent cleanup completed")
|
||||
else:
|
||||
print("No agents to cleanup")
|
||||
|
||||
|
||||
def register_agent_for_cleanup(agent_id: str):
|
||||
"""Register an agent ID for cleanup at the end of the test session"""
|
||||
_TEST_AGENTS_TO_CLEANUP.append(agent_id)
|
||||
Binary file not shown.
@@ -5,6 +5,7 @@ from pydantic import BaseModel
|
||||
|
||||
from llama_cloud_services.extract import LlamaExtract, ExtractionAgent, SourceText
|
||||
from tests.extract.util import load_test_dotenv
|
||||
from .conftest import register_agent_for_cleanup
|
||||
|
||||
load_test_dotenv()
|
||||
|
||||
@@ -27,7 +28,7 @@ class TestSchema(BaseModel):
|
||||
|
||||
# Test data paths
|
||||
TEST_DIR = Path(__file__).parent / "data"
|
||||
TEST_PDF = TEST_DIR / "slide" / "saas_slide.pdf"
|
||||
TEST_PDF = TEST_DIR / "api_test" / "noisebridge_receipt.pdf"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -58,7 +59,7 @@ def test_schema_dict():
|
||||
|
||||
@pytest.fixture
|
||||
def test_agent(llama_extract, test_agent_name, test_schema_dict, request):
|
||||
"""Creates a test agent and cleans it up after the test"""
|
||||
"""Creates a test agent and collects it for cleanup at the end of all tests"""
|
||||
test_id = request.node.nodeid
|
||||
test_hash = hex(hash(test_id))[-8:]
|
||||
base_name = test_agent_name
|
||||
@@ -86,13 +87,11 @@ def test_agent(llama_extract, test_agent_name, test_schema_dict, request):
|
||||
print(f"Warning: Failed to cleanup existing agent: {e}")
|
||||
|
||||
agent = llama_extract.create_agent(name=name, data_schema=schema)
|
||||
yield agent
|
||||
|
||||
# Cleanup after test
|
||||
try:
|
||||
llama_extract.delete_agent(agent.id)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to delete agent {agent.id}: {e}")
|
||||
# Add agent to cleanup list via conftest helper
|
||||
register_agent_for_cleanup(agent.id)
|
||||
|
||||
yield agent
|
||||
|
||||
|
||||
class TestLlamaExtract:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import pytest
|
||||
import shutil
|
||||
from typing import Optional, cast
|
||||
from fsspec.implementations.local import LocalFileSystem
|
||||
from httpx import AsyncClient
|
||||
|
||||
@@ -20,11 +21,15 @@ def test_simple_page_text() -> None:
|
||||
assert len(result[0].text) > 0
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def markdown_parser() -> LlamaParse:
|
||||
@pytest.fixture(params=[None, 2])
|
||||
def markdown_parser(request: pytest.FixtureRequest) -> LlamaParse:
|
||||
if os.environ.get("LLAMA_CLOUD_API_KEY", "") == "":
|
||||
pytest.skip("LLAMA_CLOUD_API_KEY not set")
|
||||
return LlamaParse(result_type="markdown", ignore_errors=False)
|
||||
return LlamaParse(
|
||||
result_type="markdown",
|
||||
ignore_errors=False,
|
||||
partition_pages=cast(Optional[int], request.param),
|
||||
)
|
||||
|
||||
|
||||
def test_simple_page_markdown(markdown_parser: LlamaParse) -> None:
|
||||
@@ -35,8 +40,6 @@ def test_simple_page_markdown(markdown_parser: LlamaParse) -> None:
|
||||
|
||||
|
||||
def test_simple_page_markdown_bytes(markdown_parser: LlamaParse) -> None:
|
||||
markdown_parser = LlamaParse(result_type="markdown", ignore_errors=False)
|
||||
|
||||
filepath = "tests/test_files/attention_is_all_you_need.pdf"
|
||||
with open(filepath, "rb") as f:
|
||||
file_bytes = f.read()
|
||||
@@ -51,8 +54,6 @@ def test_simple_page_markdown_bytes(markdown_parser: LlamaParse) -> None:
|
||||
|
||||
|
||||
def test_simple_page_markdown_buffer(markdown_parser: LlamaParse) -> None:
|
||||
markdown_parser = LlamaParse(result_type="markdown", ignore_errors=False)
|
||||
|
||||
filepath = "tests/test_files/attention_is_all_you_need.pdf"
|
||||
with open(filepath, "rb") as f:
|
||||
# client must provide extra_info with file_name
|
||||
@@ -161,9 +162,12 @@ async def test_mixing_input_types() -> None:
|
||||
os.environ.get("LLAMA_CLOUD_API_KEY", "") == "",
|
||||
reason="LLAMA_CLOUD_API_KEY not set",
|
||||
)
|
||||
@pytest.mark.parametrize("partition_pages", [None, 2])
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_images() -> None:
|
||||
parser = LlamaParse(result_type="markdown", take_screenshot=True)
|
||||
async def test_download_images(partition_pages: Optional[int]) -> None:
|
||||
parser = LlamaParse(
|
||||
result_type="markdown", take_screenshot=True, partition_pages=partition_pages
|
||||
)
|
||||
filepath = "tests/test_files/attention_is_all_you_need.pdf"
|
||||
json_result = await parser.aget_json([filepath])
|
||||
|
||||
@@ -175,3 +179,17 @@ async def test_download_images() -> None:
|
||||
|
||||
await parser.aget_images(json_result, download_path)
|
||||
assert len(os.listdir(download_path)) == len(json_result[0]["pages"][0]["images"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("split_by_page,expected", [(True, 4), (False, 1)])
|
||||
async def test_multiple_page_markdown(
|
||||
markdown_parser: LlamaParse,
|
||||
split_by_page: bool,
|
||||
expected: int,
|
||||
) -> None:
|
||||
markdown_parser.split_by_page = split_by_page
|
||||
filepath = "tests/test_files/TOS.pdf"
|
||||
result = await markdown_parser.aload_data(filepath)
|
||||
assert len(result) == expected
|
||||
assert all(len(doc.text) > 0 for doc in result)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import tempfile
|
||||
import os
|
||||
import pytest
|
||||
from typing import Optional
|
||||
from llama_cloud_services import LlamaParse
|
||||
from llama_cloud_services.parse.types import JobResult
|
||||
|
||||
@@ -15,16 +16,23 @@ def chart_file_path() -> str:
|
||||
return "tests/test_files/attention_is_all_you_need_chart.pdf"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def multiple_page_path() -> str:
|
||||
return "tests/test_files/TOS.pdf"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
os.environ.get("LLAMA_CLOUD_API_KEY", "") == "",
|
||||
reason="LLAMA_CLOUD_API_KEY not set",
|
||||
)
|
||||
async def test_basic_parse_result(file_path: str):
|
||||
@pytest.mark.parametrize("partition_pages", [None, 2])
|
||||
async def test_basic_parse_result(file_path: str, partition_pages: Optional[int]):
|
||||
parser = LlamaParse(
|
||||
take_screenshot=True,
|
||||
auto_mode=True,
|
||||
fast_mode=False,
|
||||
partition_pages=partition_pages,
|
||||
)
|
||||
result = await parser.aparse(file_path)
|
||||
|
||||
@@ -96,10 +104,12 @@ async def test_link_parse_result(file_path: str):
|
||||
os.environ.get("LLAMA_CLOUD_API_KEY", "") == "",
|
||||
reason="LLAMA_CLOUD_API_KEY not set",
|
||||
)
|
||||
@pytest.mark.skip(reason="TODO: Needs to be fixed in prod. Raising 500 error.")
|
||||
async def test_parse_structured_output(file_path: str):
|
||||
parser = LlamaParse(
|
||||
structured_output=True,
|
||||
structured_output_json_schema_name="imFeelingLucky",
|
||||
invalidate_cache=True,
|
||||
)
|
||||
result = await parser.aparse(file_path)
|
||||
assert isinstance(result, JobResult)
|
||||
@@ -142,8 +152,11 @@ async def test_parse_layout(file_path: str):
|
||||
os.environ.get("LLAMA_CLOUD_API_KEY", "") == "",
|
||||
reason="LLAMA_CLOUD_API_KEY not set",
|
||||
)
|
||||
def test_parse_multiple_files(file_path: str, chart_file_path: str):
|
||||
parser = LlamaParse()
|
||||
@pytest.mark.parametrize("partition_pages", [None, 2])
|
||||
def test_parse_multiple_files(
|
||||
file_path: str, chart_file_path: str, partition_pages: Optional[int]
|
||||
):
|
||||
parser = LlamaParse(partition_pages=partition_pages)
|
||||
result = parser.parse([file_path, chart_file_path])
|
||||
|
||||
assert isinstance(result, list)
|
||||
@@ -152,3 +165,40 @@ def test_parse_multiple_files(file_path: str, chart_file_path: str):
|
||||
assert isinstance(result[1], JobResult)
|
||||
assert result[0].file_name == file_path
|
||||
assert result[1].file_name == chart_file_path
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
os.environ.get("LLAMA_CLOUD_API_KEY", "") == "",
|
||||
reason="LLAMA_CLOUD_API_KEY not set",
|
||||
)
|
||||
@pytest.mark.parametrize("partition_pages", [None, 2])
|
||||
async def test_multiple_page_parse_result(
|
||||
multiple_page_path: str, partition_pages: Optional[int]
|
||||
):
|
||||
parser = LlamaParse(
|
||||
take_screenshot=True,
|
||||
auto_mode=True,
|
||||
fast_mode=False,
|
||||
partition_pages=partition_pages,
|
||||
)
|
||||
results = await parser.aparse(multiple_page_path)
|
||||
if partition_pages is None:
|
||||
assert isinstance(results, JobResult)
|
||||
results = [results]
|
||||
else:
|
||||
assert isinstance(results, list)
|
||||
|
||||
for result in results:
|
||||
assert isinstance(result, JobResult)
|
||||
assert result.job_id is not None
|
||||
assert result.file_name == multiple_page_path
|
||||
assert len(result.pages) > 0
|
||||
|
||||
assert result.pages[0].text is not None
|
||||
assert len(result.pages[0].text) > 0
|
||||
|
||||
assert result.pages[0].md is not None
|
||||
assert len(result.pages[0].md) > 0
|
||||
|
||||
assert result.pages[0].md != result.pages[0].text
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
import pytest
|
||||
|
||||
|
||||
from llama_cloud_services.parse.utils import expand_target_pages, partition_pages
|
||||
|
||||
|
||||
def test_expand_target_pages() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
list(expand_target_pages("x"))
|
||||
with pytest.raises(ValueError):
|
||||
list(expand_target_pages("1-2-3"))
|
||||
with pytest.raises(ValueError):
|
||||
list(expand_target_pages("2-1"))
|
||||
result = list(expand_target_pages("0,2-3,5,8-10"))
|
||||
assert result == [0, 2, 3, 5, 8, 9, 10]
|
||||
|
||||
|
||||
def test_partion_pages() -> None:
|
||||
pages = [0, 2, 3, 5, 8, 9, 10]
|
||||
with pytest.raises(ValueError):
|
||||
list(partition_pages(pages, 0))
|
||||
result = list(partition_pages(pages, 3))
|
||||
assert result == ["0,2-3", "5,8-9", "10"]
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
list(partition_pages(pages, 3, 0))
|
||||
result = list(partition_pages(pages, 3, max_pages=5))
|
||||
assert result == ["0,2-3", "5,8"]
|
||||
result = list(partition_pages(pages, 3, max_pages=10))
|
||||
assert result == ["0,2-3", "5,8-9", "10"]
|
||||
@@ -7,7 +7,8 @@ from llama_cloud_services.report import LlamaReport, ReportClient
|
||||
|
||||
# Skip tests if no API key is set
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not os.getenv("LLAMA_CLOUD_API_KEY"), reason="No API key provided"
|
||||
not os.getenv("LLAMA_CLOUD_API_KEY") or os.getenv("CI") == "true",
|
||||
reason="No API key provided",
|
||||
)
|
||||
|
||||
|
||||
|
||||
Binary file not shown.
@@ -0,0 +1,66 @@
|
||||
from pydantic import BaseModel
|
||||
from llama_cloud_services.utils import check_extra_params
|
||||
|
||||
|
||||
class MyModel(BaseModel):
|
||||
name: str
|
||||
age: int
|
||||
email: str
|
||||
is_active: bool
|
||||
|
||||
|
||||
def test_check_extra_params_no_extra():
|
||||
"""Test when all parameters are valid - should return empty lists."""
|
||||
data = {"name": "John", "age": 25, "email": "john@example.com", "is_active": True}
|
||||
|
||||
extra_params, suggestions = check_extra_params(MyModel, data)
|
||||
|
||||
assert extra_params == []
|
||||
assert suggestions == []
|
||||
|
||||
|
||||
def test_check_extra_params_with_typos():
|
||||
"""Test when there are extra parameters that are close to valid ones (typos)."""
|
||||
data = {
|
||||
"name": "John",
|
||||
"age": 25,
|
||||
"emial": "john@example.com", # typo: emial instead of email
|
||||
"is_activ": True, # typo: is_activ instead of is_active
|
||||
"address": "123 Main St", # completely different parameter
|
||||
}
|
||||
|
||||
extra_params, suggestions = check_extra_params(MyModel, data)
|
||||
|
||||
assert len(extra_params) == 3
|
||||
assert "emial" in extra_params
|
||||
assert "is_activ" in extra_params
|
||||
assert "address" in extra_params
|
||||
|
||||
# Check that typo suggestions are provided
|
||||
assert len(suggestions) == 3
|
||||
assert "Did you mean 'email' instead of 'emial'?" in suggestions[0]
|
||||
assert "Did you mean 'is_active' instead of 'is_activ'?" in suggestions[1]
|
||||
assert "check the documentation or update the package" in suggestions[2]
|
||||
|
||||
|
||||
def test_check_extra_params_completely_invalid():
|
||||
"""Test when there are extra parameters with no close matches."""
|
||||
data = {
|
||||
"name": "John",
|
||||
"xyz": "invalid",
|
||||
"random_field": 123,
|
||||
"completely_different": True,
|
||||
}
|
||||
|
||||
extra_params, suggestions = check_extra_params(MyModel, data)
|
||||
|
||||
assert len(extra_params) == 3
|
||||
assert "xyz" in extra_params
|
||||
assert "random_field" in extra_params
|
||||
assert "completely_different" in extra_params
|
||||
|
||||
# All suggestions should be generic (no close matches)
|
||||
assert len(suggestions) == 3
|
||||
for suggestion in suggestions:
|
||||
assert "check the documentation or update the package" in suggestion
|
||||
assert "Did you mean" not in suggestion
|
||||
Reference in New Issue
Block a user