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

import requests

Expand All @@ -33,6 +33,26 @@ 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)


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,6 +78,8 @@ 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
Expand All @@ -80,16 +102,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 +128,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 +160,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 +192,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 +221,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 +241,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 Down Expand Up @@ -256,7 +282,9 @@ 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"]
logger.info(
"ApiClient Successfully registered your run on the API.\n\n"
Expand All @@ -282,7 +310,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 +326,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
40 changes: 40 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
126 changes: 126 additions & 0 deletions tests/test_api_client_session.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
"""
Connection-reuse and timeout tests for ApiClient.

These run against a stdlib HTTP server on loopback rather than requests_mock,
because requests_mock replaces the transport adapter and therefore never opens
a real connection, which is exactly what is under test here. No traffic leaves
the machine.
"""

import threading
import unittest
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

from codecarbon.core.api_client import ApiClient

CONF = {
"os": "linux",
"python_version": "3.12",
"codecarbon_version": "3.0",
"cpu_count": 8,
"cpu_model": "CPU",
"gpu_count": 0,
"gpu_model": "",
"longitude": 0.0,
"latitude": 0.0,
"region": "EU",
"provider": "none",
"ram_total_size": 16.0,
"tracking_mode": "machine",
}

EMISSION = {
"duration": 5,
"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,
}


class _Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1" # keep-alive, so pooling is observable

def log_message(self, *args):
pass

def _serve(self):
self.server.state["requests"] += 1
length = int(self.headers.get("Content-Length", 0) or 0)
if length:
self.rfile.read(length)
body = b'{"id": "run-1"}'
self.send_response(201)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)

do_GET = _serve
do_POST = _serve


class _Server(ThreadingHTTPServer):
daemon_threads = True
allow_reuse_address = True

def __init__(self, state):
self.state = state
super().__init__(("127.0.0.1", 0), _Handler)

def process_request(self, request, client_address):
self.state["connections"] += 1
super().process_request(request, client_address)


class TestSessionReuse(unittest.TestCase):
def setUp(self):
self.state = {"requests": 0, "connections": 0}
self.server = _Server(self.state)
threading.Thread(target=self.server.serve_forever, daemon=True).start()
self.url = f"http://127.0.0.1:{self.server.server_address[1]}"
self.addCleanup(self.server.server_close)
self.addCleanup(self.server.shutdown)
self.api = ApiClient(
endpoint_url=self.url,
experiment_id="exp-1",
conf=CONF,
create_run_automatically=False,
)
self.addCleanup(self.api.close)
self.api.run_id = "run-1"

def test_sequential_calls_reuse_one_connection(self):
for _ in range(50):
self.assertTrue(self.api.add_emission(dict(EMISSION)))

self.assertEqual(self.state["requests"], 50)
self.assertEqual(self.state["connections"], 1)

def test_close_is_idempotent(self):
self.api.add_emission(dict(EMISSION))
self.api.close()
self.api.close()


class TestTimeout(unittest.TestCase):
def test_requests_get_a_connect_and_read_timeout(self):
api = ApiClient(endpoint_url="http://test.com", create_run_automatically=False)
self.addCleanup(api.close)
seen = {}

def fake_get(url, json, timeout, headers):
seen["timeout"] = timeout
return type("R", (), {"status_code": 200, "json": lambda self: {}})()

api._request(fake_get, "http://test.com/x")
self.assertEqual(seen["timeout"], (3.05, 10))


if __name__ == "__main__":
unittest.main()
Loading