Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 73 additions & 16 deletions bcb/odata/framework.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,31 @@
_METADATA_CACHE: dict[str, "ODataMetadata"] = {}
_METADATA_CACHE_LOCK = threading.RLock()


def _load_json_object(text: str, *, context: str) -> dict[str, Any]:
try:
data = json.loads(text)
except json.JSONDecodeError as ex:
raise ODataError(f"{context} returned invalid JSON: {ex}") from ex
if not isinstance(data, dict):
raise ODataError(f"{context} returned invalid JSON payload: expected object")
return data


def _required_field(data: dict[str, Any], field: str, *, context: str) -> Any:
try:
return data[field]
except KeyError as ex:
raise ODataError(f"{context} response missing required field {field!r}") from ex


def _load_xml_document(content: bytes, *, context: str) -> Any:
try:
return etree.parse(BytesIO(content))
except etree.XMLSyntaxError as ex:
raise ODataError(f"{context} returned invalid XML: {ex}") from ex


# Edm.Boolean
# Edm.Byte
# Edm.Date
Expand Down Expand Up @@ -278,14 +303,24 @@ class ODataMetadata:
def __init__(self, url: str) -> None:
self.url = url
self._load_document()
_xpath = "edmx:DataServices/edm:Schema"
schema = self.doc.xpath(_xpath, namespaces=self.namespaces)[0]
self.namespace: str = schema.attrib["Namespace"]
self._used_elements: list[str] = []
self._parse_entities(schema)
self._parse_entity_sets(schema)
self._parse_functions(schema)
self._parse_function_imports(schema)
try:
_xpath = "edmx:DataServices/edm:Schema"
schemas = self.doc.xpath(_xpath, namespaces=self.namespaces)
if not schemas:
raise ODataError(f"OData metadata {self.url} missing schema")
schema = schemas[0]
self.namespace = schema.attrib["Namespace"]
self._used_elements: list[str] = []
self._parse_entities(schema)
self._parse_entity_sets(schema)
self._parse_functions(schema)
self._parse_function_imports(schema)
except ODataError:
raise
except (KeyError, IndexError, TypeError) as ex:
raise ODataError(
f"OData metadata {self.url} has invalid structure: {ex}"
) from ex

def _load_document(self) -> None:
logger.debug(f"Fetching OData metadata from {self.url}")
Expand All @@ -306,7 +341,7 @@ def _load_document(self) -> None:
rate_limit_cls=ODataError,
server_error_cls=ODataError,
)
self.doc = etree.parse(BytesIO(res.content))
self.doc = _load_xml_document(res.content, context=f"OData metadata {self.url}")

def _parse_entity(self, entity_element: Any, namespace: str) -> ODataEntity:
name = entity_element.attrib["Name"]
Expand Down Expand Up @@ -411,11 +446,27 @@ def __init__(self, url: str) -> None:
rate_limit_cls=ODataError,
server_error_cls=ODataError,
)
self.api_data: dict[str, Any] = json.loads(res.text)
self.endpoints: list[ODataEndPoint] = [
ODataEndPoint(**x) for x in self.api_data["value"]
]
self._odata_context_url: str = self.api_data["@odata.context"]
context = f"OData service {self.url}"
self.api_data = _load_json_object(res.text, context=context)
value = _required_field(self.api_data, "value", context=context)
if not isinstance(value, list):
raise ODataError("OData service response field 'value' must be a list")
endpoints = []
for endpoint in value:
if not isinstance(endpoint, dict):
raise ODataError(
"OData service response field 'value' must contain objects"
)
endpoints.append(ODataEndPoint(**endpoint))
self.endpoints = endpoints
odata_context = _required_field(
self.api_data, "@odata.context", context=context
)
if not isinstance(odata_context, str):
raise ODataError(
"OData service response field '@odata.context' must be a string"
)
self._odata_context_url = odata_context

# Use cached metadata if available, otherwise create and cache new one
with _METADATA_CACHE_LOCK:
Expand Down Expand Up @@ -560,7 +611,10 @@ def reset(self) -> None:
self._params = {}

def collect(self) -> Any:
return json.loads(self.text())
url = self.odata_url()
data = _load_json_object(self.text(), context=f"OData query {url}")
_required_field(data, "value", context=f"OData query {url}")
return data

async def async_text(self) -> str:
"""Async version of text(). Fetches OData response using shared client."""
Expand Down Expand Up @@ -592,7 +646,10 @@ async def async_text(self) -> str:

async def async_collect(self) -> Any:
"""Async version of collect(). Awaits async_text() and parses JSON."""
return json.loads(await self.async_text())
url = self.odata_url()
data = _load_json_object(await self.async_text(), context=f"OData query {url}")
_required_field(data, "value", context=f"OData query {url}")
return data

