Skip to content
Open
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
76 changes: 59 additions & 17 deletions codecarbon/core/api_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,8 @@
# from httpx import AsyncClient
import dataclasses
import json
from datetime import timedelta, tzinfo
import time
from datetime import datetime, timedelta, tzinfo

import requests

Expand All @@ -33,6 +34,29 @@ def get_datetime_with_timezone():
return str(arrow.now().isoformat())


# (connect, read) seconds, replacing a flat 2s that timed out on a loaded API.
_TIMEOUT = (3.05, 10)
# Seconds to wait after a failed run creation before trying again, so a down
# API costs one blocking call per minute instead of one per measurement.
_RUN_CREATE_COOLDOWN = 60


def _measurement_timestamp(carbon_emission: dict) -> str:
"""
Offset-aware ISO timestamp of *when the measurement was taken*, taken from
EmissionsData.timestamp. Falls back to now for hand-built payloads that
carry no usable timestamp.
"""
try:
return (
datetime.fromisoformat(carbon_emission["timestamp"])
.astimezone()
.isoformat()
)
except (KeyError, TypeError, ValueError):
return get_datetime_with_timezone()


class ApiClient: # (AsyncClient)
"""
This class call the Code Carbon API
Expand All @@ -58,11 +82,14 @@ def __init__(
:create_run_automatically: If False, do not create a run. To use API in read only mode.
"""
# super().__init__(base_url=endpoint_url) # (AsyncClient)
# A Session so the socket and TLS handshake are reused across calls.
self._session = requests.Session()
self.url = endpoint_url
self.experiment_id = experiment_id
self.api_key = api_key
self.conf = conf
self.access_token = access_token
self._run_create_failed_at = None
if self.experiment_id is not None and create_run_automatically:
self._create_run(self.experiment_id)

Expand All @@ -80,16 +107,20 @@ def _request(self, method, url, payload=None, expected_status=200):
Call the API and return the response, raising on anything that is not
the status code the API answers on success.

:method: the requests function to call, for example requests.get
:method: the session function to call, for example self._session.get
:payload: the JSON body to send, if any
:expected_status: the http code the API returns when the call succeeds
"""
headers = self._get_headers()
response = method(url=url, json=payload, timeout=2, headers=headers)
response = method(url=url, json=payload, timeout=_TIMEOUT, headers=headers)
if response.status_code != expected_status:
self._raise_api_error(url, payload or {}, response)
return response

def close(self):
"""Release the pooled sockets. Safe to call more than once."""
self._session.close()

def set_access_token(self, token: str):
"""This method sets the access token to be used for the API.
Args:
Expand All @@ -102,14 +133,14 @@ def check_auth(self):
Check API access to user account
"""
url = self.url + "/auth/check"
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def get_list_organizations(self):
"""
List all organizations
"""
url = self.url + "/organizations"
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def check_organization_exists(self, organization_name: str):
"""
Expand All @@ -134,30 +165,30 @@ def create_organization(self, organization: OrganizationCreate):
return organization
else:
return self._request(
requests.post, url, payload=payload, expected_status=201
self._session.post, url, payload=payload, expected_status=201
).json()

def get_organization(self, organization_id):
"""
Get an organization
"""
url = self.url + "/organizations/" + organization_id
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def update_organization(self, organization: OrganizationCreate):
"""
Update an organization
"""
payload = dataclasses.asdict(organization)
url = self.url + "/organizations/" + organization.id
return self._request(requests.patch, url, payload=payload).json()
return self._request(self._session.patch, url, payload=payload).json()

