diff --git a/skyrl/tinker/extra/skyrl_train_inference_forwarding.py b/skyrl/tinker/extra/skyrl_train_inference_forwarding.py index 321fdd212e..eb908f4314 100644 --- a/skyrl/tinker/extra/skyrl_train_inference_forwarding.py +++ b/skyrl/tinker/extra/skyrl_train_inference_forwarding.py @@ -36,7 +36,7 @@ def __init__(self, engine_config: EngineConfig, db_engine): connect=10.0, read=engine_config.forwarding_inference_timeout_sec, write=300.0, - pool=300.0, + pool=engine_config.forwarding_inference_timeout_sec, ), limits=httpx.Limits( max_connections=max_conn, diff --git a/tests/tinker/test_inference_forwarding_config.py b/tests/tinker/test_inference_forwarding_config.py index 091a49efb3..add902df80 100644 --- a/tests/tinker/test_inference_forwarding_config.py +++ b/tests/tinker/test_inference_forwarding_config.py @@ -21,7 +21,7 @@ def test_forwarding_timeout_reads_environment(monkeypatch) -> None: assert config.forwarding_inference_timeout_sec == 1800.0 -def test_forwarding_client_uses_configured_timeout() -> None: +def test_forwarding_client_uses_configured_read_and_pool_timeout() -> None: config = EngineConfig( base_model="test-model", forwarding_inference_timeout_sec=1800.0, @@ -34,7 +34,7 @@ def test_forwarding_client_uses_configured_timeout() -> None: assert timeout.connect == 10.0 assert timeout.read == 1800.0 assert timeout.write == 300.0 - assert timeout.pool == 300.0 + assert timeout.pool == 1800.0 @pytest.mark.asyncio