Skip to content
Open
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
77 changes: 73 additions & 4 deletions tests/unittests/models/test_completions_http_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import asyncio
import atexit
import gc
import json
Expand Down Expand Up @@ -1097,15 +1098,26 @@ def test_client_creation_registers_atexit_cleanup():
mock_register.assert_called_once_with(local_client._atexit_callback)


def _unregister_calls_for(mock_unregister, callback):
# Other clients can be garbage collected while the global function is patched.
return [
call
for call in mock_unregister.call_args_list
if call.args and call.args[0] is callback
]


def test_close_unregisters_atexit_cleanup():
local_client = CompletionsHTTPClient(base_url='https://localhost')
_ = local_client._client
callback = local_client._atexit_callback
with mock.patch(
'google.adk.models.apigee_llm.atexit.unregister'
'google.adk.models.apigee_llm.atexit.unregister', wraps=atexit.unregister
) as mock_unregister:
local_client.close()
mock_unregister.assert_called_once_with(callback)
assert _unregister_calls_for(mock_unregister, callback) == [
mock.call(callback)
]


@pytest.mark.asyncio
Expand All @@ -1114,13 +1126,70 @@ async def test_aclose_unregisters_atexit_cleanup():
_ = local_client._client
callback = local_client._atexit_callback
with mock.patch(
'google.adk.models.apigee_llm.atexit.unregister'
'google.adk.models.apigee_llm.atexit.unregister', wraps=atexit.unregister
) as mock_unregister:
await local_client.aclose()
mock_unregister.assert_called_once_with(callback)
assert _unregister_calls_for(mock_unregister, callback) == [
mock.call(callback)
]
assert local_client._client.is_closed


@pytest.mark.asyncio
async def test_aclose_preserves_live_peer_during_other_client_collection():
"""Closing one client preserves another client's pool and exit handler."""
live_peer = CompletionsHTTPClient(base_url='https://localhost')
live_httpx_client = live_peer._client
live_callback = live_peer._atexit_callback

old_client = CompletionsHTTPClient(base_url='https://localhost')
_ = old_client._client
old_callback = old_client._atexit_callback
old_client._cycle = old_client
old_ref = weakref.ref(old_client)

local_client = CompletionsHTTPClient(base_url='https://localhost')
_ = local_client._client
callback = local_client._atexit_callback
with mock.patch(
'google.adk.models.apigee_llm.atexit.unregister', wraps=atexit.unregister
) as mock_unregister:
del old_client
gc.collect()
await local_client.aclose()

assert old_ref() is None
assert _unregister_calls_for(mock_unregister, old_callback) == [
mock.call(old_callback)
]
assert _unregister_calls_for(mock_unregister, callback) == [
mock.call(callback)
]
assert _unregister_calls_for(mock_unregister, live_callback) == []
assert not live_httpx_client.is_closed

await live_peer.aclose()
await asyncio.sleep(0)


@pytest.mark.asyncio
async def test_repeated_close_and_aclose_close_transport_once():
"""Repeated sync and async closes do not close the HTTP transport twice."""
local_client = CompletionsHTTPClient(base_url='https://localhost')
httpx_client = local_client._client
with mock.patch.object(
httpx_client._transport, 'aclose', wraps=httpx_client._transport.aclose
) as mock_transport_close:
local_client.close()
await local_client.aclose()
local_client.close()
await local_client.aclose()
await asyncio.sleep(0)

mock_transport_close.assert_awaited_once_with()
assert httpx_client.is_closed


def test_close_without_client_does_not_touch_atexit():
local_client = CompletionsHTTPClient(base_url='https://localhost')
with mock.patch(
Expand Down
Loading