|
9 | 9 | This mirrors _langgraph_async.py lines 62-78 and 100-127. |
10 | 10 | """ |
11 | 11 |
|
| 12 | +import asyncio |
| 13 | +from typing import override |
12 | 14 | from datetime import datetime |
13 | 15 |
|
14 | 16 | import pytest |
@@ -478,3 +480,116 @@ async def test_auto_send_created_at_forwarded(): |
478 | 480 | await auto_send(_gen(events), task_id="task1", tracer=None, streaming=streaming, created_at=dt) |
479 | 481 |
|
480 | 482 | assert all(ts == dt for ts in streaming.recorded_created_at) |
| 483 | + |
| 484 | + |
| 485 | +class _BlockingCtx(_FakeCtx): |
| 486 | + """A context whose stream_update never returns (a stalled backend). |
| 487 | +
|
| 488 | + Sets `blocked` once stream_update is awaited so the test can cancel exactly |
| 489 | + while delivery is suspended on the backend rather than on the source. |
| 490 | + """ |
| 491 | + |
| 492 | + def __init__(self, sink, content_type, initial_content, blocked): |
| 493 | + super().__init__(sink, content_type, initial_content) |
| 494 | + self.blocked = blocked |
| 495 | + |
| 496 | + @override |
| 497 | + async def stream_update(self, update): |
| 498 | + self.sink.append(("update", update)) |
| 499 | + self.blocked.set() |
| 500 | + await asyncio.Event().wait() |
| 501 | + |
| 502 | + |
| 503 | +class _BlockingStreaming(_FakeStreaming): |
| 504 | + """_FakeStreaming whose contexts block forever inside stream_update.""" |
| 505 | + |
| 506 | + def __init__(self): |
| 507 | + super().__init__() |
| 508 | + self.blocked = asyncio.Event() |
| 509 | + |
| 510 | + @override |
| 511 | + def streaming_task_message_context(self, task_id, initial_content, streaming_mode="coalesced", created_at=None): |
| 512 | + ctype = getattr(initial_content, "type", None) |
| 513 | + self.sink.append(("ctx", ctype)) |
| 514 | + self.recorded_created_at.append(created_at) |
| 515 | + return _BlockingCtx(self.sink, ctype, initial_content, self.blocked) |
| 516 | + |
| 517 | + |
| 518 | +class _PlainAsyncIterator: |
| 519 | + """An async iterator with no aclose (the AsyncIterator contract minimum).""" |
| 520 | + |
| 521 | + def __init__(self, events): |
| 522 | + self._events = iter(events) |
| 523 | + |
| 524 | + def __aiter__(self): |
| 525 | + return self |
| 526 | + |
| 527 | + async def __anext__(self): |
| 528 | + try: |
| 529 | + return next(self._events) |
| 530 | + except StopIteration: |
| 531 | + raise StopAsyncIteration from None |
| 532 | + |
| 533 | + |
| 534 | +@pytest.mark.asyncio |
| 535 | +async def test_auto_send_closes_source_when_cancelled_mid_delivery(): |
| 536 | + """Cancelling delivery while the backend blocks must close the event source. |
| 537 | +
|
| 538 | + The turn object pins its event generator, so GC cannot rescue it: when |
| 539 | + auto_send returns with the source still suspended at a yield, the tap's |
| 540 | + finally never runs and the harness CLI subprocess leaks. The source is held |
| 541 | + by a local here for the whole test, and nothing calls gc.collect(). |
| 542 | + """ |
| 543 | + streaming = _BlockingStreaming() |
| 544 | + closed: list[bool] = [] |
| 545 | + |
| 546 | + async def _recording_source(): |
| 547 | + try: |
| 548 | + yield StreamTaskMessageStart( |
| 549 | + type="start", |
| 550 | + index=0, |
| 551 | + content=TextContent(type="text", author="agent", content=""), |
| 552 | + ) |
| 553 | + yield StreamTaskMessageDelta( |
| 554 | + type="delta", |
| 555 | + index=0, |
| 556 | + delta=TextDelta(type="text", text_delta="hi"), |
| 557 | + ) |
| 558 | + yield StreamTaskMessageDone(type="done", index=0) |
| 559 | + finally: |
| 560 | + closed.append(True) |
| 561 | + |
| 562 | + source = _recording_source() |
| 563 | + task = asyncio.create_task(auto_send(source, task_id="task1", tracer=None, streaming=streaming)) |
| 564 | + await asyncio.wait_for(streaming.blocked.wait(), timeout=5) |
| 565 | + task.cancel() |
| 566 | + with pytest.raises(asyncio.CancelledError): |
| 567 | + await task |
| 568 | + |
| 569 | + assert closed == [True] |
| 570 | + assert ("close", "text") in [(s[0], s[1]) for s in streaming.sink] |
| 571 | + |
| 572 | + |
| 573 | +@pytest.mark.asyncio |
| 574 | +async def test_auto_send_accepts_source_without_aclose(): |
| 575 | + """A plain async iterator (no aclose) must deliver exactly as before.""" |
| 576 | + streaming = _FakeStreaming() |
| 577 | + events = [ |
| 578 | + StreamTaskMessageStart( |
| 579 | + type="start", |
| 580 | + index=0, |
| 581 | + content=TextContent(type="text", author="agent", content=""), |
| 582 | + ), |
| 583 | + StreamTaskMessageDelta( |
| 584 | + type="delta", |
| 585 | + index=0, |
| 586 | + delta=TextDelta(type="text", text_delta="Hi"), |
| 587 | + ), |
| 588 | + StreamTaskMessageDone(type="done", index=0), |
| 589 | + ] |
| 590 | + result = await auto_send(_PlainAsyncIterator(events), task_id="task1", tracer=None, streaming=streaming) |
| 591 | + |
| 592 | + assert result.final_text == "Hi" |
| 593 | + kinds = [s[0] for s in streaming.sink] |
| 594 | + assert kinds.count("open") == 1 |
| 595 | + assert kinds.count("close") == 1 |
0 commit comments