def list_projects_from_organization(self, organization_id):
"""
List all projects
"""
url = self.url + "/organizations/" + organization_id + "/projects"
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def create_project(self, project: ProjectCreate):
"""
Expand All @@ -166,15 +197,15 @@ def create_project(self, project: ProjectCreate):
payload = dataclasses.asdict(project)
url = self.url + "/projects"
return self._request(
requests.post, url, payload=payload, expected_status=201
self._session.post, url, payload=payload, expected_status=201
).json()

def get_project(self, project_id):
"""
Get a project
"""
url = self.url + "/projects/" + project_id
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def add_emission(self, carbon_emission: dict):
assert self.experiment_id is not None
Expand All @@ -195,7 +226,7 @@ def add_emission(self, carbon_emission: dict):
)
return False
emission = EmissionCreate(
timestamp=get_datetime_with_timezone(),
timestamp=_measurement_timestamp(carbon_emission),
run_id=self.run_id,
duration=int(carbon_emission["duration"]),
emissions_sum=carbon_emission["emissions"],
Expand All @@ -215,7 +246,7 @@ def add_emission(self, carbon_emission: dict):
try:
payload = dataclasses.asdict(emission)
url = self.url + "/emissions"
self._request(requests.post, url, payload=payload, expected_status=201)
self._request(self._session.post, url, payload=payload, expected_status=201)
logger.debug(f"ApiClient - Successful upload emission {payload} to {url}")
except requests.exceptions.HTTPError:
# Already logged by _raise_api_error, do not log it twice.
Expand All @@ -235,6 +266,14 @@ def _create_run(self, experiment_id: str):
"ApiClient FATAL The ApiClient._create_run() needs an experiment_id !"
)
return None
if (
self._run_create_failed_at is not None
and time.monotonic() - self._run_create_failed_at < _RUN_CREATE_COOLDOWN
):
logger.debug("ApiClient - run creation failed recently, not retrying yet")
return None
# Cleared on success below; set now so every failure path is covered.
self._run_create_failed_at = time.monotonic()
try:
run = RunCreate(
timestamp=get_datetime_with_timezone(),
Expand All @@ -256,8 +295,11 @@ def _create_run(self, experiment_id: str):
)
payload = dataclasses.asdict(run)
url = self.url + "/runs"
r = self._request(requests.post, url, payload=payload, expected_status=201)
r = self._request(
self._session.post, url, payload=payload, expected_status=201
)
self.run_id = r.json()["id"]
self._run_create_failed_at = None
logger.info(
"ApiClient Successfully registered your run on the API.\n\n"
+ f"Run ID: {self.run_id}\n"
Expand All @@ -282,7 +324,7 @@ def list_experiments_from_project(self, project_id: str):
List all experiments for a project
"""
url = self.url + "/projects/" + project_id + "/experiments"
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def set_experiment(self, experiment_id: str):
"""
Expand All @@ -298,15 +340,15 @@ def add_experiment(self, experiment: ExperimentCreate):
payload = dataclasses.asdict(experiment)
url = self.url + "/experiments"
return self._request(
requests.post, url, payload=payload, expected_status=201
self._session.post, url, payload=payload, expected_status=201
).json()

def get_experiment(self, experiment_id):
"""
Get an experiment by id
"""
url = self.url + "/experiments/" + experiment_id
return self._request(requests.get, url).json()
return self._request(self._session.get, url).json()

def _raise_api_error(self, url, payload, response):
"""
Expand Down
3 changes: 3 additions & 0 deletions codecarbon/output_methods/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,9 @@ def __init__(
)
self.run_id = self.api.run_id

def exit(self) -> None:
self.api.close()

def _ensure_api_run(self) -> None:
if self.api.run_id is None and self.api.experiment_id is not None:
self.api._create_run(self.api.experiment_id)
Expand Down
60 changes: 60 additions & 0 deletions tests/test_api_call.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import dataclasses
import unittest
from datetime import datetime
from uuid import uuid4

import requests
Expand Down Expand Up @@ -261,6 +262,45 @@ def test_add_emission_skips_short_duration(self):
)
)

def test_add_emission_keeps_measurement_timestamp(self):
"""The row must carry when it was measured, not when it was sent."""
payload = {
"duration": 10,
"emissions": 1.0,
"emissions_rate": 1.0,
"cpu_power": 1.0,
"gpu_power": 0.0,
"ram_power": 0.5,
"cpu_energy": 0.1,
"gpu_energy": 0.0,
"ram_energy": 0.1,
"energy_consumed": 0.2,
}
with requests_mock.Mocker() as m:
m.post("http://test.com/emissions", status_code=201)
api = ApiClient(
endpoint_url="http://test.com",
experiment_id="exp-1",
conf=conf,
create_run_automatically=False,
)
api.run_id = "run-1"

# naive timestamp, as produced by EmissionsData
assert api.add_emission({**payload, "timestamp": "2020-01-01T00:00:00"})
sent = datetime.fromisoformat(m.last_request.json()["timestamp"])
self.assertEqual(
sent.replace(tzinfo=None).isoformat(), "2020-01-01T00:00:00"
)
self.assertIsNotNone(sent.tzinfo)

# missing / unparseable timestamps fall back to now
for bad in ({}, {"timestamp": None}, {"timestamp": "222"}):
assert api.add_emission({**payload, **bad})
sent = datetime.fromisoformat(m.last_request.json()["timestamp"])
self.assertIsNotNone(sent.tzinfo)
self.assertGreater(sent.year, 2020)

def test_add_emission_raises_on_unsuccessful_post(self):
with requests_mock.Mocker() as m:
m.post("http://test.com/emissions", text="bad", status_code=500)
Expand Down Expand Up @@ -303,6 +343,26 @@ def test_create_run_raises_on_unsuccessful_status(self):
api._create_run("experiment_id")
self.assertIsNone(api.run_id)

def test_failed_create_run_is_not_retried_during_cooldown(self):
with requests_mock.Mocker() as m:
runs = m.post("http://test.com/runs", text="down", status_code=503)
api = ApiClient(
endpoint_url="http://test.com",
experiment_id="experiment_id",
api_key="Toto",
conf=conf,
create_run_automatically=False,
)
with self.assertRaises(requests.exceptions.HTTPError):
api._create_run("experiment_id")
self.assertIsNone(api._create_run("experiment_id"))
self.assertEqual(runs.call_count, 1)

api._run_create_failed_at -= 3600 # cooldown elapsed
runs = m.post("http://test.com/runs", json={"id": "run-1"}, status_code=201)
self.assertEqual(api._create_run("experiment_id"), "run-1")
self.assertIsNone(api._run_create_failed_at)

def test_create_run_raises_on_unexpected_2xx_status(self):
with requests_mock.Mocker() as m:
m.post("http://test.com/runs", json={}, status_code=200)
Expand Down
Loading
Loading