Skip to content

Commit dfe2268

Browse files
committed
Switch from requests to aiohttp
I would like us to eventually move to Tornado (from fastapi) to be consistent with what we use in the JupyterHub ecosystem. One part of that is to make all network stuff async, and this moves us from `requests` to `aiohttp`, making everything async as needed as we go. This isn't complete, a WIP.
1 parent 1fb3ef2 commit dfe2268

6 files changed

Lines changed: 80 additions & 84 deletions

File tree

pyproject.toml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,3 +31,6 @@ escapism = { git = "https://github.com/jupyterhub/escapism", tag = "1.0.1" }
3131
[pytest]
3232
log_cli = true
3333
log_cli_level = "INFO"
34+
35+
[tool.pytest.ini_options]
36+
asyncio_mode = "auto"

src/jupyterhub_cost_monitoring/app.py

Lines changed: 8 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from datetime import timedelta
22

3+
import aiohttp
34
import requests
45
from fastapi import FastAPI, HTTPException, Query
56
from fastapi.responses import Response
@@ -27,6 +28,7 @@
2728
app = FastAPI()
2829
app.add_middleware(MetricsMiddleware)
2930
logger = get_logger(__name__)
31+
client_session = aiohttp.ClientSession()
3032

3133

3234
@app.get("/")
@@ -256,7 +258,7 @@ def total_costs_per_group(
256258

257259

258260
@app.get("/costs-per-user")
259-
def costs_per_user(
261+
async def costs_per_user(
260262
from_date: str | None = Query(
261263
None, alias="from", description="Start date in YYYY-MM-DDTHH:MMZ format"
262264
),
@@ -315,22 +317,16 @@ def costs_per_user(
315317
# Get per-user costs by combining AWS costs with Prometheus usage data
316318
results = []
317319
for ug in usergroup:
318-
try:
319-
per_user_costs = query_total_costs_per_user(
320-
date_range, hub, component, user, ug, limit
321-
)
322-
except requests.exceptions.HTTPError as e:
323-
response = e.response
324-
raise HTTPException(status_code=response.status_code, detail=response.text)
325-
except Exception as e:
326-
raise HTTPException(status_code=500, detail=f"{e}")
320+
per_user_costs = await query_total_costs_per_user(
321+
client_session, date_range, hub, component, user, ug, limit
322+
)
327323
results.extend(per_user_costs)
328324

329325
return results
330326

331327

332328
@app.get("/total-usage")
333-
def total_usage(
329+
async def total_usage(
334330
from_date: str | None = Query(
335331
None, alias="from", description="Start date in YYYY-MM-DDTHH:MMZ format"
336332
),
@@ -359,7 +355,7 @@ def total_usage(
359355
user = None
360356

361357
try:
362-
return query_usage(date_range, hub, component, user)
358+
return await query_usage(client_session, date_range, hub, component, user)
363359
except requests.exceptions.HTTPError as e:
364360
response = e.response
365361
raise HTTPException(status_code=response.status_code, detail=response.text)

src/jupyterhub_cost_monitoring/query_cost_aws.py

Lines changed: 14 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,8 @@
66
import functools
77
from pprint import pformat
88

9+
import aiohttp
910
import boto3
10-
import requests
1111

1212
from .cache import ttl_lru_cache
1313
from .const_cost_aws import (
@@ -514,7 +514,8 @@ def query_total_costs_per_component(
514514

515515

516516
@ttl_lru_cache(seconds_to_live=3600)
517-
def query_total_costs_per_user(
517+
async def query_total_costs_per_user(
518+
client: aiohttp.ClientSession,
518519
date_range: DateRange,
519520
hub: str = None,
520521
component: str = None,
@@ -556,15 +557,13 @@ def query_total_costs_per_user(
556557
# Get user usage percentages from Prometheus using the same DateRange object
557558
# This ensures we query the same logical date range for both AWS and Prometheus,
558559
# accounting for their different date range semantics (exclusive vs inclusive)
559-
try:
560-
usage_shares = query_usage(
561-
date_range,
562-
hub_name=hub,
563-
component_name=component,
564-
user_name=user,
565-
)
566-
except requests.exceptions.ConnectionError:
567-
raise
560+
usage_shares = await query_usage(
561+
client,
562+
date_range,
563+
hub_name=hub,
564+
component_name=component,
565+
user_name=user,
566+
)
568567
results = []
569568
for entry in usage_shares:
570569
d = entry["date"]
@@ -577,7 +576,7 @@ def query_total_costs_per_user(
577576
) # Adjust usage share to cost
578577
results.append(entry)
579578
results = [x for x in results if x["hub"] != "binder"] # Exclude binder hubs
580-
user_groups = query_user_groups(date_range, hub, user)
579+
user_groups = await query_user_groups(client, date_range, hub, user)
581580
seen = set()
582581
list_groups = []
583582
# Ensure uniquely keyed entries when double-counting group costs
@@ -630,7 +629,8 @@ def query_total_costs_per_user(
630629

631630

632631
@ttl_lru_cache(seconds_to_live=3600)
633-
def query_total_costs_per_group(
632+
async def query_total_costs_per_group(
633+
client: aiohttp.ClientSession,
634634
date_range: DateRange,
635635
):
636636
"""
@@ -642,11 +642,7 @@ def query_total_costs_per_group(
642642
Returns:
643643
List of dicts with keys: date, usergroup and cost.
644644
"""
645-
try:
646-
results = query_total_costs_per_user(date_range=date_range)
647-
except Exception as e:
648-
logger.exception(f"HTTP request failed: {e}")
649-
raise
645+
results = await query_total_costs_per_user(client, date_range=date_range)
650646
response = {}
651647
for r in results:
652648
key = (r["date"], r["usergroup"])

src/jupyterhub_cost_monitoring/query_usage.py

Lines changed: 31 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,8 @@
66
from collections import defaultdict
77
from datetime import datetime, timedelta, timezone
88

9+
import aiohttp
910
import escapism
10-
import requests
1111
from yarl import URL
1212

1313
from .cache import ttl_lru_cache
@@ -23,7 +23,9 @@
2323
prometheus_password = os.environ.get("PROMETHEUS_PASSWORD", "")
2424

2525

26-
def query_prometheus(query: str, date_range: DateRange, step: str) -> requests.Response:
26+
async def query_prometheus(
27+
client: aiohttp.ClientSession, query: str, date_range: DateRange, step: str
28+
) -> dict:
2729
"""
2830
Query the Prometheus server with the given query over a date range.
2931
@@ -42,9 +44,7 @@ def query_prometheus(query: str, date_range: DateRange, step: str) -> requests.R
4244
scheme="http", host=prometheus_host, port=prometheus_port
4345
)
4446
if prometheus_username != "" and prometheus_password != "":
45-
prometheus_auth = requests.auth.HTTPBasicAuth(
46-
prometheus_username, prometheus_password
47-
)
47+
prometheus_auth = aiohttp.BasicAuth(prometheus_username, prometheus_password)
4848
else:
4949
prometheus_auth = None
5050
parameters = {
@@ -53,15 +53,16 @@ def query_prometheus(query: str, date_range: DateRange, step: str) -> requests.R
5353
"end": to_date,
5454
"step": step,
5555
}
56-
query_api = URL(prometheus_api.with_path("/api/v1/query_range"))
57-
with requests.get(query_api, params=parameters, auth=prometheus_auth) as response:
56+
query_api = prometheus_api.with_path("/api/v1/query_range").with_query(parameters)
57+
async with client.get(query_api, auth=prometheus_auth) as response:
5858
logger.info(f"Querying Prometheus: {response.url}")
5959
response.raise_for_status()
60-
result = response.json()
60+
result = await response.json()
6161
return result
6262

6363

64-
def query_usage(
64+
async def query_usage(
65+
client: aiohttp.ClientSession,
6566
date_range: DateRange,
6667
hub_name: str | None,
6768
component_name: str | None,
@@ -84,23 +85,18 @@ def query_usage(
8485
if component_name is None:
8586
# Query all components defined in USAGE_MAP
8687
for component, params in USAGE_MAP.items():
87-
try:
88-
response = query_prometheus(
89-
params["query"], date_range, step=params["step"]
90-
)
91-
except requests.exceptions.RequestException:
92-
raise
88+
response = await query_prometheus(
89+
client, params["query"], date_range, step=params["step"]
90+
)
9391
result.extend(_process_response(response, component))
9492
else:
9593
# Query specific component only
96-
try:
97-
response = query_prometheus(
98-
USAGE_MAP[component_name]["query"],
99-
date_range,
100-
step=USAGE_MAP[component_name]["step"],
101-
)
102-
except requests.exceptions.RequestException:
103-
raise
94+
response = await query_prometheus(
95+
client,
96+
USAGE_MAP[component_name]["query"],
97+
date_range,
98+
step=USAGE_MAP[component_name]["step"],
99+
)
104100
result.extend(_process_response(response, component_name))
105101
# Calculate daily cost factors from absolute usage totals)
106102
result = _calculate_daily_cost_factors(result, hub_name=hub_name)
@@ -111,9 +107,9 @@ def query_usage(
111107

112108

113109
def _process_response(
114-
response: requests.Response,
110+
response: dict,
115111
component_name: str,
116-
) -> dict:
112+
) -> list[dict]:
117113
"""
118114
Process the response from the Prometheus server to extract absolute usage data.
119115
@@ -256,7 +252,8 @@ def _calculate_daily_cost_factors(
256252

257253

258254
@ttl_lru_cache(seconds_to_live=3600)
259-
def query_user_groups(
255+
async def query_user_groups(
256+
client: aiohttp.ClientSession,
260257
hub_name: str | None = None,
261258
user_name: str | None = None,
262259
group_name: str | None = None,
@@ -266,17 +263,13 @@ def query_user_groups(
266263
"""
267264
now_date = get_now_date() - timedelta(days=1)
268265
date_range = DateRange(start_date=now_date, end_date=now_date)
269-
try:
270-
response = query_prometheus(USER_GROUP_INFO, date_range, step="1d")
271-
except requests.exceptions.RequestException as e:
272-
logger.exception(f"HTTP request failed: {e}")
273-
raise
266+
response = await query_prometheus(client, USER_GROUP_INFO, date_range, step="1d")
274267
result = _process_user_groups(response, hub_name, user_name, group_name)
275268
return result
276269

277270

278271
def _process_user_groups(
279-
response: requests.Response,
272+
response: dict,
280273
hub_name: str | None = None,
281274
user_name: str | None = None,
282275
group_name: str | None = None,
@@ -306,16 +299,13 @@ def _process_user_groups(
306299

307300

308301
@ttl_lru_cache(seconds_to_live=3600)
309-
def query_users_with_multiple_groups(
302+
async def query_users_with_multiple_groups(
303+
client: aiohttp.ClientSession,
310304
date_range: DateRange,
311305
hub_name: str | None = None,
312306
user_name: str | None = None,
313307
) -> list[dict]:
314-
try:
315-
response = query_user_groups(hub_name=hub_name, user_name=user_name)
316-
except requests.exceptions.RequestException as e:
317-
logger.exception(f"HTTP request failed: {e}")
318-
raise
308+
response = await query_user_groups(client, hub_name=hub_name, user_name=user_name)
319309
grouped = defaultdict(
320310
lambda: {"username": None, "hub": None, "usergroups": [], "has_multiple": False}
321311
)
@@ -340,16 +330,13 @@ def query_users_with_multiple_groups(
340330

341331

342332
@ttl_lru_cache(seconds_to_live=3600)
343-
def query_users_with_no_groups(
333+
async def query_users_with_no_groups(
334+
client: aiohttp.ClientSession,
344335
date_range: DateRange,
345336
hub_name: str | None = None,
346337
user_name: str | None = None,
347338
) -> list[dict]:
348-
try:
349-
response = query_user_groups(hub_name=hub_name, user_name=user_name)
350-
except requests.exceptions.RequestException as e:
351-
logger.exception(f"HTTP request failed: {e}")
352-
raise
339+
response = await query_user_groups(client, hub_name=hub_name, user_name=user_name)
353340
grouped = defaultdict(lambda: {"username": None, "hub": None})
354341
for entry in response:
355342
key = (entry["username"], entry["hub"])

tests/conftest.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from datetime import datetime, timezone
44
from unittest.mock import patch
55

6+
import aiohttp
67
import boto3
78
import pytest
89
from botocore.stub import Stubber
@@ -14,6 +15,12 @@
1415
# Usage and cost data fixtures for test_cost.py
1516

1617

18+
@pytest.fixture(scope="session")
19+
async def client_session():
20+
async with aiohttp.ClientSession() as client:
21+
yield client
22+
23+
1724
@pytest.fixture(scope="function")
1825
def input_data_usage():
1926
with open("tests/data/test_data_usage.json") as f:

0 commit comments

Comments
 (0)