def text(self) -> str:
params = self._build_parameters()
Expand Down
24 changes: 24 additions & 0 deletions tests/test_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,30 @@ async def test_odata_query_async_status_error_raises(httpx_mock):
await ep.query().limit(1).async_text()


async def test_odata_query_async_malformed_json_raises(httpx_mock):
httpx_mock.add_response(
url="https://olinda.bcb.gov.br/olinda/servico/Expectativas/versao/v1/odata/",
text=ODATA_SERVICE_ROOT_JSON,
status_code=200,
)
httpx_mock.add_response(
url="https://olinda.bcb.gov.br/olinda/servico/Expectativas/versao/v1/odata/$metadata",
content=ODATA_METADATA_XML,
status_code=200,
)
httpx_mock.add_response(
url=re.compile(r".*ExpectativasMercadoAnuais.*"),
text="not json",
status_code=200,
)

api = Expectativas()
ep = api.get_endpoint("ExpectativasMercadoAnuais")

with pytest.raises(ODataError, match="OData query.*invalid JSON"):
await ep.query().limit(1).async_collect()


async def test_odata_query_async_collect(httpx_mock):
"""Test ODataQuery.async_collect() returns DataFrame."""
httpx_mock.add_response(
Expand Down
84 changes: 84 additions & 0 deletions tests/test_odata.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,28 @@ def test_service_root_connection_error_raises_odata_error(httpx_mock):
Expectativas()


def test_service_root_malformed_json_raises_odata_error(httpx_mock):
httpx_mock.add_response(
url=EXPECTATIVAS_BASE_URL,
text="not json",
status_code=200,
)

with pytest.raises(ODataError, match="OData service.*invalid JSON"):
Expectativas()


def test_service_root_missing_required_fields_raises_odata_error(httpx_mock):
httpx_mock.add_response(
url=EXPECTATIVAS_BASE_URL,
text="{}",
status_code=200,
)

with pytest.raises(ODataError, match="missing required field 'value'"):
Expectativas()


def test_metadata_status_error_raises_odata_error(httpx_mock):
httpx_mock.add_response(
url=EXPECTATIVAS_BASE_URL,
Expand All @@ -102,6 +124,38 @@ def test_metadata_status_error_raises_odata_error(httpx_mock):
Expectativas()


def test_metadata_malformed_xml_raises_odata_error(httpx_mock):
httpx_mock.add_response(
url=EXPECTATIVAS_BASE_URL,
text=ODATA_SERVICE_ROOT_JSON,
status_code=200,
)
httpx_mock.add_response(
url=EXPECTATIVAS_METADATA_URL,
text="<edmx:Edmx><bad",
status_code=200,
)

with pytest.raises(ODataError, match="OData metadata.*invalid XML"):
Expectativas()


def test_metadata_missing_schema_raises_odata_error(httpx_mock):
httpx_mock.add_response(
url=EXPECTATIVAS_BASE_URL,
text=ODATA_SERVICE_ROOT_JSON,
status_code=200,
)
httpx_mock.add_response(
url=EXPECTATIVAS_METADATA_URL,
text='<?xml version="1.0"?><root />',
status_code=200,
)

with pytest.raises(ODataError, match="OData metadata.*missing schema"):
Expectativas()


def test_query_status_error_raises_odata_error(httpx_mock):
add_service_mocks(httpx_mock)
httpx_mock.add_response(
Expand All @@ -117,6 +171,36 @@ def test_query_status_error_raises_odata_error(httpx_mock):
ep.query().limit(1).collect()


def test_query_malformed_json_raises_odata_error(httpx_mock):
add_service_mocks(httpx_mock)
httpx_mock.add_response(
url=ENTITY_URL_PATTERN,
text="not json",
status_code=200,
)

api = Expectativas()
ep = api.get_endpoint("ExpectativasMercadoAnuais")

with pytest.raises(ODataError, match="OData query.*invalid JSON"):
ep.query().limit(1).collect()


def test_query_missing_value_raises_odata_error(httpx_mock):
add_service_mocks(httpx_mock)
httpx_mock.add_response(
url=ENTITY_URL_PATTERN,
text='{"unexpected": []}',
status_code=200,
)

api = Expectativas()
ep = api.get_endpoint("ExpectativasMercadoAnuais")

with pytest.raises(ODataError, match="missing required field 'value'"):
ep.query().limit(1).collect()


# ---------------------------------------------------------------------------
# ODataProperty operator overloading
# ---------------------------------------------------------------------------
Expand Down
Loading