diff --git a/pageindex/cloud_api.py b/pageindex/cloud_api.py index 222e1ae30..3004955e0 100644 --- a/pageindex/cloud_api.py +++ b/pageindex/cloud_api.py @@ -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. @@ -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() diff --git a/tests/test_client.py b/tests/test_client.py index d8eec89b9..82cf63815 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -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( @@ -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 (