From 0da11ebba1d34e44e55948bd4f023300aec82c7a Mon Sep 17 00:00:00 2001 From: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> Date: Thu, 1 Oct 2026 09:30:45 -0400 Subject: [PATCH] fix: close unconsumed cloud chat responses Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> --- pageindex/chat_stream.py | 22 ++++-- pageindex/cloud_api.py | 35 ++++++++- pageindex/local_chat.py | 3 +- tests/test_client.py | 152 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 204 insertions(+), 8 deletions(-) diff --git a/pageindex/chat_stream.py b/pageindex/chat_stream.py index ae4b98567..85cc23a7e 100644 --- a/pageindex/chat_stream.py +++ b/pageindex/chat_stream.py @@ -1,7 +1,7 @@ """chat(stream=True)'s return type: one run, one view — text or events.""" from __future__ import annotations -from typing import Any, Iterator, Optional +from typing import Any, Callable, Iterator, Optional from .errors import PageIndexAPIError @@ -12,12 +12,14 @@ class ChatStream: the typed process event dicts. One underlying run — consume exactly one view; call chat() again for the other.""" - def __init__(self, text, events): + def __init__(self, text, events, + on_close: Optional[Callable[[], None]] = None): self._text = text # () -> Iterator[str] self._events = events # () -> Iterator[dict], or the refusal text self._view: Optional[str] = None self._it: Any = None self._closed = False + self._on_close = on_close def _claim(self, view: str) -> None: if self._view is not None and self._view != view: @@ -30,6 +32,8 @@ def __iter__(self) -> "ChatStream": return self def __next__(self) -> str: + if self._closed: + raise StopIteration self._claim("text") if self._it is None: if self._closed: @@ -47,6 +51,8 @@ def events(self) -> Iterator[dict]: not merely reading the attribute — claims the view, so debugger panes and getattr probing stay side-effect free.""" def consume(): + if self._closed: + return if isinstance(self._events, str): raise PageIndexAPIError(self._events) self._claim("events") @@ -63,7 +69,13 @@ def close(self) -> None: """Stop the run: closes the open view, and the stream is dead afterwards, like a closed generator (own-model chat: a run never consumed never starts).""" + if self._closed: + return self._closed = True - close = getattr(self._it, "close", None) - if close is not None: - close() + try: + close = getattr(self._it, "close", None) + if close is not None: + close() + finally: + if self._on_close is not None: + self._on_close() diff --git a/pageindex/cloud_api.py b/pageindex/cloud_api.py index 222e1ae30..9fe879640 100644 --- a/pageindex/cloud_api.py +++ b/pageindex/cloud_api.py @@ -9,6 +9,36 @@ from .naming import sanitize_filename, validate_folder_name +class _ClosingChatIterator: + """Own the eagerly opened response, even before parsing starts.""" + + def __init__(self, iterator, response: requests.Response): + self._iterator = iterator + self._response = response + self._closed = False + + def __iter__(self): + return self + + def __next__(self): + if self._closed: + raise StopIteration + try: + return next(self._iterator) + except BaseException: + self.close() + raise + + def close(self) -> None: + if self._closed: + return + self._closed = True + try: + self._iterator.close() + finally: + self._response.close() + + def _enc(value: str) -> str: """URL-encode a path segment (ids may contain / ? # or spaces).""" return urllib.parse.quote(str(value), safe="") @@ -333,9 +363,10 @@ def chat_completions( if stream: if stream_metadata: - return self._stream_chat_response_raw(response) + iterator = self._stream_chat_response_raw(response) else: - return self._stream_chat_response(response) + iterator = self._stream_chat_response(response) + return _ClosingChatIterator(iterator, response) else: return response.json() diff --git a/pageindex/local_chat.py b/pageindex/local_chat.py index 42c7f5176..f29ec5760 100644 --- a/pageindex/local_chat.py +++ b/pageindex/local_chat.py @@ -902,7 +902,8 @@ def run_cloud_chat_stream(chunks, events=("chat events are produced by the in-process agent, " "which the managed chat endpoint does not serve — " "construct the client with chat_model=... (or a chat= " - "model) to run the agent in your process.")) + "model) to run the agent in your process."), + on_close=getattr(chunks, "close", None)) def run_chat_stream(client, messages, doc_id=None, model=None, diff --git a/tests/test_client.py b/tests/test_client.py index d8eec89b9..7a50af1c2 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1624,6 +1624,105 @@ def test_chat_completions_local_needs_openai_agents(local_client, monkeypatch): # ── cloud mode: request wiring ── + +@pytest.fixture +def cloud_stream_ownership_endpoint(monkeypatch): + """Retain real unread HTTP/1.1 responses to verify explicit cleanup.""" + import http.server + import threading + import pageindex.cloud_api as cloud_api + + replies = [] + responses = [] + + class Handler(http.server.BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def log_message(self, *args): + pass + + def handle(self): + try: + super().handle() + except ConnectionResetError: + pass # cancelling an unread response can reset the connection + + def do_POST(self): + self.rfile.read(int(self.headers["Content-Length"])) + body = replies.pop(0) + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + server = http.server.ThreadingHTTPServer(("127.0.0.1", 0), Handler) + server.daemon_threads = True + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + real_post = cloud_api.requests.post + + def post(*args, **kwargs): + response = real_post(*args, **kwargs) + responses.append(response) + return response + + monkeypatch.setattr(cloud_api.requests, "post", post) + client = PageIndexClient(api_key="local-test-only") + client.BASE_URL = f"http://127.0.0.1:{server.server_port}" + try: + yield client, replies, responses + finally: + for response in responses: + response.close() + server.shutdown() + server.server_close() + thread.join() + + +@pytest.mark.parametrize("view", ["text", "raw", "answer", "protocol"]) +@pytest.mark.parametrize("consume", ["none", "partial", "complete", "error"]) +def test_cloud_chat_stream_closes_owned_native_response( + cloud_stream_ownership_endpoint, view, consume): + client, replies, responses = cloud_stream_ownership_endpoint + chunk = {"choices": [{"delta": {"content": "Partial"}}]} + body = 'data: ' + json.dumps(chunk) + '\n\n' + if consume == "error": + body += 'data: {"error":{"message":"native failure"}}\n\n' + else: + body += 'data: ' + json.dumps(chunk) + '\n\ndata: [DONE]\n\n' + # Keep unread bytes beyond requests' first buffer, including after + # the terminal event, so cleanup must close the owned native socket. + body += ": " + "padding " * 200 + "\n\n" + replies.append(body.encode()) + if view == "answer": + stream = client.chat("q", stream=True, show_process=False) + elif view == "protocol": + stream = client.chat("q", stream=True, protocol="chat_completions") + else: + stream = client.chat_completions("q", stream=True, + stream_metadata=view == "raw") + response = responses[-1] + assert not response.raw.closed + socket = response.raw._fp.fp.raw._sock + assert socket.fileno() >= 0 + if consume in ("partial", "error"): + assert next(stream) == (chunk if view in ("raw", "protocol") + else "Partial") + if consume == "complete": + assert list(stream) == ([chunk, chunk] if view in ("raw", "protocol") + else ["Partial", "Partial"]) + elif consume == "error": + with pytest.raises(PageIndexAPIError, match="native failure"): + list(stream) + stream.close() + stream.close() + assert response.raw.closed + assert socket.fileno() == -1 + with pytest.raises(StopIteration): + next(stream) + + class FakeResponse: def __init__(self, payload=None, status_code=200, text="", content=b"{}", lines=None): @@ -3313,3 +3412,56 @@ def token_counter(model=None, text=None, **_): monkeypatch.setattr(litellm, "token_counter", lambda model=None, text=None, **_: 3 if model else 7) assert pageindex.utils.count_tokens("x", model="m") == 3 # the model's own count wins when it works + + +@pytest.mark.parametrize("raw", [False, True]) +def test_cloud_chat_close_releases_response_when_parser_close_fails( + cloud_stream_ownership_endpoint, monkeypatch, raw): + from pageindex.cloud_api import CloudAPI + + class BrokenParser: + def __next__(self): + raise AssertionError("close must not start parsing") + + def close(self): + raise RuntimeError("parser close failed") + + method = "_stream_chat_response_raw" if raw else "_stream_chat_response" + monkeypatch.setattr(CloudAPI, method, lambda self, response: BrokenParser()) + client, replies, responses = cloud_stream_ownership_endpoint + replies.append(b'data: {"choices":[]}\n\n') + stream = client.chat_completions("q", stream=True, stream_metadata=raw) + response = responses[-1] + socket = response.raw._fp.fp.raw._sock + assert socket.fileno() >= 0 + with pytest.raises(RuntimeError, match="parser close failed"): + stream.close() + assert response.raw.closed + assert socket.fileno() == -1 + stream.close() + with pytest.raises(StopIteration): + next(stream) + + +def test_chat_stream_close_preserves_unstarted_own_model_run(): + from pageindex.chat_stream import ChatStream + + started = [] + + def text(): + started.append("text") + return iter(["answer"]) + + def events(): + started.append("events") + return iter([{"type": "answer", "delta": "answer"}]) + + stream = ChatStream(text=text, events=events) + stream.close() + stream.close() + assert started == [] + assert list(stream) == [] + events_stream = ChatStream(text=text, events=events) + events_stream.close() + assert list(events_stream.events) == [] + assert started == []