Compare commits

...

1 Commits

Author SHA1 Message Date
Clelia (Astra) Bertelli e110334273 feat: add mocking for agent data (#1028)
* wip: add mocking for agent data

* chore: add tests; fix: various fixes

* chore: fix schema generation null handling
2025-11-25 21:59:17 +01:00
4 changed files with 563 additions and 2 deletions
@@ -178,6 +178,9 @@ def _generate_value(schema: Any, rng: random.Random, depth: int) -> Any:
rng.randint(0, 1_000_000), sentences=max(1, length // 5)
)
if schema_type == "null":
return None
if "oneOf" in schema:
option = rng.choice(schema["oneOf"])
return _generate_value(option, rng, depth + 1)
@@ -0,0 +1,335 @@
from __future__ import annotations
import httpx
import re
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Dict
from ._deterministic import utcnow, hash_schema
if TYPE_CHECKING:
from .server import FakeLlamaCloudServer
@dataclass
class StoredAgentData:
data: dict[str, Any]
id: str
collection: str
deployment_name: str
def __getattr__(self, name: str) -> Any:
return self.data.get(name)
def __setattr__(self, name: str, value: Any) -> None:
if name in ("data", "id", "collection", "deployment_name"):
super().__setattr__(name, value)
else:
self.data[name] = value
@classmethod
def from_request_data(cls, data: dict[str, Any]) -> "StoredAgentData":
return cls(
data=data.get("data", {}),
collection=data.get("collection", "default"),
deployment_name=data.get("deployment_name", ""),
id=hash_schema(data.get("data", {}))[:7],
)
def apply_filter(data: dict, filters: dict) -> bool:
"""Check if data matches all filters"""
ops = {
"gt": lambda a, b: a > b,
"gte": lambda a, b: a >= b,
"lt": lambda a, b: a < b,
"lte": lambda a, b: a <= b,
"eq": lambda a, b: a == b,
"ne": lambda a, b: a != b,
"in": lambda a, b: a in b,
"nin": lambda a, b: a not in b,
}
for key, condition in filters.items():
if key not in data:
return False
if isinstance(condition, dict):
for op, value in condition.items():
if op in ops:
if not ops[op](data[key], value):
return False
else:
return False
else:
if data[key] != condition:
return False
return True
class FakeAgentDataNamespace:
def __init__(
self,
*,
server: "FakeLlamaCloudServer",
) -> None:
self._server = server
self.stored: list[StoredAgentData] = []
self.routes: Dict[str, Any] = {}
def _create_data(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request=request)
data = StoredAgentData.from_request_data(payload)
self.stored.append(data)
response = {
"data": data.data,
"collection": data.collection,
"deployment_name": data.deployment_name,
"created_at": utcnow().isoformat(),
"updated_at": None,
"id": data.id,
"project_id": None,
"organization_id": None,
}
return self._server.json_response(response, status_code=200)
def _delete_data_by_query(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request=request)
delete_count = 0
if (filters := payload.get("filter")) is not None:
to_keep = []
for data in self.stored:
if data.collection == payload.get(
"collection", "default"
) and data.deployment_name == payload.get("deployment_name"):
if not apply_filter(data.data, filters):
to_keep.append(data)
else:
delete_count += 1
self.stored = to_keep
return self._server.json_response(
{"deleted_count": delete_count}, status_code=200
)
def _delete_data_by_id(self, request: httpx.Request) -> httpx.Response:
item_id = self._find_item_id(request=request)
if not item_id:
return self._server.json_response(
{
"detail": "An item_id path parameter is required to perform this operation"
},
status_code=400,
)
if not any(data.id == item_id for data in self.stored):
return self._server.json_response(
{"detail": f"No data with ID: {item_id}"}, status_code=404
)
self.stored = [data for data in self.stored if data.id != item_id]
return self._server.json_response({}, status_code=200)
def _get_data_by_id(self, request: httpx.Request) -> httpx.Response:
item_id = self._find_item_id(request=request)
if not item_id:
return self._server.json_response(
{
"detail": "An item_id path parameter is required to perform this operation"
},
status_code=400,
)
data = [data for data in self.stored if data.id == item_id]
if data:
response = {
"data": data[0].data,
"collection": data[0].collection,
"deployment_name": data[0].deployment_name,
"created_at": utcnow().isoformat(),
"updated_at": None,
"id": data[0].id,
"project_id": None,
"organization_id": None,
}
return self._server.json_response(response, status_code=200)
else:
return self._server.json_response(
{"detail": f"No data with ID: {item_id}"}, status_code=404
)
def _search_data(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request=request)
found = []
if (filters := payload.get("filter")) is not None:
for data in self.stored:
if data.collection == payload.get(
"collection", "default"
) and data.deployment_name == payload.get("deployment_name"):
if apply_filter(data.data, filters):
found.append(
{
"data": data.data,
"collection": data.collection,
"deployment_name": data.deployment_name,
"created_at": utcnow().isoformat(),
"updated_at": None,
"id": data.id,
"project_id": None,
"organization_id": None,
}
)
else:
for data in self.stored:
if data.collection == payload.get(
"collection", "default"
) and data.deployment_name == payload.get("deployment_name"):
found.append(
{
"data": data.data,
"collection": data.collection,
"deployment_name": data.deployment_name,
"created_at": utcnow().isoformat(),
"updated_at": None,
"id": data.id,
"project_id": None,
"organization_id": None,
}
)
return self._server.json_response(
{"items": found, "next_page_token": None, "total_size": len(found)},
status_code=200,
)
def _update_data(self, request: httpx.Request) -> httpx.Response:
item_id = self._find_item_id(request=request)
payload = self._server.json(request=request)
if not item_id:
return self._server.json_response(
{
"detail": "An item_id path parameter is required to perform this operation"
},
status_code=400,
)
updated = None
for i, data in enumerate(self.stored):
if data.id == item_id:
updated = data
updated.data = payload.get("data", data.data)
self.stored[i] = updated
print(updated)
if updated is not None:
response = {
"data": updated.data,
"collection": updated.collection,
"deployment_name": updated.deployment_name,
"created_at": None,
"updated_at": utcnow().isoformat(),
"id": updated.id,
"project_id": None,
"organization_id": None,
}
status_code = 200
else:
response = {"detail": f"Record with id {item_id} not found"}
status_code = 404
return self._server.json_response(response, status_code=status_code)
def _aggregate_data(self, request: httpx.Request) -> httpx.Response:
payload = self._server.json(request=request)
add_count = payload.get("count", False)
group_bys: list[str] = payload.get("group_by", [])
groups: dict[str, dict[str, list[dict]]] = {key: {} for key in group_bys}
if (filters := payload.get("filter")) is not None:
for data in self.stored:
if data.collection == payload.get(
"collection", "default"
) and data.deployment_name == payload.get("deployment_name"):
if apply_filter(data.data, filters):
for key in group_bys:
if key in data.data and data.data[key] in groups[key]:
groups[key][data.data[key]].append(data.data)
elif key in data.data and data.data[key] not in groups[key]:
groups[key][data.data[key]] = [data.data]
else:
for data in self.stored:
if data.collection == payload.get(
"collection", "default"
) and data.deployment_name == payload.get("deployment_name"):
for key in group_bys:
if key in data.data and data.data[key] in groups[key]:
groups[key][data.data[key]].append(data.data)
elif key in data.data and data.data[key] not in groups[key]:
groups[key][data.data[key]] = [data.data]
response: dict[str, Any] = {
"items": [],
"next_page_token": None,
"total_size": 0,
}
for k in groups:
if len(groups[k]) > 0:
for v in groups[k]:
if groups[k][v]:
first_element = groups[k][v][0]
else:
first_element = None
response["items"].append(
{
"first_item": first_element,
"count": len(groups[k][v]) if add_count else None,
"group_key": {k: v},
}
)
response["total_size"] = len(response["items"])
return self._server.json_response(response, status_code=200)
def _find_item_id(self, request: httpx.Request) -> str | None:
matchgroups = re.search(r"/agent-data\/([^\/]+)$", request.url.path)
return matchgroups.group(1) if matchgroups is not None else None
def register(self) -> None:
server = self._server
route = server.add_route(
"POST",
"/api/v1/beta/agent-data",
self._create_data,
namespace="create_item",
)
self.routes["stateless_run"] = route
self.stateless_run = route
server.add_route(
"POST",
"/api/v1/beta/agent-data/:aggregate",
self._aggregate_data,
namespace="untyped_aggregate",
alias="aggregate",
)
server.add_route(
"POST",
"/api/v1/beta/agent-data/:delete",
self._delete_data_by_query,
namespace="delete",
)
server.add_route(
"POST",
"/api/v1/beta/agent-data/:search",
self._search_data,
namespace="untyped_search",
alias="search",
)
server.add_route(
"DELETE",
"/api/v1/beta/agent-data/{item_id}",
self._delete_data_by_id,
namespace="delete_item",
)
server.add_route(
"GET",
"/api/v1/beta/agent-data/{item_id}",
self._get_data_by_id,
namespace="untyped_get_item",
alias="get_item",
)
server.add_route(
"PUT",
"/api/v1/beta/agent-data/{item_id}",
self._update_data,
namespace="update_item",
)
@@ -12,7 +12,7 @@ from .classify import FakeClassifyNamespace
from .extract import FakeExtractNamespace
from .files import FakeFilesNamespace
from .parse import FakeParseNamespace
from .agent_data import FakeAgentDataNamespace
Handler = Callable[[httpx.Request], httpx.Response]
@@ -33,7 +33,7 @@ class FakeLlamaCloudServer:
default_organization_id: str = "org-test",
) -> None:
self.base_urls = tuple(base_urls or (self.DEFAULT_BASE_URL,))
selected = namespaces or ("files", "extract", "parse", "classify")
selected = namespaces or ("files", "extract", "parse", "classify", "agent_data")
self._namespace_names = {name.lower() for name in selected}
self._upload_base_url = upload_base_url or self.DEFAULT_UPLOAD_BASE
self._download_base_url = download_base_url or self.DEFAULT_DOWNLOAD_BASE
@@ -51,6 +51,7 @@ class FakeLlamaCloudServer:
self.extract = FakeExtractNamespace(server=self, files=self.files)
self.parse = FakeParseNamespace(server=self)
self.classify = FakeClassifyNamespace(server=self, files=self.files)
self.agent_data = FakeAgentDataNamespace(server=self)
# Context management ----------------------------------------------
def install(self) -> "FakeLlamaCloudServer":
@@ -164,6 +165,8 @@ class FakeLlamaCloudServer:
self.parse.register()
if "classify" in self._namespace_names:
self.classify.register()
if "agent_data" in self._namespace_names:
self.agent_data.register()
self._registered = True
@@ -5,9 +5,12 @@ from pathlib import Path
import pytest
from llama_cloud import ExtractConfig
from llama_cloud.types import ExtractMode
from llama_cloud.core.api_error import ApiError
from llama_cloud_services.extract import LlamaExtract
from llama_cloud_services.parse import LlamaParse
from llama_cloud_services.beta.agent_data import AsyncAgentDataClient
from llama_cloud_services.testing_utils import FakeLlamaCloudServer
from llama_cloud_services.testing_utils._deterministic import hash_schema
from pydantic import BaseModel, Field
@@ -89,3 +92,220 @@ def test_parse_load_data_returns_documents(
assert documents
assert "(page 1)" in documents[0].text
@pytest.mark.asyncio
async def test_agent_data_create(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data = Receipt(merchant="Test Inc", total=1000)
item = await client.create_item(data)
assert item.id == hash_schema(data)[:7]
assert item.data.merchant == data.merchant and item.data.total == data.total
assert item.collection == "extracted_data"
assert item.deployment_name == "extraction_agent"
@pytest.mark.asyncio
async def test_agent_data_update(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data = Receipt(merchant="Test Inc", total=1000)
item = await client.create_item(data)
assert item.id is not None
updated_data = Receipt(merchant="Testing Inc", total=1100)
updated_item = await client.update_item(item_id=item.id, data=updated_data)
# ensure that the data actually changed
assert (
updated_item.data.merchant == updated_data.merchant
and updated_item.data.total == updated_data.total
)
# make sure nothing else changed
assert updated_item.id == item.id
assert updated_item.collection == item.collection
assert updated_item.deployment_name == item.deployment_name
@pytest.mark.asyncio
async def test_agent_data_search(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
item1 = await client.create_item(data1)
item2 = await client.create_item(data2)
item3 = await client.create_item(data3)
result = await client.search(filter={"merchant": {"eq": "Test Inc"}})
assert result.total == 2
assert any(item.id == item1.id for item in result.items) and any(
item.id == item2.id for item in result.items
)
assert all(item.data.merchant == "Test Inc" for item in result.items)
result1 = await client.search(filter={"total": {"lt": 1200}})
assert result.total == 2
assert any(item.id == item1.id for item in result1.items) and any(
item.id == item3.id for item in result1.items
)
assert all(item.data.total < 1200 for item in result1.items)
@pytest.mark.asyncio
async def test_agent_data_aggregate(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
await client.create_item(data1)
await client.create_item(data2)
await client.create_item(data3)
result = await client.aggregate(
filter={"merchant": {"eq": "Test Inc"}},
group_by=["merchant"],
count=True,
)
# filtering for 'Test Inc' on merchant means that only data with 'Test Inc' are left, meaning that there is only one group of data for merchant, i.e. the 'Test Inc' group
assert len(result.items) == 1
assert result.items[0].count == 2
assert result.items[0].first_item is not None
assert result.items[0].first_item.merchant == data1.merchant
assert result.items[0].first_item.total == data1.total
assert result.items[0].group_key == {"merchant": "Test Inc"}
result = await client.aggregate(
group_by=["merchant"],
count=True,
)
assert len(result.items) == 2
assert result.items[0].count == 2
assert result.items[0].first_item is not None
assert result.items[0].first_item.merchant == data1.merchant
assert result.items[0].first_item.total == data1.total
assert result.items[0].group_key == {"merchant": "Test Inc"}
assert result.items[1].count == 1
assert result.items[1].first_item is not None
assert result.items[1].first_item.merchant == data3.merchant
assert result.items[1].first_item.total == data3.total
assert result.items[1].group_key == {"merchant": "Testing Inc"}
@pytest.mark.asyncio
async def test_agent_data_get(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
item1 = await client.create_item(data1)
assert item1.id is not None
item2 = await client.create_item(data2)
assert item2.id is not None
item = await client.get_item(item1.id)
assert item.collection == item1.collection
assert item.deployment_name == item1.deployment_name
assert item.data.merchant == data1.merchant
assert item.data.total == data1.total
# using this pattern instead of `with pytest.raise` for more granual control over the error itself
try:
notitem = await client.get_item(item2.id + "thisdoesnotexist")
e = None
except ApiError as err:
e = err
notitem = None
assert notitem is None
assert e is not None
assert e.status_code == 404
assert e.body == {"detail": f"No data with ID: {item2.id+'thisdoesnotexist'}"}
@pytest.mark.asyncio
async def test_agent_data_delete_by_id(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data = Receipt(merchant="Test Inc", total=1300)
item = await client.create_item(data)
assert item.id is not None
await client.delete_item(item.id)
# using this pattern instead of `with pytest.raise` for more granual control over the error itself
try:
notitem = await client.get_item(item.id)
e = None
except ApiError as err:
e = err
notitem = None
assert notitem is None
assert e is not None
assert e.status_code == 404
assert e.body == {"detail": f"No data with ID: {item.id}"}
# using this pattern instead of `with pytest.raise` for more granual control over the error itself
try:
await client.delete_item(item.id)
e = None
except ApiError as err:
e = err
assert e is not None
assert e.status_code == 404
assert e.body == {"detail": f"No data with ID: {item.id}"}
@pytest.mark.asyncio
async def test_agent_data_delete_by_query(fake_server: FakeLlamaCloudServer):
with fake_server as _:
client = AsyncAgentDataClient(
Receipt,
collection="extracted_data",
deployment_name="extraction_agent",
token="fake-api-key",
)
data1 = Receipt(merchant="Test Inc", total=1000)
data2 = Receipt(merchant="Test Inc", total=1300)
data3 = Receipt(merchant="Testing Inc", total=1100)
item1 = await client.create_item(data1)
item2 = await client.create_item(data2)
item3 = await client.create_item(data3)
result = await client.delete(filter={"merchant": {"eq": "Test Inc"}})
assert result == 2
for item in (item1, item2):
assert item.id is not None
try:
notitem = await client.get_item(item.id)
e = None
except ApiError as err:
e = err
notitem = None
assert notitem is None
assert e is not None
assert e.status_code == 404
assert e.body == {"detail": f"No data with ID: {item.id}"}
assert item3.id is not None
itemfound = await client.get_item(item3.id)
assert itemfound.id == item3.id