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
84 changes: 47 additions & 37 deletions pageindex/cloud_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,26 @@ def _enc(value: str) -> str:
return urllib.parse.quote(str(value), safe="")


def _sse_data(response: requests.Response) -> Iterator[str]:
"""Complete SSE data frames; fields join with LF and EOF discards a
pending frame. SSE is UTF-8 regardless of the response charset."""
data_lines = []
first_line = True
for raw_line in response.iter_lines():
line = raw_line.decode("utf-8", errors="replace")
if first_line:
line = line.removeprefix("\ufeff")
first_line = False
if not line:
if data_lines:
yield "\n".join(data_lines)
data_lines = []
continue
field, _, value = line.partition(":")
if field == "data":
data_lines.append(value.removeprefix(" "))


class CloudAPI:
"""
Python SDK client for the PageIndex API.
Expand Down Expand Up @@ -350,49 +370,39 @@ def _stream_chat_response(self, response: requests.Response) -> Iterator[str]:
str: Content chunks from the streaming response
"""
try:
for line in response.iter_lines():
if line:
line = line.decode('utf-8')
if line.startswith('data: '):
data = line[6:]
if data == '[DONE]':
break

try:
chunk = json.loads(data)
if chunk.get("error"):
raise PageIndexAPIError(
"Chat completion failed mid-stream: "
f"{chunk['error']}")
choices = chunk.get("choices") or [{}]
content = choices[0].get("delta", {}).get("content", "")
if content:
yield content
except json.JSONDecodeError:
continue
for data in _sse_data(response):
if data == '[DONE]':
break
try:
chunk = json.loads(data)
if chunk.get("error"):
raise PageIndexAPIError(
"Chat completion failed mid-stream: "
f"{chunk['error']}")
choices = chunk.get("choices") or [{}]
content = choices[0].get("delta", {}).get("content", "")
if content:
yield content
except json.JSONDecodeError:
continue
finally:
response.close()

def _stream_chat_response_raw(self, response: requests.Response) -> Iterator[Dict[str, Any]]:
"""Streaming chat completion with full metadata, including citation events."""
try:
for line in response.iter_lines():
if line:
line = line.decode('utf-8')
if line.startswith('data: '):
data = line[6:]
if data == '[DONE]':
break

try:
chunk = json.loads(data)
if chunk.get("error"):
raise PageIndexAPIError(
"Chat completion failed mid-stream: "
f"{chunk['error']}")
yield chunk
except json.JSONDecodeError:
continue
for data in _sse_data(response):
if data == '[DONE]':
break
try:
chunk = json.loads(data)
if chunk.get("error"):
raise PageIndexAPIError(
"Chat completion failed mid-stream: "
f"{chunk['error']}")
yield chunk
except json.JSONDecodeError:
continue
finally:
response.close()

Expand Down
119 changes: 119 additions & 0 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -2051,15 +2051,132 @@ def handler(method, url, kw):
assert client.get_folder_path("a") == "B/A"



@pytest.fixture
def cloud_sse_endpoint(monkeypatch):
"""Exercise requests and the public SDK against native SSE bytes."""
import http.server
import threading
import pageindex.cloud_api as cloud_api

replies = []
responses = []

class Handler(http.server.BaseHTTPRequestHandler):
def log_message(self, *args):
pass

def do_POST(self):
self.rfile.read(int(self.headers["Content-Length"]))
body, content_type = replies.pop(0)
self.send_response(200)
self.send_header("Content-Type", content_type)
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()


def _native_cloud_view(client, view):
if view == "answer":
return client.chat("q", stream=True, show_process=False)
return client.chat_completions("q", stream=True,
stream_metadata=view == "raw")


@pytest.mark.parametrize("view", ["text", "raw", "answer"])
@pytest.mark.parametrize("newline", ["\n", "\r\n", "\r"])
@pytest.mark.parametrize("content_type", ["text/event-stream",
"text/event-stream; charset=iso-8859-1"])
def test_cloud_chat_stream_native_sse_frames(cloud_sse_endpoint, view,
newline, content_type):
client, replies, responses = cloud_sse_endpoint
first = {"choices": [{"delta": {"content": "Résumé"}}]}
second = {"choices": [{"delta": {"content": " 世界"}}]}
citation = {"object": "chat.completion.citations", "citations": [{"page": 1}]}
body = (
"\ufeff: keepalive\n\nevent: message\ndata:bad-json\n\n"
"data:" + json.dumps(first, ensure_ascii=False) + "\n\n"
'data: {"choices": [\n: comment inside an event\n'
'data:{"delta": {"content": " 世界"}}]}\n\n'
'data: {"object": "chat.completion.citations",\n'
'data: "citations": [{"page": 1}]}\n\n'
'data\n\ndata:[DONE]\n\n'
'data:{"choices":[{"delta":{"content":"after done"}}]}\n\n'
).replace("\n", newline).encode("utf-8")
replies.append((body, content_type))
result = list(_native_cloud_view(client, view))
assert result == ([first, second, citation] if view == "raw"
else ["Résumé", " 世界"])
assert responses[-1].raw.closed


@pytest.mark.parametrize("view", ["text", "raw", "answer"])
@pytest.mark.parametrize("error_frame", [
'data:{"error":{"message":"native stream failed"}}\n\n',
'data: {"error": {\ndata:"message":"native stream failed"}}\n\n',
])
def test_cloud_chat_stream_native_sse_error(cloud_sse_endpoint, view,
error_frame):
client, replies, responses = cloud_sse_endpoint
partial = {"choices": [{"delta": {"content": "Partial"}}]}
body = ('data: ' + json.dumps(partial) + '\n\n' + error_frame).encode()
replies.append((body, "text/event-stream"))
stream = _native_cloud_view(client, view)
assert next(stream) == (partial if view == "raw" else "Partial")
with pytest.raises(PageIndexAPIError, match="native stream failed"):
list(stream)
assert responses[-1].raw.closed


@pytest.mark.parametrize("view", ["text", "raw", "answer"])
def test_cloud_chat_stream_discards_unterminated_sse_event(
cloud_sse_endpoint, view):
client, replies, responses = cloud_sse_endpoint
complete = {"choices": [{"delta": {"content": "Complete frame"}}]}
body = ('data: ' + json.dumps(complete) + '\n\n'
'data: {"error":{"message":"unfinished event"}}\n').encode()
replies.append((body, "text/event-stream"))
assert list(_native_cloud_view(client, view)) == (
[complete] if view == "raw" else ["Complete frame"])
assert responses[-1].raw.closed


def test_cloud_chat_stream_parsing(cloud, monkeypatch):
client, calls, fake = cloud
lines = [
b'data: {"choices": [{"delta": {"role": "assistant", "content": ""}}]}',
b"",
b'data: {"choices": [{"delta": {"content": "Hi"}}]}',
b"",
b'data: {"object": "chat.completion.citations", "citations": []}',
b"",
b'data: {"choices": [{"delta": {"content": " there"}}]}',
b"",
b"data: [DONE]",
b"",
]
_patch_requests(monkeypatch, lambda m, url, kw: FakeResponse(lines=lines))
pieces = list(client.chat_completions(
Expand All @@ -2079,7 +2196,9 @@ def test_cloud_chat_stream_error_chunk_raises(cloud, monkeypatch):
client, calls, fake = cloud
lines = [
b'data: {"choices": [{"delta": {"content": "Partial"}}]}',
b"",
b'data: {"error": {"message": "boom", "type": "internal_error"}}',
b"",
]
_patch_requests(monkeypatch, lambda m, url, kw: FakeResponse(lines=lines))
for stream in (
Expand Down