Skip to content
Open
Show file tree
Hide file tree
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
22 changes: 17 additions & 5 deletions pageindex/chat_stream.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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")
Expand All @@ -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()
35 changes: 33 additions & 2 deletions pageindex/cloud_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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="")
Expand Down Expand Up @@ -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()

Expand Down
3 changes: 2 additions & 1 deletion pageindex/local_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
152 changes: 152 additions & 0 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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 == []