d4b53bfb08
Docker / build (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Build / build (push) Has been cancelled
350 lines
12 KiB
Python
350 lines
12 KiB
Python
"""Real installed SDK + isolated HTTP transport; SSE is simulated, not Google evidence."""
|
|
|
|
import asyncio
|
|
import json
|
|
from types import SimpleNamespace
|
|
|
|
import httpx
|
|
import pytest
|
|
from google import genai
|
|
from google.genai import types
|
|
from langchain_core.messages import HumanMessage
|
|
|
|
from EvoScientist.llm.gemini_interactions import (
|
|
GeminiInteractionsChatModel,
|
|
create_gemini_interactions_model,
|
|
)
|
|
|
|
EVENTS = [
|
|
{
|
|
"event_type": "content.start",
|
|
"index": 0,
|
|
"content": {"type": "text", "text": ""},
|
|
},
|
|
{
|
|
"event_type": "content.delta",
|
|
"index": 0,
|
|
"delta": {"type": "text", "text": "hello"},
|
|
},
|
|
{
|
|
"event_type": "content.delta",
|
|
"index": 0,
|
|
"delta": {"type": "text", "text": " world"},
|
|
},
|
|
{"event_type": "content.stop", "index": 0},
|
|
{
|
|
"event_type": "content.start",
|
|
"index": 1,
|
|
"content": {
|
|
"type": "function_call",
|
|
"id": "call1",
|
|
"name": "lookup",
|
|
"arguments": {"q": "test"},
|
|
},
|
|
},
|
|
{"event_type": "content.stop", "index": 1},
|
|
{
|
|
"event_type": "interaction.complete",
|
|
"interaction": {
|
|
"id": "fixture1",
|
|
"status": "completed",
|
|
"model": "gemini-test",
|
|
"usage": {
|
|
"total_input_tokens": 4,
|
|
"total_cached_tokens": 0,
|
|
"total_output_tokens": 2,
|
|
"total_tokens": 6,
|
|
},
|
|
},
|
|
},
|
|
]
|
|
|
|
|
|
class Wire(httpx.SyncByteStream, httpx.AsyncByteStream):
|
|
def __init__(self, events, block=False):
|
|
self.events = events
|
|
self.block = block
|
|
self.closed = False
|
|
self.delivered = 0
|
|
self.waiting = asyncio.Event()
|
|
|
|
def __iter__(self):
|
|
for event in self.events:
|
|
self.delivered += 1
|
|
yield ("data: " + json.dumps(event) + "\n\n").encode()
|
|
|
|
async def __aiter__(self):
|
|
for chunk in self:
|
|
yield chunk
|
|
if self.block:
|
|
self.waiting.set()
|
|
await asyncio.Event().wait()
|
|
|
|
def close(self):
|
|
self.closed = True
|
|
|
|
async def aclose(self):
|
|
self.closed = True
|
|
|
|
|
|
def setup_wire(monkeypatch, events=EVENTS, block=False, status=200):
|
|
wire = Wire(events, block)
|
|
requests = []
|
|
clients = []
|
|
|
|
def handle(request):
|
|
requests.append(json.loads(request.content))
|
|
return httpx.Response(
|
|
status, headers={"content-type": "text/event-stream"}, stream=wire
|
|
)
|
|
|
|
def client(_self):
|
|
transport = httpx.MockTransport(handle)
|
|
sdk = genai.Client(
|
|
api_key="local-fixture-not-a-credential",
|
|
http_options=types.HttpOptions(
|
|
base_url="https://isolated.invalid",
|
|
client_args={"transport": transport},
|
|
async_client_args={"transport": transport},
|
|
),
|
|
)
|
|
clients.append(sdk)
|
|
return sdk
|
|
|
|
monkeypatch.setattr(GeminiInteractionsChatModel, "_client", client)
|
|
model = create_gemini_interactions_model(model="gemini-test", api_key="fixture")
|
|
return model, wire, requests, clients
|
|
|
|
|
|
async def consume(model, mode):
|
|
if mode == "stream":
|
|
async for _ in model._astream([HumanMessage(content="hi")]):
|
|
pass
|
|
elif mode == "sync":
|
|
model.invoke("hi")
|
|
else:
|
|
await model.ainvoke("hi")
|
|
|
|
|
|
async def test_truncated_eof_is_protocol_error(monkeypatch):
|
|
for mode in ("stream", "sync", "async"):
|
|
model, wire, requests, clients = setup_wire(monkeypatch, EVENTS[:3])
|
|
with pytest.raises(RuntimeError, match=r"^MODEL_PROVIDER_PROTOCOL_ERROR$"):
|
|
await consume(model, mode)
|
|
assert requests[0]["stream"] is True
|
|
assert wire.closed
|
|
assert clients[0]._api_client._httpx_client.is_closed
|
|
|
|
|
|
async def test_stream_completion_close_and_original_errors(monkeypatch):
|
|
from google.genai._interactions import BadRequestError
|
|
|
|
model, wire, requests, clients = setup_wire(monkeypatch)
|
|
chunks = [chunk async for chunk in model._astream([HumanMessage(content="hi")])]
|
|
assert [chunk.message.content for chunk in chunks[:2]] == ["hello", " world"]
|
|
assert chunks[-1].message.usage_metadata["total_tokens"] == 6
|
|
assert chunks[-2].message.tool_calls[0]["args"] == {"q": "test"}
|
|
assert requests[0]["stream"] is True
|
|
assert wire.closed
|
|
assert clients[0]._api_client._async_httpx_client.is_closed
|
|
|
|
model, wire, _, clients = setup_wire(monkeypatch)
|
|
stream = model._astream([HumanMessage(content="hi")])
|
|
await anext(stream)
|
|
await stream.aclose()
|
|
assert wire.closed
|
|
assert clients[0]._api_client._async_httpx_client.is_closed
|
|
|
|
for mode in ("stream", "sync", "async"):
|
|
for status, events, error in (
|
|
(400, [], BadRequestError),
|
|
(
|
|
200,
|
|
[{"event_type": "error", "error": {"message": "fixture"}}],
|
|
RuntimeError,
|
|
),
|
|
):
|
|
model, wire, requests, clients = setup_wire(
|
|
monkeypatch, events, status=status
|
|
)
|
|
with pytest.raises(error) as caught:
|
|
await consume(model, mode)
|
|
if status == 400:
|
|
assert caught.value.status_code == 400
|
|
else:
|
|
assert str(caught.value) == "MODEL_PROVIDER_PROTOCOL_ERROR"
|
|
assert requests[0]["stream"] is True
|
|
assert wire.closed
|
|
assert clients[0]._api_client._httpx_client.is_closed
|
|
if mode != "sync":
|
|
assert clients[0]._api_client._async_httpx_client.is_closed
|
|
|
|
|
|
def assert_result(result):
|
|
assert result.content[0] == {"type": "text", "text": "hello world"}
|
|
assert result.tool_calls[0]["args"] == {"q": "test"}
|
|
assert result.usage_metadata["total_tokens"] == 6
|
|
assert result.additional_kwargs["gemini_interaction_content"][1]["id"] == "call1"
|
|
|
|
|
|
def test_sync_invoke_uses_sdk_stream(monkeypatch):
|
|
model, wire, requests, clients = setup_wire(monkeypatch)
|
|
result = model.invoke([HumanMessage(content="hi")])
|
|
assert requests[0]["stream"] is True
|
|
assert_result(result)
|
|
assert wire.closed
|
|
assert clients[0]._api_client._httpx_client.is_closed
|
|
|
|
|
|
async def test_async_invoke_aggregates_sdk_stream(monkeypatch):
|
|
model, wire, requests, clients = setup_wire(monkeypatch)
|
|
result = await model.ainvoke([HumanMessage(content="hi")])
|
|
assert requests[0]["stream"] is True
|
|
assert_result(result)
|
|
assert wire.closed
|
|
assert clients[0]._api_client._async_httpx_client.is_closed
|
|
assert clients[0]._api_client._httpx_client.is_closed
|
|
|
|
|
|
async def test_incremental_stream_cancellation_closes_owned_clients(monkeypatch):
|
|
model, wire, requests, clients = setup_wire(monkeypatch, EVENTS[:3], block=True)
|
|
stream = model._astream([HumanMessage(content="hi")])
|
|
first = await anext(stream)
|
|
assert first.message.content == "hello"
|
|
assert wire.delivered == 2
|
|
assert requests[0]["stream"] is True
|
|
assert (await anext(stream)).message.content == " world"
|
|
task = asyncio.ensure_future(anext(stream))
|
|
await asyncio.wait_for(wire.waiting.wait(), 2)
|
|
task.cancel()
|
|
with pytest.raises(asyncio.CancelledError):
|
|
await task
|
|
assert wire.closed
|
|
assert clients[0]._api_client._async_httpx_client.is_closed
|
|
assert clients[0]._api_client._httpx_client.is_closed
|
|
|
|
|
|
@pytest.mark.parametrize("mode", ["sync", "async", "stream"])
|
|
@pytest.mark.parametrize("failure", ["protocol", "cancel", "success"])
|
|
async def test_cleanup_failures_preserve_primary_and_attempt_all(
|
|
monkeypatch, caplog, mode, failure
|
|
):
|
|
"""Independent fault injection, not a replay of transport success tests."""
|
|
primary = {
|
|
"protocol": RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR"),
|
|
"cancel": asyncio.CancelledError("primary cancellation"),
|
|
"success": None,
|
|
}[failure]
|
|
attempts = []
|
|
secret = "cleanup-secret-must-not-be-logged"
|
|
|
|
def fail_close(resource):
|
|
attempts.append(resource)
|
|
raise OSError(secret)
|
|
|
|
class BrokenStream:
|
|
def __iter__(self):
|
|
yield from EVENTS
|
|
if primary is not None:
|
|
raise primary
|
|
|
|
async def __aiter__(self):
|
|
for event in self:
|
|
yield event
|
|
|
|
def close(self):
|
|
fail_close("stream")
|
|
|
|
class AsyncBrokenStream:
|
|
__aiter__ = BrokenStream.__aiter__
|
|
__iter__ = BrokenStream.__iter__
|
|
|
|
async def close(self):
|
|
fail_close("stream")
|
|
|
|
async def create_async(**request):
|
|
assert request["stream"] is True
|
|
return AsyncBrokenStream()
|
|
|
|
def create_sync(**request):
|
|
assert request["stream"] is True
|
|
return BrokenStream()
|
|
|
|
async def close_async():
|
|
fail_close("async_client")
|
|
|
|
client = SimpleNamespace(
|
|
interactions=SimpleNamespace(create=create_sync),
|
|
aio=SimpleNamespace(
|
|
interactions=SimpleNamespace(create=create_async), aclose=close_async
|
|
),
|
|
close=lambda: fail_close("client"),
|
|
)
|
|
monkeypatch.setattr(GeminiInteractionsChatModel, "_client", lambda self: client)
|
|
model = create_gemini_interactions_model(model="gemini-test", api_key="fixture")
|
|
expected = type(primary) if primary is not None else RuntimeError
|
|
operation = (
|
|
model._agenerate([HumanMessage(content="hi")])
|
|
if mode == "async" and failure == "cancel"
|
|
else consume(model, mode)
|
|
)
|
|
with pytest.raises(expected) as caught:
|
|
await operation
|
|
if primary is not None:
|
|
assert caught.value is primary
|
|
else:
|
|
assert str(caught.value) == "MODEL_PROVIDER_CLEANUP_ERROR"
|
|
resources = (
|
|
["stream", "client"] if mode == "sync" else ["stream", "async_client", "client"]
|
|
)
|
|
assert attempts == resources
|
|
records = [r for r in caplog.records if r.name.endswith("gemini_interactions")]
|
|
assert len(records) == len(resources)
|
|
assert all(r.exc_info is None for r in records)
|
|
assert secret not in caplog.text
|
|
for resource, record in zip(resources, records, strict=True):
|
|
assert (
|
|
record.getMessage() == f"MODEL_PROVIDER_CLEANUP_ERROR resource={resource}"
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize("mode", ["async", "stream"])
|
|
async def test_public_task_cancellation_with_cleanup_failure(monkeypatch, caplog, mode):
|
|
model, wire, requests, clients = setup_wire(monkeypatch, EVENTS[:3], block=True)
|
|
attempts = []
|
|
original_sync = genai.Client.close
|
|
original_async = genai.client.AsyncClient.aclose
|
|
|
|
def close_sync(self):
|
|
attempts.append("client")
|
|
original_sync(self)
|
|
raise OSError("synthetic-cleanup-secret")
|
|
|
|
async def close_async(self):
|
|
attempts.append("async_client")
|
|
await original_async(self)
|
|
raise OSError("synthetic-cleanup-secret")
|
|
|
|
monkeypatch.setattr(genai.Client, "close", close_sync)
|
|
monkeypatch.setattr(genai.client.AsyncClient, "aclose", close_async)
|
|
|
|
async def public_operation():
|
|
if mode == "async":
|
|
await model.ainvoke("hi")
|
|
else:
|
|
async for _ in model.astream("hi"):
|
|
pass
|
|
|
|
task = asyncio.create_task(public_operation())
|
|
await asyncio.wait_for(wire.waiting.wait(), 2)
|
|
task.cancel("public-cancel-marker")
|
|
with pytest.raises(asyncio.CancelledError, match="public-cancel-marker"):
|
|
await task
|
|
assert task.cancelled()
|
|
assert attempts == ["async_client", "client"]
|
|
assert requests[0]["stream"] is True
|
|
assert wire.closed
|
|
assert clients[0]._api_client._async_httpx_client.is_closed
|
|
assert clients[0]._api_client._httpx_client.is_closed
|
|
assert "synthetic-cleanup-secret" not in caplog.text
|