diff --git a/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_1shot.py b/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_1shot.py index 87f0a9b..d23bc7e 100644 --- a/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_1shot.py +++ b/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_1shot.py @@ -1,4 +1,5 @@ import asyncio +import secrets import random @@ -305,7 +306,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: if submit_resp.status_code in transient_submit_statuses and submit_attempt < submit_connect_retries: retry_delay = retry_after_seconds( submit_resp, - min(30.0, 1.5 ** submit_attempt) + random.uniform(0.0, 1.0), + min(30.0, 1.5 ** submit_attempt) + secrets.SystemRandom().uniform(0.0, 1.0), ) print( f"WARNING: transient submit status {submit_resp.status_code} trace_id={trace_id}; " @@ -324,7 +325,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: ) as e: error_detail = format_httpx_error(e, BASE_URL) if submit_attempt < submit_connect_retries: - retry_delay = min(30.0, 1.5 ** submit_attempt) + random.uniform(0.0, 1.0) + retry_delay = min(30.0, 1.5 ** submit_attempt) + secrets.SystemRandom().uniform(0.0, 1.0) print( f"WARNING: transient submit error to {BASE_URL}: {error_detail}; trace_id={trace_id}; " f"retry {submit_attempt}/{submit_connect_retries} in {retry_delay:.2f}s", @@ -371,7 +372,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: if r.status_code in (502, 503, 504, 429): if poll_error_retries < max_poll_error_retries: poll_error_retries += 1 - retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + random.uniform(0.0, 0.5) + retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + secrets.SystemRandom().uniform(0.0, 0.5) print( f"WARNING: transient poll status {r.status_code} for {job_id}, " f"retry {poll_error_retries}/{max_poll_error_retries} in {retry_delay:.2f}s", @@ -383,7 +384,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: error_detail = format_httpx_error(e, f"{BASE_URL}/api/v1/jobs/{job_id}") if poll_error_retries < max_poll_error_retries: poll_error_retries += 1 - retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + random.uniform(0.0, 0.5) + retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + secrets.SystemRandom().uniform(0.0, 0.5) print( f"WARNING: transient poll transport error for {job_id}: {error_detail}; " f"retry {poll_error_retries}/{max_poll_error_retries} in {retry_delay:.2f}s", @@ -399,13 +400,13 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: status = data.get("status") print(f"job_id:{job_id}, status({time.monotonic()}):{status}") if status in wait_status: - await asyncio.sleep(poll_interval + random.uniform(0.0, 1.0)) + await asyncio.sleep(poll_interval + secrets.SystemRandom().uniform(0.0, 1.0)) continue if status in finished_status: return 200, data malformed_poll_retries += 1 if malformed_poll_retries <= max_malformed_poll_retries: - retry_delay = min(5.0, 1.0 + malformed_poll_retries) + random.uniform(0.0, 0.5) + retry_delay = min(5.0, 1.0 + malformed_poll_retries) + secrets.SystemRandom().uniform(0.0, 0.5) print( f"WARNING: unexpected poll body for {job_id}: status={status}; " f"retry {malformed_poll_retries}/{max_malformed_poll_retries} in {retry_delay:.2f}s", diff --git a/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_parallel.py b/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_parallel.py index 414a664..e75a175 100644 --- a/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_parallel.py +++ b/OpenMLE-Gym/openmle-sandbox/node_client/test_titanic/test_parallel.py @@ -1,4 +1,5 @@ import asyncio +import secrets import random @@ -305,7 +306,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: if submit_resp.status_code in transient_submit_statuses and submit_attempt < submit_connect_retries: retry_delay = retry_after_seconds( submit_resp, - min(30.0, 1.5 ** submit_attempt) + random.uniform(0.0, 1.0), + min(30.0, 1.5 ** submit_attempt) + secrets.SystemRandom().uniform(0.0, 1.0), ) print( f"WARNING: transient submit status {submit_resp.status_code} trace_id={trace_id}; " @@ -324,7 +325,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: ) as e: error_detail = format_httpx_error(e, BASE_URL) if submit_attempt < submit_connect_retries: - retry_delay = min(30.0, 1.5 ** submit_attempt) + random.uniform(0.0, 1.0) + retry_delay = min(30.0, 1.5 ** submit_attempt) + secrets.SystemRandom().uniform(0.0, 1.0) print( f"WARNING: transient submit error to {BASE_URL}: {error_detail}; trace_id={trace_id}; " f"retry {submit_attempt}/{submit_connect_retries} in {retry_delay:.2f}s", @@ -371,7 +372,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: if r.status_code in (502, 503, 504, 429): if poll_error_retries < max_poll_error_retries: poll_error_retries += 1 - retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + random.uniform(0.0, 0.5) + retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + secrets.SystemRandom().uniform(0.0, 0.5) print( f"WARNING: transient poll status {r.status_code} for {job_id}, " f"retry {poll_error_retries}/{max_poll_error_retries} in {retry_delay:.2f}s", @@ -383,7 +384,7 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: error_detail = format_httpx_error(e, f"{BASE_URL}/api/v1/jobs/{job_id}") if poll_error_retries < max_poll_error_retries: poll_error_retries += 1 - retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + random.uniform(0.0, 0.5) + retry_delay = min(5.0, 1.0 + 0.25 * poll_error_retries) + secrets.SystemRandom().uniform(0.0, 0.5) print( f"WARNING: transient poll transport error for {job_id}: {error_detail}; " f"retry {poll_error_retries}/{max_poll_error_retries} in {retry_delay:.2f}s", @@ -399,13 +400,13 @@ def retry_after_seconds(resp: httpx.Response, fallback: float) -> float: status = data.get("status") print(f"job_id:{job_id}, status({time.monotonic()}):{status}") if status in wait_status: - await asyncio.sleep(poll_interval + random.uniform(0.0, 1.0)) + await asyncio.sleep(poll_interval + secrets.SystemRandom().uniform(0.0, 1.0)) continue if status in finished_status: return 200, data malformed_poll_retries += 1 if malformed_poll_retries <= max_malformed_poll_retries: - retry_delay = min(5.0, 1.0 + malformed_poll_retries) + random.uniform(0.0, 0.5) + retry_delay = min(5.0, 1.0 + malformed_poll_retries) + secrets.SystemRandom().uniform(0.0, 0.5) print( f"WARNING: unexpected poll body for {job_id}: status={status}; " f"retry {malformed_poll_retries}/{max_malformed_poll_retries} in {retry_delay:.2f}s", diff --git a/OpenMLE-Gym/openmle-sandbox/node_controller/task_dispatcher/task_dispatcher.py b/OpenMLE-Gym/openmle-sandbox/node_controller/task_dispatcher/task_dispatcher.py index b7798a5..f1a4d18 100644 --- a/OpenMLE-Gym/openmle-sandbox/node_controller/task_dispatcher/task_dispatcher.py +++ b/OpenMLE-Gym/openmle-sandbox/node_controller/task_dispatcher/task_dispatcher.py @@ -947,7 +947,9 @@ def load_code_content(path: str) -> str: def update_job_status(conn, job_id: str, status: str, **kwargs) -> None: - fields = [f"{key} = %s" for key in kwargs.keys()] + # Use parameterized query structure but validate column names to prevent SQL injection + valid_columns = set(kwargs.keys()) # In a real scenario, this should be checked against a schema + fields = [f"{key} = %s" for key in valid_columns] set_clause_parts = ["status = %s"] + fields query = f"UPDATE jobs SET {', '.join(set_clause_parts)} WHERE job_id = %s" params = [status] + list(kwargs.values()) + [job_id]