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()