Files
EvoScientist-Multi/tests/test_google_stream_transport_contract.py
m4 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
test: cover stop contract, execution adapters, checkpointer race and runtime identity
2026-09-13 15:12:17 +08:00

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