From 7b566a8f1fc46c29f715d784f6eb5a710a7b069c Mon Sep 17 00:00:00 2001 From: David Berenstein Date: Thu, 20 Aug 2026 08:12:57 +0200 Subject: [PATCH] feat(api): reuse connections, add timeouts, and stamp emissions with measurement time `ApiClient` called the module-level `requests` functions, so every call opened a new TCP connection and TLS handshake. A tracker sending one measurement per tick opened 50 connections for 50 uploads; with a `Session` it opens 1. The flat 2s timeout is replaced with `(3.05, 10)`: the old value timed out against a healthy but loaded API, while a hung endpoint still cannot block the scheduler thread for long. `ApiClient.add_emission` also discarded `carbon_emission["timestamp"]` and called `get_datetime_with_timezone()` instead, so every row stored the moment the payload was built rather than the moment it was measured. Harmless while the two are milliseconds apart, wrong by the full latency as soon as a send is slow, retried or queued. The measurement timestamp is now normalised to offset-aware, falling back to now when the payload carries none or an unparseable one (the method is public and takes a plain dict). CSV output is untouched. Co-Authored-By: Claude Opus 5 (1M context) --- codecarbon/core/api_client.py | 62 +++++++++++---- codecarbon/output_methods/http.py | 3 + tests/test_api_call.py | 40 ++++++++++ tests/test_api_client_session.py | 126 ++++++++++++++++++++++++++++++ 4 files changed, 214 insertions(+), 17 deletions(-) create mode 100644 tests/test_api_client_session.py diff --git a/codecarbon/core/api_client.py b/codecarbon/core/api_client.py index bc2e0974e..450cdabbc 100644 --- a/codecarbon/core/api_client.py +++ b/codecarbon/core/api_client.py @@ -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 @@ -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 @@ -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 @@ -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: @@ -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): """ @@ -134,7 +160,7 @@ 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): @@ -142,7 +168,7 @@ 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): """ @@ -150,14 +176,14 @@ def update_organization(self, organization: OrganizationCreate): """ 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): """ @@ -166,7 +192,7 @@ 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): @@ -174,7 +200,7 @@ 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 @@ -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"], @@ -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. @@ -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" @@ -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): """ @@ -298,7 +326,7 @@ 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): @@ -306,7 +334,7 @@ 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): """ diff --git a/codecarbon/output_methods/http.py b/codecarbon/output_methods/http.py index e0ff710b1..b7e8a5896 100644 --- a/codecarbon/output_methods/http.py +++ b/codecarbon/output_methods/http.py @@ -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) diff --git a/tests/test_api_call.py b/tests/test_api_call.py index 31e25c039..481c76111 100644 --- a/tests/test_api_call.py +++ b/tests/test_api_call.py @@ -1,5 +1,6 @@ import dataclasses import unittest +from datetime import datetime from uuid import uuid4 import requests @@ -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) diff --git a/tests/test_api_client_session.py b/tests/test_api_client_session.py new file mode 100644 index 000000000..085881785 --- /dev/null +++ b/tests/test_api_client_session.py @@ -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()