Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
Next Next commit
test cases and more robust err handling
  • Loading branch information
UmanShahzad committed Dec 21, 2020
commit 4347bfddd5134502cc5182f7fb305141e175e740
12 changes: 4 additions & 8 deletions ipinfo/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,10 +179,9 @@ def getBatchDetails(
timeout_total is not None
and time.time() - start_time > timeout_total
):
if raise_on_fail:
raise TimeoutExceededError()
else:
return result
return handler_utils.return_or_fail(
raise_on_fail, TimeoutExceededError(), result
)

chunk = lookup_addresses[i : i + batch_size]

Expand All @@ -197,10 +196,7 @@ def getBatchDetails(
raise RequestQuotaExceededError()
response.raise_for_status()
except Exception as e:
if raise_on_fail:
raise e
else:
return result
return handler_utils.return_or_fail(raise_on_fail, e, result)

# fill cache
json_response = response.json()
Expand Down
110 changes: 71 additions & 39 deletions ipinfo/handler_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,13 @@
import json
import os
import sys
import time

import aiohttp

from .cache.default import DefaultCache
from .details import Details
from .exceptions import RequestQuotaExceededError
from .exceptions import RequestQuotaExceededError, TimeoutExceededError
from .handler_utils import (
API_URL,
COUNTRY_FILE_DEFAULT,
Expand Down Expand Up @@ -197,49 +198,80 @@ async def getBatchDetails(
url = API_URL + "/batch"
headers = handler_utils.get_headers(self.access_token)
headers["content-type"] = "application/json"
reqs = []
for i in range(0, len(lookup_addresses), batch_size):
chunk = lookup_addresses[i : i + batch_size]

# do http req
reqs.append(
self.httpsess.post(
url,
data=json.dumps(chunk),
headers=headers,
timeout=timeout_per_batch,
)

# prepare coroutines that will make reqs and update results.
reqs = [
self._do_batch_req(
lookup_addresses[i : i + batch_size],
url,
headers,
timeout_per_batch,
raise_on_fail,
result,
)
for i in range(0, len(lookup_addresses), batch_size)
]

try:
_, pending = await asyncio.wait(
{*reqs},
timeout=timeout_total,
return_when=asyncio.FIRST_EXCEPTION,
)

resps = await asyncio.wait_for(
asyncio.gather(*reqs, return_exceptions=raise_on_fail),
timeout_total
)
for resp in resps:
# gather data
try:
if resp.status == 429:
raise RequestQuotaExceededError()
resp.raise_for_status()
except Exception as e:
if raise_on_fail:
raise e
else:
return result

json_resp = await resp.json()

# format & fill up cache
for ip_address, details in json_resp.items():
if isinstance(details, dict):
handler_utils.format_details(details, self.countries)
self.cache[ip_address] = details

# merge cached results with new lookup
result.update(json_resp)
# if all done, return result.
if len(pending) == 0:
return result

# if some had a timeout, first cancel timed out stuff and wait for
# cleanup. then exit with return_or_fail.
for co in pending:
try:
co.cancel()
await co
except asyncio.CancelledError:
pass

return handler_utils.return_or_fail(
raise_on_fail, TimeoutExceededError(), result
)
except Exception as e:
return handler_utils.return_or_fail(raise_on_fail, e, result)

return result

async def _do_batch_req(
self, chunk, url, headers, timeout_per_batch, raise_on_fail, result
):
"""
Coroutine which will do the actual POST request for getBatchDetails.
"""
resp = await self.httpsess.post(
url,
data=json.dumps(chunk),
headers=headers,
timeout=timeout_per_batch,
)

# gather data
try:
if resp.status == 429:
raise RequestQuotaExceededError()
resp.raise_for_status()
except Exception as e:
return handler_utils.return_or_fail(raise_on_fail, e, None)

json_resp = await resp.json()

# format & fill up cache
for ip_address, details in json_resp.items():
if isinstance(details, dict):
handler_utils.format_details(details, self.countries)
self.cache[ip_address] = details

# merge cached results with new lookup
result.update(json_resp)

def _ensure_aiohttp_ready(self):
"""Ensures aiohttp internal state is initialized."""
if self.httpsess:
Expand Down
10 changes: 10 additions & 0 deletions ipinfo/handler_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,3 +83,13 @@ def read_country_names(countries_file=None):
countries_json = f.read()

return json.loads(countries_json)


def return_or_fail(raise_on_fail, e, v):
"""
Either throws `e` if `raise_on_fail` or else returns `v`.
"""
if raise_on_fail:
raise e
else:
return v
14 changes: 13 additions & 1 deletion tests/handler_async_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from ipinfo.details import Details
from ipinfo.handler_async import AsyncHandler
from ipinfo import handler_utils
import ipinfo
import pytest


Expand Down Expand Up @@ -117,10 +118,21 @@ def _check_batch_details(ips, details, token):
assert "domains" in d


@pytest.mark.parametrize("batch_size", [None, 2, 3])
@pytest.mark.parametrize("batch_size", [None, 1, 2, 3])
@pytest.mark.asyncio
async def test_get_batch_details(batch_size):
handler, token, ips = _prepare_batch_test()
details = await handler.getBatchDetails(ips, batch_size=batch_size)
_check_batch_details(ips, details, token)
await handler.deinit()


@pytest.mark.parametrize("batch_size", [None, 1, 2, 3])
@pytest.mark.asyncio
async def test_get_batch_details_total_timeout(batch_size):
handler, token, ips = _prepare_batch_test()
with pytest.raises(ipinfo.exceptions.TimeoutExceededError):
await handler.getBatchDetails(
ips, batch_size=batch_size, timeout_total=0.001
)
await handler.deinit()
12 changes: 11 additions & 1 deletion tests/handler_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from ipinfo.details import Details
from ipinfo.handler import Handler
from ipinfo import handler_utils
import ipinfo
import pytest


Expand Down Expand Up @@ -112,8 +113,17 @@ def _check_batch_details(ips, details, token):
assert "domains" in d


@pytest.mark.parametrize("batch_size", [None, 2, 3])
@pytest.mark.parametrize("batch_size", [None, 1, 2, 3])
def test_get_batch_details(batch_size):
handler, token, ips = _prepare_batch_test()
details = handler.getBatchDetails(ips, batch_size=batch_size)
_check_batch_details(ips, details, token)


@pytest.mark.parametrize("batch_size", [1, 2])
def test_get_batch_details_total_timeout(batch_size):
handler, token, ips = _prepare_batch_test()
with pytest.raises(ipinfo.exceptions.TimeoutExceededError):
handler.getBatchDetails(
ips, batch_size=batch_size, timeout_total=0.001
)