From c441ef7e756505c15318e8f7ea946be1ee1d8ff3 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:41:53 +0200 Subject: [PATCH 01/16] feat(types): add conversation streaming config and builder - introduce typed stream callbacks and message chunks for conversation search - add decorator-based StreamConfigBuilder and wire streaming params into search --- src/typesense/types/document.py | 168 +++++++++++++++++++++++++++++++- 1 file changed, 163 insertions(+), 5 deletions(-) diff --git a/src/typesense/types/document.py b/src/typesense/types/document.py index ee44b04..f1307f5 100644 --- a/src/typesense/types/document.py +++ b/src/typesense/types/document.py @@ -586,6 +586,162 @@ class NLLanguageParameters(typing.TypedDict): nl_query_debug: typing.NotRequired[bool] +class MessageChunk(typing.TypedDict): + """ + A single chunk from a conversation stream response. + + Attributes: + conversation_id (str): ID of the conversation. + message (str): Message content for this chunk. + """ + + conversation_id: str + message: str + + +class StreamConfig(typing.Generic[TDoc], typing.TypedDict, total=False): + """ + Configuration for streaming conversation search responses. + + Attributes: + on_chunk: Callback invoked for each streamed chunk (conversation_id, message). + on_complete: Callback invoked when the stream completes with the full search response. + on_error: Callback invoked if an error occurs during streaming. + """ + + on_chunk: typing.Callable[[MessageChunk], None] + on_complete: "OnCompleteCallback[TDoc]" + on_error: typing.Callable[[BaseException], None] + + +OnChunkCallback = typing.Callable[[MessageChunk], None] + + +class OnCompleteCallback(typing.Protocol[TDoc]): + def __call__(self, response: "SearchResponse[TDoc]") -> None: ... + + +OnErrorCallback = typing.Callable[[BaseException], None] + + +class StreamConfigBuilder(typing.Generic[TDoc]): + """ + Builder for StreamConfig using decorators. + + Example: + >>> stream = StreamConfigBuilder() + >>> + >>> @stream.on_chunk + ... def handle_chunk(chunk: MessageChunk) -> None: + ... print(chunk["message"], end="", flush=True) + >>> + >>> @stream.on_complete + ... def handle_complete(response: dict) -> None: + ... print(f"Done! Found {response.get('found', 0)}") + >>> + >>> response = await client.collections["docs"].documents.search({ + ... "q": "query", + ... "query_by": "content", + ... "conversation_stream": True, + ... "stream_config": stream, + ... }) + """ + + def __init__(self) -> None: + """Initialize an empty StreamConfigBuilder.""" + self._on_chunk: OnChunkCallback | None = None + self._on_complete: OnCompleteCallback[TDoc] | None = None + self._on_error: OnErrorCallback | None = None + + def on_chunk(self, func: OnChunkCallback) -> OnChunkCallback: + """ + Decorator to register an on_chunk callback. + + Args: + func: Callback invoked for each streamed message chunk. + + Returns: + The original function (unmodified). + """ + self._on_chunk = func + return func + + def on_complete(self, func: OnCompleteCallback[TDoc]) -> OnCompleteCallback[TDoc]: + """ + Decorator to register an on_complete callback. + + Args: + func: Callback invoked when streaming completes with the full response. + + Returns: + The original function (unmodified). + """ + self._on_complete = func + return func + + def on_error(self, func: OnErrorCallback) -> OnErrorCallback: + """ + Decorator to register an on_error callback. + + Args: + func: Callback invoked if an error occurs during streaming. + + Returns: + The original function (unmodified). + """ + self._on_error = func + return func + + def build(self) -> StreamConfig[TDoc]: + """ + Build the StreamConfig dictionary. + + Returns: + A StreamConfig with the registered callbacks. + """ + config: StreamConfig[TDoc] = {} + if self._on_chunk is not None: + config["on_chunk"] = self._on_chunk + if self._on_complete is not None: + config["on_complete"] = self._on_complete + if self._on_error is not None: + config["on_error"] = self._on_error + return config + + def get( + self, + key: typing.Literal["on_chunk", "on_complete", "on_error"], + default: typing.Callable[..., None] | None = None, + ) -> typing.Callable[..., None] | None: + """ + Get a callback by key (for compatibility with dict-like access). + + Args: + key: The callback name ('on_chunk', 'on_complete', or 'on_error'). + default: Default value if the callback is not set. + + Returns: + The callback function or the default value. + """ + return self.build().get(key, default) + + +class ConversationStreamParameters(typing.Generic[TDoc], typing.TypedDict): + """ + Parameters for conversational search streaming. + + Attributes: + conversation_stream (bool): When true, the search response is streamed (SSE). + stream_config: Callbacks for stream events. Not sent to the API. + Can be a StreamConfig dict or a StreamConfigBuilder instance. + """ + + conversation_stream: typing.NotRequired[bool] + stream_config: typing.NotRequired[ + typing.Union[StreamConfig[TDoc], StreamConfigBuilder[TDoc]] + ] + + class SearchParameters( RequiredSearchParameters, QueryParameters, @@ -598,11 +754,13 @@ class SearchParameters( TypoToleranceParameters, CachingParameters, NLLanguageParameters, + ConversationStreamParameters[TDoc], + typing.Generic[TDoc], ): """Parameters for searching documents.""" -class MultiSearchParameters(SearchParameters): +class MultiSearchParameters(SearchParameters[TDoc], typing.Generic[TDoc]): """ Parameters for performing a [Federated/Multi-Search](https://typesense.org/docs/26.0/api/federated-multi-search.html#federated-multi-search). @@ -867,7 +1025,7 @@ class LLMResponse(typing.TypedDict): model: str -class ParsedNLQuery(typing.TypedDict): +class ParsedNLQuery(typing.Generic[TDoc], typing.TypedDict): """ Schema for a parsed natural language query. @@ -879,8 +1037,8 @@ class ParsedNLQuery(typing.TypedDict): """ parse_time_ms: int - generated_params: SearchParameters - augmented_params: SearchParameters + generated_params: SearchParameters[TDoc] + augmented_params: SearchParameters[TDoc] llm_response: typing.NotRequired[LLMResponse] @@ -912,7 +1070,7 @@ class SearchResponse(typing.Generic[TDoc], typing.TypedDict): hits: typing.List[Hit[TDoc]] grouped_hits: typing.NotRequired[typing.List[GroupedHit[TDoc]]] conversation: typing.NotRequired[Conversation] - parsed_nl_query: typing.NotRequired[ParsedNLQuery] + parsed_nl_query: typing.NotRequired[ParsedNLQuery[TDoc]] class DeleteSingleDocumentParameters(typing.TypedDict): From c6af4df9e69b9931ddf03f6fa016bd02cf67af45 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:42:47 +0200 Subject: [PATCH 02/16] feat(streaming): add sse stream parsing and chunk combiner - parse conversation stream sse lines into message chunks or search responses - combine streamed chunks into a final search response for async calls --- src/typesense/stream_handlers.py | 164 +++++++++++++++++++++++++++++++ 1 file changed, 164 insertions(+) create mode 100644 src/typesense/stream_handlers.py diff --git a/src/typesense/stream_handlers.py b/src/typesense/stream_handlers.py new file mode 100644 index 0000000..45e38d5 --- /dev/null +++ b/src/typesense/stream_handlers.py @@ -0,0 +1,164 @@ +""" +SSE stream parsing and chunk combining for conversation search streaming. + +This module contains pure logic for parsing server-sent event lines from +conversation_stream responses and combining message chunks into a final +search response. Used by async API calls. +""" + +import json +import sys + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +from typesense.types.document import MessageChunk + +JSONPrimitive: typing.TypeAlias = typing.Union[str, int, float, bool, None] +JSONValue: typing.TypeAlias = typing.Union[ + JSONPrimitive, typing.Dict[str, "JSONValue"], typing.List["JSONValue"] +] +JSONDict: typing.TypeAlias = typing.Dict[str, JSONValue] + +_SEARCH_RESPONSE_KEYS = frozenset( + {"results", "found", "hits", "page", "search_time_ms"} +) + +StreamChunk: typing.TypeAlias = typing.Union[MessageChunk, JSONDict] + + +def parse_sse_line(line: str) -> typing.Optional[StreamChunk]: + """ + Parse a single SSE line into a MessageChunk, search response dict, or None. + + Handles: + - Empty lines and "data: [DONE]" -> None + - "data: {...}" -> parse JSON, return MessageChunk or search response + - Raw JSON line starting with "{" -> same + - Plain text -> return chunk with conversation_id="unknown", message=line + + Returns: + MessageChunk for conversation chunks, dict for search responses, or None to skip. + """ + line = line.strip() + if not line or line == "data: [DONE]": + return None + + # SSE format: "data: {...}" + if line.startswith("data:"): + content = line[5:].strip() + return _parse_data_content(content) + + # Raw JSON + if line.startswith("{"): + return _parse_json_content(line) + + return _chunk_from_text(line) + + +def _parse_data_content(content: str) -> typing.Optional[StreamChunk]: + """Parse the content after 'data:' into a MessageChunk, search response, or None.""" + if not content: + return None + if content.startswith("{"): + return _parse_json_content(content) + return _chunk_from_text(content) + + +def _parse_json_content(raw: str) -> StreamChunk: + """Parse a JSON string into a MessageChunk or search response dict.""" + try: + data = json.loads(raw) + except json.JSONDecodeError: + return _chunk_from_text(raw) + if not isinstance(data, dict): + return _chunk_from_text(json.dumps(data)) + + parsed = typing.cast(JSONDict, data) + conversation_id = parsed.get("conversation_id") + message = parsed.get("message") + nested_conversation = parsed.get("conversation") + + if conversation_id is None or message is None: + if isinstance(nested_conversation, dict): + nested_conversation_id = nested_conversation.get("conversation_id") + nested_message = nested_conversation.get("message") + if conversation_id is None and nested_conversation_id is not None: + conversation_id = nested_conversation_id + if message is None and nested_message is not None: + message = nested_message + + if conversation_id is None: + parsed["conversation_id"] = "unknown" + elif not isinstance(conversation_id, str): + parsed["conversation_id"] = str(conversation_id) + else: + parsed["conversation_id"] = conversation_id + + if message is None: + parsed["message"] = "" + elif not isinstance(message, str): + parsed["message"] = str(message) + else: + parsed["message"] = message + + return parsed + + +def _is_search_response_dict(data: typing.Mapping[str, JSONValue]) -> bool: + """Check if a dict is a search response (has found, hits, results, etc.).""" + return bool(set(data.keys()) & _SEARCH_RESPONSE_KEYS) + + +def is_message_chunk(chunk: JSONValue) -> bool: + """Return True if chunk is a conversation message chunk (has conversation_id and message).""" + if not isinstance(chunk, dict): + return False + if "message" not in chunk or "conversation_id" not in chunk: + return False + return not _is_search_response_dict(chunk) + + +def is_complete_search_response(chunk: JSONValue) -> bool: + """Return True if chunk looks like a full search response (has hits, found, etc.).""" + if not isinstance(chunk, dict) or not chunk: + return False + keys = set(chunk.keys()) + return bool(keys & _SEARCH_RESPONSE_KEYS) + + +def combine_stream_chunks( + chunks: typing.Sequence[StreamChunk], +) -> JSONDict: + """ + Combine streamed chunks into a single search response. + + - If no chunks, return empty dict. + - If one chunk, return it. + - If we have message chunks (conversation_id + message), find the metadata + chunk (complete search response) and return it; otherwise return last chunk + if it is complete. + - For regular search streams, return the last chunk if it is a complete response. + """ + if not chunks: + return {} + if len(chunks) == 1: + return typing.cast(JSONDict, chunks[0]) + + message_chunks = [c for c in chunks if is_message_chunk(c)] + if message_chunks: + for chunk in chunks: + if is_complete_search_response(chunk): + return typing.cast(JSONDict, chunk) + return typing.cast(JSONDict, chunks[-1]) + + last = chunks[-1] + if is_complete_search_response(last): + return typing.cast(JSONDict, last) + return typing.cast(JSONDict, last) + + +def _chunk_from_text(text: str) -> MessageChunk: + return {"conversation_id": "unknown", "message": text} From 16d0301a475a4a196b9e4e32e6467e7a87749398 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:54:36 +0200 Subject: [PATCH 03/16] feat(async): support streaming conversation search over sse - add async sse handling with chunk parsing, callbacks, and final response combine - wire stream_config and conversation_stream through async search api --- src/typesense/async_/api_call.py | 103 ++++++++++++++++++++++++++++++ src/typesense/async_/documents.py | 12 +++- 2 files changed, 114 insertions(+), 1 deletion(-) diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index bfaf6de..f4b1063 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -32,6 +32,7 @@ """ import asyncio +import json import sys from types import MappingProxyType, TracebackType @@ -59,6 +60,14 @@ ) from typesense.node_manager import NodeManager from typesense.request_handler import RequestHandler +from typesense.stream_handlers import ( + JSONDict, + StreamChunk, + combine_stream_chunks, + is_message_chunk, + parse_sse_line, +) +from typesense.types.document import StreamConfig if sys.version_info >= (3, 11): import typing @@ -225,6 +234,8 @@ async def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[False], params: typing.Union[TParams, None] = None, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, ) -> str: """ Execute an async GET request to the Typesense API. @@ -246,6 +257,8 @@ async def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[True] = True, params: typing.Union[TParams, None] = None, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, ) -> TEntityDict: """ Execute an async GET request to the Typesense API. @@ -266,6 +279,8 @@ async def get( entity_type: typing.Type[TEntityDict], as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, params: typing.Union[TParams, None] = None, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, ) -> typing.Union[TEntityDict, str]: """ Execute an async GET request to the Typesense API. @@ -285,6 +300,8 @@ async def get( entity_type, as_json, params=params, + stream_config=stream_config, + is_streaming_request=is_streaming_request, ) @typing.overload @@ -453,6 +470,8 @@ async def _execute_request( as_json: typing.Literal[True], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> TEntityDict: """Execute an async request with retry logic.""" @@ -466,6 +485,8 @@ async def _execute_request( as_json: typing.Literal[False], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> str: """Execute an async request with retry logic.""" @@ -478,6 +499,8 @@ async def _execute_request( as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: """ @@ -509,6 +532,10 @@ async def _execute_request( node, url, request_kwargs = self._prepare_request_params(endpoint, **kwargs) try: + if is_streaming_request and method == "GET": + return await self._handle_streaming_get( + url, entity_type, stream_config, **request_kwargs + ) return await self._make_request_and_process_response( method, node, @@ -521,6 +548,13 @@ async def _execute_request( raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) + if is_streaming_request and stream_config: + on_error = stream_config.get("on_error") + if on_error: + try: + on_error(server_error) + except Exception: + pass if num_retries < self.config.num_retries: await asyncio.sleep(self.config.retry_interval_seconds) return await self._execute_request( @@ -530,6 +564,8 @@ async def _execute_request( as_json, last_exception=server_error, num_retries=num_retries + 1, + stream_config=stream_config, + is_streaming_request=is_streaming_request, **kwargs, ) @@ -559,6 +595,73 @@ async def _make_request_and_process_response( else typing.cast(str, request_response) ) + async def _handle_streaming_get( + self, + url: str, + entity_type: typing.Type[TEntityDict], + stream_config: StreamConfig[TEntityDict] | None, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> TEntityDict: + """Perform an async streaming GET, parse SSE lines, invoke callbacks, return combined result.""" + headers: typing.Dict[str, str] = { + self.request_handler.api_key_header_name: self.config.api_key, + "Accept": "text/event-stream", + } + headers.update(self.config.additional_headers) + extra_headers = kwargs.get("headers") + if extra_headers: + headers.update(extra_headers) + + params = kwargs.get("params") + content: typing.Union[str, bytes, None] = None + if body := kwargs.get("data"): + if isinstance(body, (str, bytes)): + content = body + else: + content = json.dumps(body) + + all_chunks: typing.List[StreamChunk] = [] + async with self._client.stream( + "GET", + url, + params=params, + content=content, + headers=headers, + timeout=self.config.connection_timeout_seconds, + ) as response: + if response.status_code < 200 or response.status_code >= 300: + await response.aread() + error_message = self.request_handler._get_error_message(response) + raise self.request_handler._get_exception(response.status_code)( + response.status_code, + error_message, + ) + async for line in response.aiter_lines(): + chunk = parse_sse_line(line) + if chunk is not None: + all_chunks.append(chunk) + if stream_config and is_message_chunk(chunk): + on_chunk = stream_config.get("on_chunk") + if on_chunk: + try: + on_chunk(chunk) + except Exception: + pass + + self.node_manager.set_node_health( + self.node_manager.get_node(), + is_healthy=True, + ) + final: JSONDict = combine_stream_chunks(all_chunks) + if stream_config: + on_complete = stream_config.get("on_complete") + if on_complete: + try: + on_complete(typing.cast(TEntityDict, final)) + except Exception: + pass + return typing.cast(TEntityDict, final) + def _prepare_request_params( self, endpoint: str, diff --git a/src/typesense/async_/documents.py b/src/typesense/async_/documents.py index 399c82d..c43ac35 100644 --- a/src/typesense/async_/documents.py +++ b/src/typesense/async_/documents.py @@ -43,6 +43,7 @@ ImportResponseWithId, SearchParameters, SearchResponse, + StreamConfigBuilder, UpdateByFilterParameters, UpdateByFilterResponse, ) @@ -362,16 +363,25 @@ async def search(self, search_parameters: SearchParameters) -> SearchResponse[TD Args: search_parameters (SearchParameters): The search parameters. + Use conversation_stream=True and optionally stream_config (on_chunk, + on_complete, on_error) for conversational search streaming. Returns: SearchResponse[TDoc]: The search response containing matching documents. """ - stringified_search_params = stringify_search_params(search_parameters) + params_for_api = dict(search_parameters) + stream_config = params_for_api.pop("stream_config", None) + if isinstance(stream_config, StreamConfigBuilder): + stream_config = stream_config.build() + conversation_stream = params_for_api.get("conversation_stream") is True + stringified_search_params = stringify_search_params(params_for_api) response: SearchResponse[TDoc] = await self.api_call.get( self._endpoint_path("search"), params=stringified_search_params, entity_type=SearchResponse, as_json=True, + stream_config=stream_config, + is_streaming_request=conversation_stream, ) return response From a3f8e5e96870aa5d38b836383973dbd8dbb254d4 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:46:55 +0200 Subject: [PATCH 04/16] test(streaming): add fixtures for async and sync stream handling - provide fake sse stream responses and contexts for unit tests - add integration fixtures for conversational streaming collections and docs --- tests/fixtures/streaming_fixtures.py | 180 +++++++++++++++++++++++++++ 1 file changed, 180 insertions(+) create mode 100644 tests/fixtures/streaming_fixtures.py diff --git a/tests/fixtures/streaming_fixtures.py b/tests/fixtures/streaming_fixtures.py new file mode 100644 index 0000000..4df0d52 --- /dev/null +++ b/tests/fixtures/streaming_fixtures.py @@ -0,0 +1,180 @@ +"""Fixtures for streaming tests.""" + +import json +import os +import sys +from types import TracebackType + +import pytest +import requests + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + + +JSONPrimitive: typing.TypeAlias = typing.Union[str, int, float, bool, None] +JSONValue: typing.TypeAlias = typing.Union[ + JSONPrimitive, typing.Dict[str, "JSONValue"], typing.List["JSONValue"] +] +JSONDict: typing.TypeAlias = typing.Dict[str, JSONValue] + + +class FakeAsyncStreamResponse: + """Minimal async streaming response for httpx.AsyncClient.stream().""" + + def __init__( + self, + *, + lines: typing.Sequence[str], + status_code: int = 200, + headers: typing.Mapping[str, str] | None = None, + text: str = "", + ) -> None: + self.status_code = status_code + self._lines = list(lines) + self.headers = dict(headers or {}) + self.text = text + + async def aiter_lines(self) -> typing.AsyncIterator[str]: + for line in self._lines: + yield line + + async def aread(self) -> bytes: + return self.text.encode() + + def json(self) -> JSONDict: + return typing.cast(JSONDict, json.loads(self.text)) + + +class FakeAsyncStreamContext: + """Async context manager that yields a fake streaming response.""" + + def __init__(self, response: FakeAsyncStreamResponse) -> None: + self._response = response + + async def __aenter__(self) -> FakeAsyncStreamResponse: + return self._response + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> None: + return None + + +class FakeStreamResponse: + """Minimal streaming response for httpx.Client.stream().""" + + def __init__( + self, + *, + lines: typing.Sequence[str], + status_code: int = 200, + headers: typing.Mapping[str, str] | None = None, + text: str = "", + ) -> None: + self.status_code = status_code + self._lines = list(lines) + self.headers = dict(headers or {}) + self.text = text + + def iter_lines(self) -> typing.Iterator[str]: + for line in self._lines: + yield line + + def read(self) -> bytes: + return self.text.encode() + + def json(self) -> JSONDict: + return typing.cast(JSONDict, json.loads(self.text)) + + +class FakeStreamContext: + """Sync context manager that yields a fake streaming response.""" + + def __init__(self, response: FakeStreamResponse) -> None: + self._response = response + + def __enter__(self) -> FakeStreamResponse: + return self._response + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> None: + return None + + +@pytest.fixture(name="stream_response_async") +def stream_response_async_fixture() -> type[FakeAsyncStreamResponse]: + return FakeAsyncStreamResponse + + +@pytest.fixture(name="stream_context_async") +def stream_context_async_fixture() -> type[FakeAsyncStreamContext]: + return FakeAsyncStreamContext + + +@pytest.fixture(name="stream_response") +def stream_response_fixture() -> type[FakeStreamResponse]: + return FakeStreamResponse + + +@pytest.fixture(name="stream_context") +def stream_context_fixture() -> type[FakeStreamContext]: + return FakeStreamContext + + +@pytest.fixture(name="create_streaming_collection") +def create_streaming_collection_fixture(delete_all: None) -> str: + """Create a collection for streaming tests with an auto-embedding field.""" + open_ai_key = os.environ.get("OPEN_AI_KEY") + if not open_ai_key: + pytest.skip("OPEN_AI_KEY is required for streaming integration tests.") + url = "http://localhost:8108/collections" + headers = {"X-TYPESENSE-API-KEY": "xyz"} + collection_data = { + "name": "streaming_docs", + "fields": [ + { + "name": "title", + "type": "string", + }, + { + "name": "embedding", + "type": "float[]", + "embed": { + "from": ["title"], + "model_config": { + "model_name": "openai/text-embedding-3-small", + "api_key": open_ai_key, + }, + }, + }, + ], + } + + response = requests.post(url, headers=headers, json=collection_data, timeout=3) + response.raise_for_status() + return "streaming_docs" + + +@pytest.fixture(name="create_streaming_document") +def create_streaming_document_fixture(create_streaming_collection: str) -> str: + """Create a document for streaming tests.""" + url = "http://localhost:8108/collections/streaming_docs/documents" + headers = {"X-TYPESENSE-API-KEY": "xyz"} + document_data = { + "id": "stream-1", + "title": "Company profile", + } + + response = requests.post(url, headers=headers, json=document_data, timeout=3) + response.raise_for_status() + return "stream-1" From 4962bef1648b48e3eeeed77733c019460b52b3fb Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:49:09 +0200 Subject: [PATCH 05/16] fix(utils): fix prefix in sync client --- utils/run-unasync.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/utils/run-unasync.py b/utils/run-unasync.py index 5dd8816..7b11b6c 100644 --- a/utils/run-unasync.py +++ b/utils/run-unasync.py @@ -31,6 +31,8 @@ def collect_class_replacements(source_dir: Path) -> dict[str, str]: replacements["AsyncConcurrencyLimit"] = "ConcurrencyLimit" # Defined in the shared ``typesense.http_backend`` module, outside async_. replacements["ASYNC_CLIENT_TYPES"] = "CLIENT_TYPES" + replacements["aiter_lines"] = "iter_lines" + replacements["aread"] = "read" return replacements From 542b146d8d65e6ccd9e88c13cadf2d4c12415a70 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:47:44 +0200 Subject: [PATCH 06/16] feat(sync): generate sync client for streaming responses --- src/typesense/sync/api_call.py | 103 ++++++++++++++++++++++++++++++++ src/typesense/sync/documents.py | 12 +++- 2 files changed, 114 insertions(+), 1 deletion(-) diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index 0ffd4bc..57ff736 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -32,6 +32,7 @@ """ import time +import json import sys from types import MappingProxyType, TracebackType @@ -59,6 +60,14 @@ ) from typesense.node_manager import NodeManager from typesense.request_handler import RequestHandler +from typesense.stream_handlers import ( + JSONDict, + StreamChunk, + combine_stream_chunks, + is_message_chunk, + parse_sse_line, +) +from typesense.types.document import StreamConfig if sys.version_info >= (3, 11): import typing @@ -225,6 +234,8 @@ def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[False], params: typing.Union[TParams, None] = None, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, ) -> str: """ Execute an async GET request to the Typesense API. @@ -246,6 +257,8 @@ def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[True] = True, params: typing.Union[TParams, None] = None, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, ) -> TEntityDict: """ Execute an async GET request to the Typesense API. @@ -266,6 +279,8 @@ def get( entity_type: typing.Type[TEntityDict], as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, params: typing.Union[TParams, None] = None, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, ) -> typing.Union[TEntityDict, str]: """ Execute an async GET request to the Typesense API. @@ -285,6 +300,8 @@ def get( entity_type, as_json, params=params, + stream_config=stream_config, + is_streaming_request=is_streaming_request, ) @typing.overload @@ -453,6 +470,8 @@ def _execute_request( as_json: typing.Literal[True], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> TEntityDict: """Execute an async request with retry logic.""" @@ -466,6 +485,8 @@ def _execute_request( as_json: typing.Literal[False], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> str: """Execute an async request with retry logic.""" @@ -478,6 +499,8 @@ def _execute_request( as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, + stream_config: StreamConfig[TEntityDict] | None = None, + is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: """ @@ -509,6 +532,10 @@ def _execute_request( node, url, request_kwargs = self._prepare_request_params(endpoint, **kwargs) try: + if is_streaming_request and method == "GET": + return self._handle_streaming_get( + url, entity_type, stream_config, **request_kwargs + ) return self._make_request_and_process_response( method, node, @@ -521,6 +548,13 @@ def _execute_request( raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) + if is_streaming_request and stream_config: + on_error = stream_config.get("on_error") + if on_error: + try: + on_error(server_error) + except Exception: + pass if num_retries < self.config.num_retries: time.sleep(self.config.retry_interval_seconds) return self._execute_request( @@ -530,6 +564,8 @@ def _execute_request( as_json, last_exception=server_error, num_retries=num_retries + 1, + stream_config=stream_config, + is_streaming_request=is_streaming_request, **kwargs, ) @@ -559,6 +595,73 @@ def _make_request_and_process_response( else typing.cast(str, request_response) ) + def _handle_streaming_get( + self, + url: str, + entity_type: typing.Type[TEntityDict], + stream_config: StreamConfig[TEntityDict] | None, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> TEntityDict: + """Perform an async streaming GET, parse SSE lines, invoke callbacks, return combined result.""" + headers: typing.Dict[str, str] = { + self.request_handler.api_key_header_name: self.config.api_key, + "Accept": "text/event-stream", + } + headers.update(self.config.additional_headers) + extra_headers = kwargs.get("headers") + if extra_headers: + headers.update(extra_headers) + + params = kwargs.get("params") + content: typing.Union[str, bytes, None] = None + if body := kwargs.get("data"): + if isinstance(body, (str, bytes)): + content = body + else: + content = json.dumps(body) + + all_chunks: typing.List[StreamChunk] = [] + with self._client.stream( + "GET", + url, + params=params, + content=content, + headers=headers, + timeout=self.config.connection_timeout_seconds, + ) as response: + if response.status_code < 200 or response.status_code >= 300: + response.read() + error_message = self.request_handler._get_error_message(response) + raise self.request_handler._get_exception(response.status_code)( + response.status_code, + error_message, + ) + for line in response.iter_lines(): + chunk = parse_sse_line(line) + if chunk is not None: + all_chunks.append(chunk) + if stream_config and is_message_chunk(chunk): + on_chunk = stream_config.get("on_chunk") + if on_chunk: + try: + on_chunk(chunk) + except Exception: + pass + + self.node_manager.set_node_health( + self.node_manager.get_node(), + is_healthy=True, + ) + final: JSONDict = combine_stream_chunks(all_chunks) + if stream_config: + on_complete = stream_config.get("on_complete") + if on_complete: + try: + on_complete(typing.cast(TEntityDict, final)) + except Exception: + pass + return typing.cast(TEntityDict, final) + def _prepare_request_params( self, endpoint: str, diff --git a/src/typesense/sync/documents.py b/src/typesense/sync/documents.py index 0c7d7f7..e9225f6 100644 --- a/src/typesense/sync/documents.py +++ b/src/typesense/sync/documents.py @@ -43,6 +43,7 @@ ImportResponseWithId, SearchParameters, SearchResponse, + StreamConfigBuilder, UpdateByFilterParameters, UpdateByFilterResponse, ) @@ -362,16 +363,25 @@ def search(self, search_parameters: SearchParameters) -> SearchResponse[TDoc]: Args: search_parameters (SearchParameters): The search parameters. + Use conversation_stream=True and optionally stream_config (on_chunk, + on_complete, on_error) for conversational search streaming. Returns: SearchResponse[TDoc]: The search response containing matching documents. """ - stringified_search_params = stringify_search_params(search_parameters) + params_for_api = dict(search_parameters) + stream_config = params_for_api.pop("stream_config", None) + if isinstance(stream_config, StreamConfigBuilder): + stream_config = stream_config.build() + conversation_stream = params_for_api.get("conversation_stream") is True + stringified_search_params = stringify_search_params(params_for_api) response: SearchResponse[TDoc] = self.api_call.get( self._endpoint_path("search"), params=stringified_search_params, entity_type=SearchResponse, as_json=True, + stream_config=stream_config, + is_streaming_request=conversation_stream, ) return response From 550fde758c6bfc474f5ff706376a331f44a88372 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:48:28 +0200 Subject: [PATCH 07/16] test(streaming): add tests for streaming responses - test both async and sync version of the client - add unit tests and tests against real typesense instance --- tests/streaming_async_test.py | 454 ++++++++++++++++++++++++++++++++++ tests/streaming_test.py | 414 +++++++++++++++++++++++++++++++ 2 files changed, 868 insertions(+) create mode 100644 tests/streaming_async_test.py create mode 100644 tests/streaming_test.py diff --git a/tests/streaming_async_test.py b/tests/streaming_async_test.py new file mode 100644 index 0000000..2882c22 --- /dev/null +++ b/tests/streaming_async_test.py @@ -0,0 +1,454 @@ +"""Async streaming conversation search tests.""" + +import sys + +import pytest + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +from tests.fixtures.streaming_fixtures import ( + FakeAsyncStreamContext, + FakeAsyncStreamResponse, + JSONValue, +) +from typesense.async_.api_call import AsyncApiCall +from typesense.async_.documents import AsyncDocuments +from typesense.exceptions import ServerError +from typesense.types.document import ( + DocumentSchema, + MessageChunk, + StreamConfig, + StreamConfigBuilder, +) + + +async def test_streaming_search_invokes_on_chunk_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test that streaming search invokes on_chunk for each message chunk.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + stream_config: StreamConfig[DocumentSchema] = {"on_chunk": on_chunk} + + sse_lines = [ + 'data: {"conversation_id":"123","message":"First chunk"}', + 'data: {"conversation_id":"123","message":"Second chunk"}', + '{"found": 2, "hits": [], "page": 1, "search_time_ms": 10}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + result = await fake_async_documents.search( + { + "q": "test query", + "query_by": "title", + "conversation_stream": True, + "stream_config": stream_config, + } + ) + + assert len(chunks_received) == 2 + assert chunks_received[0]["message"] == "First chunk" + assert chunks_received[1]["message"] == "Second chunk" + assert result["found"] == 2 + + +async def test_streaming_search_handles_plain_text_lines_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Plain text lines should be treated as message chunks with unknown id.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + "Hello", + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + assert chunks_received[0]["conversation_id"] == "unknown" + assert chunks_received[0]["message"] == "Hello" + + +async def test_streaming_search_handles_missing_fields_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """JSON lines without conversation_id/message should use defaults.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + 'data: {"foo":"bar"}', + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + assert chunks_received[0]["conversation_id"] == "unknown" + assert chunks_received[0]["message"] == "" + + +async def test_streaming_search_skips_done_marker_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """data: [DONE] lines should be ignored.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + 'data: {"conversation_id":"123","message":"Chunk"}', + "data: [DONE]", + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + + +async def test_streaming_search_handles_json_array_lines_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """JSON arrays should be treated as plain text message chunks.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + 'data: ["a", "b"]', + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + assert chunks_received[0]["message"] == '["a", "b"]' + + +async def test_streaming_search_supports_builder_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test StreamConfigBuilder for streaming callbacks.""" + complete_calls: typing.List[int] = [] + + stream = StreamConfigBuilder() + + @stream.on_complete + def on_complete(response: typing.Mapping[str, JSONValue]) -> None: + found = response.get("found") + if isinstance(found, int): + complete_calls.append(found) + + sse_lines = [ + 'data: {"conversation_id":"123","message":"Hello"}', + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": stream, + } + ) + + assert complete_calls == [1] + + +async def test_stream_config_not_sent_to_api_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test that stream_config is removed from API params.""" + captured_params: typing.Dict[str, str] = {} + + sse_lines = [ + '{"found": 0, "hits": [], "page": 1, "search_time_ms": 1}', + ] + response = stream_response_async(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + if params: + captured_params.update(params) + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + stream_config: StreamConfig[DocumentSchema] = {"on_chunk": lambda _: None} + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": stream_config, + } + ) + + assert "stream_config" not in captured_params + assert captured_params.get("conversation_stream") == "true" + + +async def test_streaming_search_invokes_on_error_async( + fake_async_documents: AsyncDocuments[DocumentSchema], + stream_response_async: type[FakeAsyncStreamResponse], + stream_context_async: type[FakeAsyncStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test that streaming search invokes on_error for request failures.""" + errors: typing.List[BaseException] = [] + + def on_error(error: BaseException) -> None: + errors.append(error) + + fake_async_documents.api_call.config.num_retries = 0 + + response = stream_response_async( + lines=[], + status_code=500, + headers={"Content-Type": "application/json"}, + text='{"message": "Server error"}', + ) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeAsyncStreamContext: + return stream_context_async(response) + + monkeypatch.setattr( + fake_async_documents.api_call._client, + "stream", + fake_stream, + ) + + with pytest.raises(ServerError): + await fake_async_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_error": on_error}, + } + ) + + assert len(errors) == 1 + assert isinstance(errors[0], ServerError) + + +@pytest.mark.open_ai +async def test_actual_streaming_search_async( + actual_async_api_call: AsyncApiCall, + create_streaming_collection: str, + create_streaming_document: str, + create_conversations_model: str, +) -> None: + """Integration test against a real Typesense server with conversation streaming.""" + actual_async_documents = AsyncDocuments( + actual_async_api_call, + create_streaming_collection, + ) + chunks_received: typing.List[MessageChunk] = [] + complete_called: typing.List[bool] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + def on_complete(response: typing.Mapping[str, JSONValue]) -> None: + complete_called.append(True) + + response = await actual_async_documents.search( + { + "q": "What is this document about?", + "query_by": "embedding", + "conversation": True, + "conversation_stream": True, + "conversation_model_id": create_conversations_model, + "prefix": False, + "exclude_fields": "embedding", + "stream_config": {"on_chunk": on_chunk, "on_complete": on_complete}, + } + ) + + assert complete_called == [True] + assert len(chunks_received) > 0 + assert "found" in response or "hits" in response diff --git a/tests/streaming_test.py b/tests/streaming_test.py new file mode 100644 index 0000000..44b110b --- /dev/null +++ b/tests/streaming_test.py @@ -0,0 +1,414 @@ +"""Sync streaming conversation search tests.""" + +import sys + +import pytest + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +from tests.fixtures.streaming_fixtures import ( + FakeStreamContext, + FakeStreamResponse, + JSONValue, +) +from typesense.exceptions import ServerError +from typesense.sync.documents import Documents +from typesense.types.document import ( + DocumentSchema, + MessageChunk, + StreamConfig, + StreamConfigBuilder, +) + + +def test_streaming_search_invokes_on_chunk( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test that streaming search invokes on_chunk for each message chunk.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + stream_config: StreamConfig[DocumentSchema] = {"on_chunk": on_chunk} + + sse_lines = [ + 'data: {"conversation_id":"123","message":"First chunk"}', + 'data: {"conversation_id":"123","message":"Second chunk"}', + '{"found": 2, "hits": [], "page": 1, "search_time_ms": 10}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + result = fake_documents.search( + { + "q": "test query", + "query_by": "title", + "conversation_stream": True, + "stream_config": stream_config, + } + ) + + assert len(chunks_received) == 2 + assert chunks_received[0]["message"] == "First chunk" + assert chunks_received[1]["message"] == "Second chunk" + assert result["found"] == 2 + + +def test_streaming_search_handles_plain_text_lines( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Plain text lines should be treated as message chunks with unknown id.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + "Hello", + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + assert chunks_received[0]["conversation_id"] == "unknown" + assert chunks_received[0]["message"] == "Hello" + + +def test_streaming_search_handles_missing_fields( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """JSON lines without conversation_id/message should use defaults.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + 'data: {"foo":"bar"}', + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + assert chunks_received[0]["conversation_id"] == "unknown" + assert chunks_received[0]["message"] == "" + + +def test_streaming_search_skips_done_marker( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """data: [DONE] lines should be ignored.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + 'data: {"conversation_id":"123","message":"Chunk"}', + "data: [DONE]", + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + + +def test_streaming_search_handles_json_array_lines( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """JSON arrays should be treated as plain text message chunks.""" + chunks_received: typing.List[MessageChunk] = [] + + def on_chunk(chunk: MessageChunk) -> None: + chunks_received.append(chunk) + + sse_lines = [ + 'data: ["a", "b"]', + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_chunk": on_chunk}, + } + ) + + assert len(chunks_received) == 1 + assert chunks_received[0]["message"] == '["a", "b"]' + + +def test_streaming_search_supports_builder( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test StreamConfigBuilder for streaming callbacks.""" + complete_calls: typing.List[int] = [] + + stream = StreamConfigBuilder() + + @stream.on_complete + def on_complete(response: typing.Mapping[str, JSONValue]) -> None: + found = response.get("found") + if isinstance(found, int): + complete_calls.append(found) + + sse_lines = [ + 'data: {"conversation_id":"123","message":"Hello"}', + '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": stream, + } + ) + + assert complete_calls == [1] + + +def test_stream_config_not_sent_to_api( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test that stream_config is removed from API params.""" + captured_params: typing.Dict[str, str] = {} + + sse_lines = [ + '{"found": 0, "hits": [], "page": 1, "search_time_ms": 1}', + ] + response = stream_response(lines=sse_lines) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + if params: + captured_params.update(params) + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + stream_config: StreamConfig[DocumentSchema] = {"on_chunk": lambda _: None} + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": stream_config, + } + ) + + assert "stream_config" not in captured_params + assert captured_params.get("conversation_stream") == "true" + + +def test_streaming_search_invokes_on_error( + fake_documents: Documents[DocumentSchema], + stream_response: type[FakeStreamResponse], + stream_context: type[FakeStreamContext], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Test that streaming search invokes on_error for request failures.""" + errors: typing.List[BaseException] = [] + + def on_error(error: BaseException) -> None: + errors.append(error) + + fake_documents.api_call.config.num_retries = 0 + + response = stream_response( + lines=[], + status_code=500, + headers={"Content-Type": "application/json"}, + text='{"message": "Server error"}', + ) + + def fake_stream( + method: str, + url: str, + params: typing.Mapping[str, str] | None = None, + content: str | bytes | None = None, + headers: typing.Mapping[str, str] | None = None, + timeout: float | None = None, + ) -> FakeStreamContext: + return stream_context(response) + + monkeypatch.setattr( + fake_documents.api_call._client, + "stream", + fake_stream, + ) + + with pytest.raises(ServerError): + fake_documents.search( + { + "q": "test", + "query_by": "title", + "conversation_stream": True, + "stream_config": {"on_error": on_error}, + } + ) + + assert len(errors) == 1 + assert isinstance(errors[0], ServerError) From 6a52fcd55801934136ecdcc7705f3711adadcd89 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Fri, 6 Feb 2026 17:48:47 +0200 Subject: [PATCH 08/16] docs(examples): add examples for streaming conversations --- examples/async_conversation_streaming.py | 135 ++++++++++++++++++++++ examples/conversation_streaming.py | 141 +++++++++++++++++++++++ 2 files changed, 276 insertions(+) create mode 100644 examples/async_conversation_streaming.py create mode 100644 examples/conversation_streaming.py diff --git a/examples/async_conversation_streaming.py b/examples/async_conversation_streaming.py new file mode 100644 index 0000000..6098c5a --- /dev/null +++ b/examples/async_conversation_streaming.py @@ -0,0 +1,135 @@ +import asyncio +import os +import sys +import typing +import uuid + +curr_dir = os.path.dirname(os.path.realpath(__file__)) +repo_root = os.path.abspath(os.path.join(curr_dir, os.pardir)) +sys.path.insert(1, os.path.join(repo_root, "src")) + +import typesense + +from typesense.types.document import MessageChunk, StreamConfigBuilder + + +def require_env(name: str) -> str: + value = os.environ.get(name) + if not value: + raise RuntimeError(f"Missing required environment variable: {name}") + return value + + +async def main() -> None: + typesense_api_key = require_env("TYPESENSE_API_KEY") + openai_api_key = require_env("OPENAI_API_KEY") + + run_id = uuid.uuid4().hex + history_collection = f"streaming_history_{run_id}" + documents_collection = f"streaming_docs_{run_id}" + model_id = f"streaming_model_{run_id}" + + client = typesense.AsyncClient( + { + "api_key": typesense_api_key, + "nodes": [ + { + "host": "localhost", + "port": "8108", + "protocol": "http", + } + ], + "connection_timeout_seconds": 10, + } + ) + + try: + await client.collections.create( + { + "name": history_collection, + "fields": [ + {"name": "conversation_id", "type": "string"}, + {"name": "model_id", "type": "string"}, + {"name": "timestamp", "type": "int32"}, + {"name": "role", "type": "string", "index": False}, + {"name": "message", "type": "string", "index": False}, + ], + } + ) + + await client.collections.create( + { + "name": documents_collection, + "fields": [ + {"name": "title", "type": "string"}, + { + "name": "embedding", + "type": "float[]", + "embed": { + "from": ["title"], + "model_config": { + "model_name": "openai/text-embedding-3-small", + "api_key": openai_api_key, + }, + }, + }, + ], + } + ) + + await client.collections[documents_collection].documents.create( + {"id": "stream-1", "title": "Company profile: a developer tools firm."} + ) + await client.collections[documents_collection].documents.create( + {"id": "stream-2", "title": "Internal memo about quarterly planning."} + ) + + conversation_model = await client.conversations_models.create( + { + "id": model_id, + "model_name": "openai/gpt-4o-mini", + "history_collection": history_collection, + "api_key": openai_api_key, + "system_prompt": ( + "You are an assistant for question-answering. " + "Only use the provided context. Add some fluff about you Being an assistant built for Typesense Conversational Search and a brief overview of how it works" + ), + "max_bytes": 16384, + } + ) + + search_parameters = { + "q": "What is this document about?", + "query_by": "embedding", + "exclude_fields": "embedding", + "prefix": False, + "conversation_model_id": conversation_model["id"], + } + documents = client.collections[documents_collection].documents + + @stream.on_chunk + def on_chunk(chunk: MessageChunk) -> None: + print(chunk["message"], end="", flush=True) + + @stream.on_complete + def on_complete(response: dict) -> None: + print("\n---\nComplete response keys:", response.keys()) + + await client.collections["streaming_docs"].documents.search( + { + "q": "What is this document about?", + "query_by": "embedding", + "exclude_fields": "embedding", + "conversation": True, + "prefix": False, + "conversation_stream": True, + "conversation_model_id": conversation_model["id"], + "stream_config": stream, + } + ) + finally: + await client.api_call.aclose() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/conversation_streaming.py b/examples/conversation_streaming.py new file mode 100644 index 0000000..addb9d4 --- /dev/null +++ b/examples/conversation_streaming.py @@ -0,0 +1,141 @@ +from operator import truediv +import os +import sys +import typing +import uuid + +curr_dir = os.path.dirname(os.path.realpath(__file__)) +repo_root = os.path.abspath(os.path.join(curr_dir, os.pardir)) +sys.path.insert(1, os.path.join(repo_root, "src")) + +import typesense + +from typesense.types.document import MessageChunk, StreamConfigBuilder + + +def require_env(name: str) -> str: + value = os.environ.get(name) + if not value: + raise RuntimeError(f"Missing required environment variable: {name}") + return value + + +typesense_api_key = require_env("TYPESENSE_API_KEY") +openai_api_key = require_env("OPENAI_API_KEY") + +run_id = uuid.uuid4().hex +history_collection = f"streaming_history_{run_id}" +documents_collection = f"streaming_docs_{run_id}" +model_id = f"streaming_model_{run_id}" + +client = typesense.Client( + { + "api_key": typesense_api_key, + "nodes": [ + { + "host": "localhost", + "port": "8108", + "protocol": "http", + } + ], + "connection_timeout_seconds": 10, + } +) + +client.collections.create( + { + "name": history_collection, + "fields": [ + {"name": "conversation_id", "type": "string"}, + {"name": "model_id", "type": "string"}, + {"name": "timestamp", "type": "int32"}, + {"name": "role", "type": "string", "index": False}, + {"name": "message", "type": "string", "index": False}, + ], + } +) + +client.collections.create( + { + "name": documents_collection, + "fields": [ + {"name": "title", "type": "string"}, + { + "name": "embedding", + "type": "float[]", + "embed": { + "from": ["title"], + "model_config": { + "model_name": "openai/text-embedding-3-small", + "api_key": openai_api_key, + }, + }, + }, + ], + } +) + +client.collections[documents_collection].documents.create( + {"id": "stream-1", "title": "Company profile: a developer tools firm."} +) +client.collections[documents_collection].documents.create( + {"id": "stream-2", "title": "Internal memo about a quarterly planning meeting."} +) + +conversation_model = client.conversations_models.create( + { + "id": model_id, + "model_name": "openai/gpt-4o-mini", + "history_collection": history_collection, + "api_key": openai_api_key, + "system_prompt": ( + "You are an assistant for question-answering. " + "Only use the provided context. Add some fluff about you Being an assistant built for Typesense Conversational Search and a brief overview of how it works" + ), + "max_bytes": 16384, + } +) + +search_parameters = { + "q": "What is this document about?", + "query_by": "embedding", + "exclude_fields": "embedding", + "prefix": False, + "conversation_model_id": conversation_model["id"], +} + +# Iterate over the answer as it is generated, then read the search response. +with client.collections[documents_collection].documents.search_stream( + search_parameters, +) as answer_stream: + for chunk in answer_stream: + print(chunk["message"], end="", flush=True) + response = answer_stream.get_final_response() +print("\n---\nFound", response["found"], "documents") + +# Or pass callbacks to search(), which returns the search response at the end. +stream_config: StreamConfigBuilder[SearchResponse[typing.Any]] = StreamConfigBuilder() + + +@stream.on_chunk +def on_chunk(chunk: MessageChunk) -> None: + print(chunk["message"], end="", flush=True) + + +@stream.on_complete +def on_complete(response: dict) -> None: + print("\n---\nComplete response keys:", response.keys()) + + +client.collections[documents_collection].documents.search( + { + "q": "What is this document about?", + "query_by": "embedding", + "exclude_fields": "embedding", + "conversation": True, + "prefix": False, + "conversation_stream": True, + "conversation_model_id": conversation_model["id"], + "stream_config": stream, + } +) From 0a0f1a6e289f8fea2087777f0a76533233d69c4d Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:39:58 +0300 Subject: [PATCH 09/16] refactor: share request kwargs and status checks in request handler --- src/typesense/request_handler.py | 92 ++++++++++++++++++++------------ 1 file changed, 59 insertions(+), 33 deletions(-) diff --git a/src/typesense/request_handler.py b/src/typesense/request_handler.py index 91d9bc6..d7bf1eb 100644 --- a/src/typesense/request_handler.py +++ b/src/typesense/request_handler.py @@ -216,27 +216,7 @@ def make_request( Raises: TypesenseClientError: If the API returns an error response. """ - headers = { - self.api_key_header_name: self.config.api_key, - } - headers.update(self.config.additional_headers) - - request_kwargs: SessionFunctionKwargs[TParams, TBody] = typing.cast( - SessionFunctionKwargs[TParams, TBody], - { - "headers": headers, - "timeout": self.config.connection_timeout_seconds, - }, - ) - - if params := kwargs.get("params"): - self.normalize_params(params) - request_kwargs["params"] = params - - if body := kwargs.get("data"): - request_kwargs["content"] = ( - body if isinstance(body, (str, bytes)) else json.dumps(body) - ) + request_kwargs = self.build_request_kwargs(**kwargs) if isinstance(client, ASYNC_CLIENT_TYPES): return self._make_async_request( @@ -270,12 +250,7 @@ def _make_sync_request( headers=headers, ) - if response.status_code < 200 or response.status_code >= 300: - error_message = self._get_error_message(response) - raise self._get_exception(response.status_code)( - response.status_code, - error_message, - ) + self.raise_for_status(response) if as_json: res: TEntityDict = typing.cast(TEntityDict, response.json()) @@ -305,12 +280,7 @@ async def _make_async_request( headers=headers, ) - if response.status_code < 200 or response.status_code >= 300: - error_message = self._get_error_message(response) - raise self._get_exception(response.status_code)( - response.status_code, - error_message, - ) + self.raise_for_status(response) if as_json: res: TEntityDict = typing.cast(TEntityDict, response.json()) @@ -318,6 +288,62 @@ async def _make_async_request( return response.text + def build_request_kwargs( + self, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> SessionFunctionKwargs[TParams, TBody]: + """ + Build the headers, query parameters and body for a request. + + Args: + kwargs: The request's ``params`` and ``data``. + + Returns: + SessionFunctionKwargs: The ``headers``, ``params`` and ``content`` to send. + """ + headers = { + self.api_key_header_name: self.config.api_key, + } + headers.update(self.config.additional_headers) + + request_kwargs: SessionFunctionKwargs[TParams, TBody] = typing.cast( + SessionFunctionKwargs[TParams, TBody], + { + "headers": headers, + "timeout": self.config.connection_timeout_seconds, + }, + ) + + if params := kwargs.get("params"): + self.normalize_params(params) + request_kwargs["params"] = params + + if body := kwargs.get("data"): + request_kwargs["content"] = ( + body if isinstance(body, (str, bytes)) else json.dumps(body) + ) + + return request_kwargs + + def raise_for_status(self, response: ResponseType) -> None: + """ + Raise the client error matching a non-2xx response. + + The response body must already be read. + + Args: + response (httpx.Response | httpx2.Response): The API response. + + Raises: + TypesenseClientError: If the response status is not 2xx. + """ + if response.status_code < 200 or response.status_code >= 300: + error_message = self._get_error_message(response) + raise self._get_exception(response.status_code)( + response.status_code, + error_message, + ) + @staticmethod def normalize_params(params: typing.Mapping[str, object]) -> None: """ From f85b35617afafa77bf9f70a7e70378bb018a4945 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:40:15 +0300 Subject: [PATCH 10/16] feat(config): add stream_read_timeout_seconds for streaming requests --- src/typesense/configuration.py | 15 +++++++++++++++ tests/configuration_test.py | 3 +++ tests/configuration_validations_test.py | 5 +++++ 3 files changed, 23 insertions(+) diff --git a/src/typesense/configuration.py b/src/typesense/configuration.py index 31ba091..34df662 100644 --- a/src/typesense/configuration.py +++ b/src/typesense/configuration.py @@ -102,6 +102,12 @@ class ConfigDict(typing.TypedDict): once; further requests wait for a slot. Keep it below ``max_connections`` so a burst of slow requests cannot exhaust the pool. Defaults to no limit. + + stream_read_timeout_seconds (float): How long a streaming conversation + search waits for the next chunk before raising ``httpx.ReadTimeout``. + Replaces the read timeout for streaming requests only, since the first + chunk arrives only once the LLM starts answering. Defaults to 60, the + server's own limit for an LLM response. """ nodes: typing.List[typing.Union[str, NodeConfigDict]] @@ -124,6 +130,7 @@ class ConfigDict(typing.TypedDict): max_connections: typing.NotRequired[int] max_keepalive_connections: typing.NotRequired[int] max_concurrent_requests: typing.NotRequired[int] + stream_read_timeout_seconds: typing.NotRequired[float] class Node: @@ -216,6 +223,7 @@ class Configuration: max_connections (int): The maximum number of connections in the pool. max_keepalive_connections (int): The maximum number of idle pooled connections. max_concurrent_requests (int | None): The maximum number of requests in flight. + stream_read_timeout_seconds (float): How long a stream waits for its next chunk. """ def __init__( @@ -272,6 +280,10 @@ def __init__( self.max_concurrent_requests: typing.Optional[int] = config_dict.get( "max_concurrent_requests", ) + self.stream_read_timeout_seconds = config_dict.get( + "stream_read_timeout_seconds", + 60.0, + ) def _handle_nearest_node( self, @@ -352,6 +364,9 @@ def validate_connection_pool(config_dict: ConfigDict) -> None: "pool_timeout_seconds": config_dict.get("pool_timeout_seconds"), "max_connections": config_dict.get("max_connections"), "max_concurrent_requests": config_dict.get("max_concurrent_requests"), + "stream_read_timeout_seconds": config_dict.get( + "stream_read_timeout_seconds" + ), } for key, config_value in positive_settings.items(): if config_value is not None and config_value <= 0: diff --git a/tests/configuration_test.py b/tests/configuration_test.py index 092c93b..838400d 100644 --- a/tests/configuration_test.py +++ b/tests/configuration_test.py @@ -224,6 +224,7 @@ def test_configuration_connection_pool_defaults() -> None: "max_connections": 100, "max_keepalive_connections": 20, "max_concurrent_requests": None, + "stream_read_timeout_seconds": 60.0, } assert_to_contain_object(configuration, expected) @@ -239,6 +240,7 @@ def test_configuration_connection_pool_explicit() -> None: "max_connections": 200, "max_keepalive_connections": 50, "max_concurrent_requests": 150, + "stream_read_timeout_seconds": 120.0, }, ) @@ -247,6 +249,7 @@ def test_configuration_connection_pool_explicit() -> None: "max_connections": 200, "max_keepalive_connections": 50, "max_concurrent_requests": 150, + "stream_read_timeout_seconds": 120.0, } assert_to_contain_object(configuration, expected) diff --git a/tests/configuration_validations_test.py b/tests/configuration_validations_test.py index 8cf8061..e8fa9c2 100644 --- a/tests/configuration_validations_test.py +++ b/tests/configuration_validations_test.py @@ -217,6 +217,11 @@ def test_validate_config_dict_with_wrong_nearest_node() -> None: -1, "`max_concurrent_requests` must be greater than 0.", ), + ( + "stream_read_timeout_seconds", + 0, + "`stream_read_timeout_seconds` must be greater than 0.", + ), ( "max_keepalive_connections", -1, From 7975c39b01c4f23c7affc0dd65111b42460b5494 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:41:34 +0300 Subject: [PATCH 11/16] feat(streaming): add spec-compliant sse decoder --- src/typesense/sse.py | 208 +++++++++++++++++++++++++++++++++++++++++++ tests/sse_test.py | 117 ++++++++++++++++++++++++ 2 files changed, 325 insertions(+) create mode 100644 src/typesense/sse.py create mode 100644 tests/sse_test.py diff --git a/src/typesense/sse.py b/src/typesense/sse.py new file mode 100644 index 0000000..337eed4 --- /dev/null +++ b/src/typesense/sse.py @@ -0,0 +1,208 @@ +""" +Server-sent events (SSE) parsing for streaming responses. + +Typesense streams conversational search answers as ``text/event-stream``. This +module turns the raw response bytes into ``ServerSentEvent`` objects, following the +WHATWG parsing rules: + +- Lines end with CRLF, LF or a lone CR, even when a CRLF pair is split across + two network reads. +- Lines are split before decoding, so a multi-byte UTF-8 character split across + reads stays intact, and characters like U+2028 inside JSON never break a line. + (httpx's ``iter_lines`` splits on those, so it is not used here.) +- Multiple ``data:`` lines in one event are joined with ``\\n``. +- Lines starting with ``:`` are comments. +- An event is dispatched on a blank line; a trailing event without one is dropped. + +``iter_events`` and ``aiter_events`` are the sync and async entry points +(``utils/run-unasync.py`` maps one name to the other). +""" + +import json +import re +import sys + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +_LINE_END = re.compile(rb"\r\n|\r|\n") +_BOM = "" + + +class ServerSentEvent: + """A single dispatched server-sent event.""" + + def __init__( + self, + *, + event: str = "message", + data: str = "", + id: str = "", # noqa: A002 (the SSE field name) + retry: typing.Optional[int] = None, + ) -> None: + """ + Initialize the event. + + Args: + event (str): The event type. Defaults to ``message``. + data (str): The event data, with multiple ``data:`` lines joined by ``\\n``. + id (str): The last event ID seen on the stream. + retry (int | None): The reconnection time sent with the event, if any. + """ + self.event = event + self.data = data + self.id = id + self.retry = retry + + def json(self) -> typing.Any: + """Parse the event data as JSON.""" + return json.loads(self.data) + + def __repr__(self) -> str: + """Return a debug representation of the event.""" + return ( + f"ServerSentEvent(event={self.event!r}, data={self.data!r}, " + f"id={self.id!r}, retry={self.retry!r})" + ) + + def __eq__(self, other: object) -> bool: + """Compare two events field by field.""" + if not isinstance(other, ServerSentEvent): + return NotImplemented + return (self.event, self.data, self.id, self.retry) == ( + other.event, + other.data, + other.id, + other.retry, + ) + + +class SSEDecoder: + """Incremental decoder that turns response bytes into server-sent events.""" + + def __init__(self) -> None: + """Initialize an empty decoder.""" + self._buffer = b"" + # The previous chunk ended in ``\r``; a leading ``\n`` belongs to that line end. + self._pending_cr = False + self._at_start = True + self._event = "" + self._data: typing.List[str] = [] + self._last_event_id = "" + self._retry: typing.Optional[int] = None + + @property + def remainder(self) -> str: + """Return the bytes after the last line end, decoded, once the stream is over.""" + return self._buffer.decode("utf-8", errors="replace") + + def feed(self, chunk: bytes) -> typing.List[ServerSentEvent]: + """ + Decode a chunk of the response body. + + Args: + chunk (bytes): The next bytes read from the response. + + Returns: + List[ServerSentEvent]: The events completed by this chunk. + """ + if not chunk: + return [] + if self._pending_cr and chunk.startswith(b"\n"): + chunk = chunk[1:] + self._pending_cr = False + + buffer = self._buffer + chunk + events: typing.List[ServerSentEvent] = [] + start = 0 + for line_end in _LINE_END.finditer(buffer): + if line_end.group() == b"\r" and line_end.end() == len(buffer): + self._pending_cr = True + event = self._process_line(buffer[start : line_end.start()]) + if event is not None: + events.append(event) + start = line_end.end() + self._buffer = buffer[start:] + return events + + def _process_line(self, raw_line: bytes) -> typing.Optional[ServerSentEvent]: + """Apply one line to the pending event, returning the event on a blank line.""" + line = raw_line.decode("utf-8", errors="replace") + if self._at_start: + self._at_start = False + line = line[len(_BOM) :] if line.startswith(_BOM) else line + + if not line: + return self._dispatch() + if line.startswith(":"): + return None + + field, _, field_value = line.partition(":") + if field_value.startswith(" "): + field_value = field_value[1:] + + if field == "event": + self._event = field_value + elif field == "data": + self._data.append(field_value) + elif field == "id": + if "\0" not in field_value: + self._last_event_id = field_value + elif field == "retry": + if field_value.isascii() and field_value.isdigit(): + self._retry = int(field_value) + return None + + def _dispatch(self) -> typing.Optional[ServerSentEvent]: + """Build the pending event and reset the per-event fields.""" + event: typing.Optional[ServerSentEvent] = None + if self._data: + event = ServerSentEvent( + event=self._event or "message", + data="\n".join(self._data), + id=self._last_event_id, + retry=self._retry, + ) + self._event = "" + self._data = [] + self._retry = None + return event + + +def iter_events( + chunks: typing.Iterable[bytes], + decoder: SSEDecoder, +) -> typing.Iterator[ServerSentEvent]: + """ + Yield the server-sent events in a stream of response bytes. + + Args: + chunks (Iterable[bytes]): The response body, e.g. ``response.iter_bytes()``. + decoder (SSEDecoder): The decoder holding the parsing state. + + Yields: + ServerSentEvent: Each event, as soon as its blank line arrives. + """ + for chunk in chunks: + yield from decoder.feed(chunk) + + +async def aiter_events( + chunks: typing.AsyncIterable[bytes], + decoder: SSEDecoder, +) -> typing.AsyncIterator[ServerSentEvent]: + """ + Yield the server-sent events in an async stream of response bytes. + + Args: + chunks (AsyncIterable[bytes]): The response body, e.g. ``response.aiter_bytes()``. + decoder (SSEDecoder): The decoder holding the parsing state. + + Yields: + ServerSentEvent: Each event, as soon as its blank line arrives. + """ + async for chunk in chunks: + for event in decoder.feed(chunk): + yield event diff --git a/tests/sse_test.py b/tests/sse_test.py new file mode 100644 index 0000000..8dedbb6 --- /dev/null +++ b/tests/sse_test.py @@ -0,0 +1,117 @@ +"""Tests for the server-sent events decoder.""" + +import typing + +import pytest + +from typesense.sse import SSEDecoder, ServerSentEvent, aiter_events, iter_events + + +def decode(*chunks: bytes) -> list[ServerSentEvent]: + """Decode the chunks with a fresh decoder.""" + return list(iter_events(chunks, SSEDecoder())) + + +@pytest.mark.parametrize("line_end", [b"\n", b"\r\n", b"\r"]) +def test_line_endings(line_end: bytes) -> None: + """Test that LF, CRLF and a lone CR all end a line.""" + body = b"data: one" + line_end + line_end + b"data: two" + line_end + line_end + + assert [event.data for event in decode(body)] == ["one", "two"] + + +def test_crlf_split_across_chunks() -> None: + """Test that a CRLF split across two reads is a single line end.""" + events = decode(b"data: one\r", b"\n\r", b"\ndata: two\r\n\r\n") + + assert [event.data for event in events] == ["one", "two"] + + +def test_multibyte_character_split_across_chunks() -> None: + """Test that a UTF-8 character split across two reads is decoded intact.""" + body = 'data: {"message": "καλημέρα"}\n\n'.encode() + split_at = body.index("μ".encode()) + 1 + + events = decode(body[:split_at], body[split_at:]) + + assert events[0].json() == {"message": "καλημέρα"} + + +def test_unicode_line_separator_inside_data() -> None: + """Test that U+2028 inside JSON does not split the line.""" + events = decode('data: {"message": "a
b"}\n\n'.encode()) + + assert events[0].json() == {"message": "a
b"} + + +def test_multiline_data_is_joined_with_newlines() -> None: + """Test that the data lines of one event are joined with a newline.""" + events = decode(b"data: first\ndata:second\n\n") + + assert events[0].data == "first\nsecond" + + +def test_only_one_leading_space_is_stripped() -> None: + """Test that a single space after the colon is removed, but not more.""" + assert decode(b"data: padded\n\n")[0].data == " padded" + + +def test_comments_and_unknown_fields_are_ignored() -> None: + """Test that comment lines and unknown fields do not create events.""" + events = decode(b": keep-alive\n\nfoo: bar\ndata: real\n\n") + + assert [event.data for event in events] == ["real"] + + +def test_event_id_and_retry_fields() -> None: + """Test that event, id and retry are parsed, and id carries over.""" + events = decode( + b"event: delta\nid: 7\nretry: 1500\ndata: a\n\ndata: b\n\nretry: x1\ndata: c\n\n", + ) + + assert events == [ + ServerSentEvent(event="delta", data="a", id="7", retry=1500), + ServerSentEvent(event="message", data="b", id="7"), + ServerSentEvent(event="message", data="c", id="7"), + ] + + +def test_leading_bom_is_stripped() -> None: + """Test that a UTF-8 byte order mark at the start of the stream is ignored.""" + assert decode(b"\xef\xbb\xbfdata: x\n\n")[0].data == "x" + + +def test_trailing_event_without_blank_line_is_dropped() -> None: + """Test that an unterminated event is not dispatched, and its text is kept.""" + decoder = SSEDecoder() + + events = list(iter_events([b"data: one\n\n", b'{"message": "boom"}'], decoder)) + + assert [event.data for event in events] == ["one"] + assert decoder.remainder == '{"message": "boom"}' + + +def test_event_without_data_is_not_dispatched() -> None: + """Test that a blank line after only non-data fields dispatches nothing.""" + assert decode(b"event: ping\n\n") == [] + + +def test_byte_at_a_time() -> None: + """Test decoding when every read returns a single byte.""" + body = b'data: {"message": "hi"}\r\n\r\ndata: [DONE]\r\n\r\n' + + events = decode(*(body[index : index + 1] for index in range(len(body)))) + + assert [event.data for event in events] == ['{"message": "hi"}', "[DONE]"] + + +async def test_aiter_events() -> None: + """Test the async entry point.""" + + async def chunks() -> typing.AsyncIterator[bytes]: + for chunk in (b"data: a\n", b"\ndata: b\n\n"): + yield chunk + + events = [event.data async for event in aiter_events(chunks(), SSEDecoder())] + + assert events == ["a", "b"] From 132c1419efb018eabc7eebace465b506376e5c9a Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:56:58 +0300 Subject: [PATCH 12/16] feat(streaming): let streams hold a concurrency slot until closed --- src/typesense/concurrency_limit.py | 32 ++++++++++++++++++++++-------- 1 file changed, 24 insertions(+), 8 deletions(-) diff --git a/src/typesense/concurrency_limit.py b/src/typesense/concurrency_limit.py index 34ebdc5..76a115c 100644 --- a/src/typesense/concurrency_limit.py +++ b/src/typesense/concurrency_limit.py @@ -37,14 +37,23 @@ def __init__(self, max_concurrent_requests: typing.Optional[int]) -> None: # semaphore binds to the loop that is current when it is constructed. self._semaphore: typing.Optional[asyncio.Semaphore] = None - async def __aenter__(self) -> None: - """Wait for a free slot.""" + async def acquire(self) -> None: + """Wait for a free slot. Streams hold it until ``release`` is called.""" if self._max_concurrent_requests is None: return if self._semaphore is None: self._semaphore = asyncio.Semaphore(self._max_concurrent_requests) await self._semaphore.acquire() + def release(self) -> None: + """Release a slot taken with ``acquire``.""" + if self._semaphore is not None: + self._semaphore.release() + + async def __aenter__(self) -> None: + """Wait for a free slot.""" + await self.acquire() + async def __aexit__( self, exc_type: typing.Optional[typing.Type[BaseException]], @@ -52,8 +61,7 @@ async def __aexit__( exc_tb: typing.Optional[TracebackType], ) -> None: """Release the slot.""" - if self._semaphore is not None: - self._semaphore.release() + self.release() class ConcurrencyLimit: @@ -73,11 +81,20 @@ def __init__(self, max_concurrent_requests: typing.Optional[int]) -> None: else threading.Semaphore(max_concurrent_requests) ) - def __enter__(self) -> None: - """Wait for a free slot.""" + def acquire(self) -> None: + """Wait for a free slot. Streams hold it until ``release`` is called.""" if self._semaphore is not None: self._semaphore.acquire() + def release(self) -> None: + """Release a slot taken with ``acquire``.""" + if self._semaphore is not None: + self._semaphore.release() + + def __enter__(self) -> None: + """Wait for a free slot.""" + self.acquire() + def __exit__( self, exc_type: typing.Optional[typing.Type[BaseException]], @@ -85,5 +102,4 @@ def __exit__( exc_tb: typing.Optional[TracebackType], ) -> None: """Release the slot.""" - if self._semaphore is not None: - self._semaphore.release() + self.release() From 58167d9588a3fb3880cd520a193872b3406f4453 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:56:58 +0300 Subject: [PATCH 13/16] feat(streaming): rework streaming as search_stream and perform_stream iterators --- src/typesense/async_/api_call.py | 209 ++++++------ src/typesense/async_/documents.py | 83 ++++- src/typesense/async_/multi_search.py | 97 +++++- src/typesense/async_/stream.py | 220 +++++++++++++ src/typesense/stream_handlers.py | 164 ---------- src/typesense/sync/api_call.py | 209 ++++++------ src/typesense/sync/documents.py | 83 ++++- src/typesense/sync/multi_search.py | 97 +++++- src/typesense/sync/stream.py | 220 +++++++++++++ src/typesense/types/document.py | 194 +++++------- src/typesense/types/multi_search.py | 8 +- tests/fixtures/streaming_fixtures.py | 180 ----------- tests/streaming_async_test.py | 454 --------------------------- tests/streaming_test.py | 414 ------------------------ utils/run-unasync.py | 7 +- 15 files changed, 1074 insertions(+), 1565 deletions(-) create mode 100644 src/typesense/async_/stream.py delete mode 100644 src/typesense/stream_handlers.py create mode 100644 src/typesense/sync/stream.py delete mode 100644 tests/fixtures/streaming_fixtures.py delete mode 100644 tests/streaming_async_test.py delete mode 100644 tests/streaming_test.py diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index f4b1063..f2e27b2 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -32,7 +32,6 @@ """ import asyncio -import json import sys from types import MappingProxyType, TracebackType @@ -58,16 +57,9 @@ backend_errors, verify_option, ) +from .stream import AsyncSearchStream from typesense.node_manager import NodeManager -from typesense.request_handler import RequestHandler -from typesense.stream_handlers import ( - JSONDict, - StreamChunk, - combine_stream_chunks, - is_message_chunk, - parse_sse_line, -) -from typesense.types.document import StreamConfig +from typesense.request_handler import RequestHandler, _QueryParams if sys.version_info >= (3, 11): import typing @@ -234,8 +226,6 @@ async def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[False], params: typing.Union[TParams, None] = None, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, ) -> str: """ Execute an async GET request to the Typesense API. @@ -257,8 +247,6 @@ async def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[True] = True, params: typing.Union[TParams, None] = None, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, ) -> TEntityDict: """ Execute an async GET request to the Typesense API. @@ -279,8 +267,6 @@ async def get( entity_type: typing.Type[TEntityDict], as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, params: typing.Union[TParams, None] = None, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, ) -> typing.Union[TEntityDict, str]: """ Execute an async GET request to the Typesense API. @@ -300,8 +286,6 @@ async def get( entity_type, as_json, params=params, - stream_config=stream_config, - is_streaming_request=is_streaming_request, ) @typing.overload @@ -470,8 +454,6 @@ async def _execute_request( as_json: typing.Literal[True], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> TEntityDict: """Execute an async request with retry logic.""" @@ -485,8 +467,6 @@ async def _execute_request( as_json: typing.Literal[False], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> str: """Execute an async request with retry logic.""" @@ -499,8 +479,6 @@ async def _execute_request( as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: """ @@ -532,10 +510,6 @@ async def _execute_request( node, url, request_kwargs = self._prepare_request_params(endpoint, **kwargs) try: - if is_streaming_request and method == "GET": - return await self._handle_streaming_get( - url, entity_type, stream_config, **request_kwargs - ) return await self._make_request_and_process_response( method, node, @@ -548,13 +522,6 @@ async def _execute_request( raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) - if is_streaming_request and stream_config: - on_error = stream_config.get("on_error") - if on_error: - try: - on_error(server_error) - except Exception: - pass if num_retries < self.config.num_retries: await asyncio.sleep(self.config.retry_interval_seconds) return await self._execute_request( @@ -564,8 +531,6 @@ async def _execute_request( as_json, last_exception=server_error, num_retries=num_retries + 1, - stream_config=stream_config, - is_streaming_request=is_streaming_request, **kwargs, ) @@ -595,72 +560,126 @@ async def _make_request_and_process_response( else typing.cast(str, request_response) ) - async def _handle_streaming_get( + async def stream( + self, + method: str, + endpoint: str, + entity_type: typing.Type[TEntityDict], + params: typing.Union[TParams, None] = None, + body: typing.Union[TBody, None] = None, + ) -> AsyncSearchStream[TEntityDict]: + """ + Open a streaming request to the Typesense API. + + Failing nodes are retried like any other request until the response + headers arrive. Errors after that are raised while reading the stream and + are not retried, since part of the answer has already been read. + + Args: + method (str): The HTTP method to use. + endpoint (str): The API endpoint to call. + entity_type (Type[TEntityDict]): The type of the final response. + params (Union[TParams, None], optional): Query parameters for the request. + body (Union[TBody, None], optional): The request body. + + Returns: + AsyncSearchStream[TEntityDict]: The open stream. + """ + return await self._execute_stream_request( + method, + endpoint, + entity_type, + params=params, + data=body, + ) + + async def _execute_stream_request( self, + method: str, + endpoint: str, + entity_type: typing.Type[TEntityDict], + last_exception: typing.Union[None, Exception] = None, + num_retries: int = 0, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> AsyncSearchStream[TEntityDict]: + """Open a streaming request, failing over to other nodes like ``_execute_request``.""" + if num_retries > self.config.num_retries: + if last_exception: + raise last_exception + raise TypesenseClientError("All nodes are unhealthy") + + node, url, request_kwargs = self._prepare_request_params(endpoint, **kwargs) + + try: + return await self._open_stream( + method, node, url, entity_type, **request_kwargs + ) + except _CLIENT_ERRORS: + raise + except _SERVER_ERRORS as server_error: + self.node_manager.set_node_health(node, is_healthy=False) + if num_retries < self.config.num_retries: + await asyncio.sleep(self.config.retry_interval_seconds) + return await self._execute_stream_request( + method, + endpoint, + entity_type, + last_exception=server_error, + num_retries=num_retries + 1, + **kwargs, + ) + + async def _open_stream( + self, + method: str, + node: Node, url: str, entity_type: typing.Type[TEntityDict], - stream_config: StreamConfig[TEntityDict] | None, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], - ) -> TEntityDict: - """Perform an async streaming GET, parse SSE lines, invoke callbacks, return combined result.""" - headers: typing.Dict[str, str] = { - self.request_handler.api_key_header_name: self.config.api_key, - "Accept": "text/event-stream", - } - headers.update(self.config.additional_headers) - extra_headers = kwargs.get("headers") - if extra_headers: - headers.update(extra_headers) - - params = kwargs.get("params") - content: typing.Union[str, bytes, None] = None - if body := kwargs.get("data"): - if isinstance(body, (str, bytes)): - content = body - else: - content = json.dumps(body) - - all_chunks: typing.List[StreamChunk] = [] - async with self._client.stream( - "GET", + ) -> AsyncSearchStream[TEntityDict]: + """ + Send a streaming request to `node` and return the stream once headers arrive. + + The stream holds a concurrency slot until it is closed. Reads use + ``stream_read_timeout_seconds``, since the first piece of an answer only + arrives once the LLM starts generating it. + """ + request_kwargs = self.request_handler.build_request_kwargs(**kwargs) + headers = request_kwargs.get("headers", {}) + headers["Accept"] = "text/event-stream" + timeout = self._client.timeout + request = self._client.build_request( + method, url, - params=params, - content=content, + params=typing.cast( + typing.Optional[_QueryParams], + request_kwargs.get("params"), + ), + content=request_kwargs.get("content"), headers=headers, - timeout=self.config.connection_timeout_seconds, - ) as response: - if response.status_code < 200 or response.status_code >= 300: - await response.aread() - error_message = self.request_handler._get_error_message(response) - raise self.request_handler._get_exception(response.status_code)( - response.status_code, - error_message, - ) - async for line in response.aiter_lines(): - chunk = parse_sse_line(line) - if chunk is not None: - all_chunks.append(chunk) - if stream_config and is_message_chunk(chunk): - on_chunk = stream_config.get("on_chunk") - if on_chunk: - try: - on_chunk(chunk) - except Exception: - pass - - self.node_manager.set_node_health( - self.node_manager.get_node(), - is_healthy=True, + timeout=( + timeout.connect, + self.config.stream_read_timeout_seconds, + timeout.write, + timeout.pool, + ), ) - final: JSONDict = combine_stream_chunks(all_chunks) - if stream_config: - on_complete = stream_config.get("on_complete") - if on_complete: + + await self._concurrency_limit.acquire() + try: + response = await self._client.send(request, stream=True) + if response.status_code < 200 or response.status_code >= 300: try: - on_complete(typing.cast(TEntityDict, final)) - except Exception: - pass - return typing.cast(TEntityDict, final) + await response.aread() + finally: + await response.aclose() + self.request_handler.raise_for_status(response) + except BaseException: + self._concurrency_limit.release() + raise + + self.node_manager.set_node_health(node, is_healthy=True) + return AsyncSearchStream(response, self._concurrency_limit.release) def _prepare_request_params( self, diff --git a/src/typesense/async_/documents.py b/src/typesense/async_/documents.py index c43ac35..2b91cd5 100644 --- a/src/typesense/async_/documents.py +++ b/src/typesense/async_/documents.py @@ -21,6 +21,12 @@ from .api_call import AsyncApiCall from .document import AsyncDocument +from .stream import ( + AsyncSearchStream, + consume_stream, + notify_error, + resolve_stream_config, +) from typesense.exceptions import TypesenseClientError from typesense.logger import logger from typesense.preprocess import stringify_search_params @@ -43,7 +49,6 @@ ImportResponseWithId, SearchParameters, SearchResponse, - StreamConfigBuilder, UpdateByFilterParameters, UpdateByFilterResponse, ) @@ -361,30 +366,86 @@ async def search(self, search_parameters: SearchParameters) -> SearchResponse[TD """ Search for documents in the collection. + With ``conversation_stream`` enabled, the LLM's answer is streamed and the + callbacks in ``stream_config`` run as it arrives. To iterate over the answer + instead, use ``search_stream``. + Args: search_parameters (SearchParameters): The search parameters. - Use conversation_stream=True and optionally stream_config (on_chunk, - on_complete, on_error) for conversational search streaming. Returns: SearchResponse[TDoc]: The search response containing matching documents. """ - params_for_api = dict(search_parameters) - stream_config = params_for_api.pop("stream_config", None) - if isinstance(stream_config, StreamConfigBuilder): - stream_config = stream_config.build() - conversation_stream = params_for_api.get("conversation_stream") is True - stringified_search_params = stringify_search_params(params_for_api) + if search_parameters.get("conversation_stream"): + stream_config = resolve_stream_config( + search_parameters.get("stream_config"), + ) + try: + search_stream = await self._open_search_stream(search_parameters) + streamed_response: SearchResponse[TDoc] = await consume_stream( + search_stream, + stream_config, + ) + except Exception as error: + notify_error(stream_config, error) + raise + return streamed_response + + stringified_search_params = stringify_search_params( + { + param: param_value + for param, param_value in search_parameters.items() + if param != "stream_config" + }, + ) response: SearchResponse[TDoc] = await self.api_call.get( self._endpoint_path("search"), params=stringified_search_params, entity_type=SearchResponse, as_json=True, - stream_config=stream_config, - is_streaming_request=conversation_stream, ) return response + async def search_stream( + self, + search_parameters: SearchParameters, + ) -> AsyncSearchStream[SearchResponse[TDoc]]: + """ + Search, streaming the LLM's answer as it is generated. + + Iterate over the returned stream for the pieces of the answer, then call + its ``get_final_response`` for the search response. Use the stream as a + context manager so the connection is released if you stop early. + + ``conversation`` and ``conversation_stream`` are enabled for you; pass the + ``conversation_model_id`` to answer with. ``stream_config`` is ignored. + + Args: + search_parameters (SearchParameters): The search parameters. + + Returns: + AsyncSearchStream[SearchResponse[TDoc]]: The open stream. + """ + return await self._open_search_stream(search_parameters) + + async def _open_search_stream( + self, + search_parameters: SearchParameters, + ) -> AsyncSearchStream[SearchResponse[TDoc]]: + """Open the search request as a stream.""" + stream_params: typing.Dict[str, object] = { + "conversation": True, + **search_parameters, + "conversation_stream": True, + } + stream_params.pop("stream_config", None) + return await self.api_call.stream( + "GET", + self._endpoint_path("search"), + entity_type=SearchResponse, + params=stringify_search_params(stream_params), + ) + async def delete( self, delete_parameters: typing.Union[DeleteQueryParameters, None] = None, diff --git a/src/typesense/async_/multi_search.py b/src/typesense/async_/multi_search.py index 466ac51..d71c448 100644 --- a/src/typesense/async_/multi_search.py +++ b/src/typesense/async_/multi_search.py @@ -19,6 +19,12 @@ import sys from .api_call import AsyncApiCall +from .stream import ( + AsyncSearchStream, + consume_stream, + notify_error, + resolve_stream_config, +) from typesense.preprocess import stringify_search_params from typesense.types.document import MultiSearchCommonParameters from typesense.types.multi_search import MultiSearchRequestSchema, MultiSearchResponse @@ -89,20 +95,93 @@ async def perform( ... ], ... } ... ) + + With ``conversation_stream`` enabled in ``common_params``, the LLM's answer + is streamed and the callbacks in ``stream_config`` run as it arrives. To + iterate over the answer instead, use ``perform_stream``. + """ + if common_params and common_params.get("conversation_stream"): + stream_config = resolve_stream_config(common_params.get("stream_config")) + try: + search_stream = await self.perform_stream(search_queries, common_params) + streamed_response: MultiSearchResponse = await consume_stream( + search_stream, + stream_config, + ) + except Exception as error: + notify_error(stream_config, error) + raise + return streamed_response + + response: MultiSearchResponse = await self.api_call.post( + AsyncMultiSearch.resource_path, + body=self._search_body(search_queries), + params=_without_stream_config(common_params) if common_params else None, + as_json=True, + entity_type=MultiSearchResponse, + ) + return response + + async def perform_stream( + self, + search_queries: MultiSearchRequestSchema, + common_params: typing.Union[MultiSearchCommonParameters, None] = None, + ) -> AsyncSearchStream[MultiSearchResponse]: """ + Perform a multi-search, streaming the LLM's answer as it is generated. + + The searches' hits are combined into one context for a single answer, sent + in the response's top-level ``conversation``. Iterate over the returned + stream for the pieces of the answer, then call its ``get_final_response`` + for the multi-search response. Use the stream as a context manager so the + connection is released if you stop early. + + ``conversation`` and ``conversation_stream`` are enabled for you; pass + ``q`` and the ``conversation_model_id`` in ``common_params``, since + Typesense reads them from the query string. ``stream_config`` is ignored. + + Args: + search_queries (MultiSearchRequestSchema): The searches to perform. + common_params (Union[MultiSearchCommonParameters, None], optional): + Parameters for every search, including the conversation parameters. + + Returns: + AsyncSearchStream[MultiSearchResponse]: The open stream. + """ + stream_params: typing.Dict[str, object] = { + "conversation": True, + **_without_stream_config(common_params or {}), + "conversation_stream": True, + } + return await self.api_call.stream( + "POST", + AsyncMultiSearch.resource_path, + entity_type=MultiSearchResponse, + params=stream_params, + body=self._search_body(search_queries), + ) + + @staticmethod + def _search_body( + search_queries: MultiSearchRequestSchema, + ) -> typing.Dict[str, object]: + """Build the request body, with every search's parameters stringified.""" stringified_search_params = [ stringify_search_params(search_params) for search_params in search_queries.get("searches") ] - search_body = { + return { "searches": stringified_search_params, "union": search_queries.get("union", False), } - response: MultiSearchResponse = await self.api_call.post( - AsyncMultiSearch.resource_path, - body=search_body, - params=common_params, - as_json=True, - entity_type=MultiSearchResponse, - ) - return response + + +def _without_stream_config( + common_params: MultiSearchCommonParameters, +) -> typing.Dict[str, object]: + """Return the parameters to send, leaving out the client-side ``stream_config``.""" + return { + param: param_value + for param, param_value in common_params.items() + if param != "stream_config" + } diff --git a/src/typesense/async_/stream.py b/src/typesense/async_/stream.py new file mode 100644 index 0000000..e84ec3b --- /dev/null +++ b/src/typesense/async_/stream.py @@ -0,0 +1,220 @@ +""" +Streamed conversational search responses. + +With ``conversation_stream`` enabled, Typesense sends the LLM's answer as +server-sent events while it is generated, then the full search response: + + data: {"conversation_id": "...", "message": "The"} + data: {"conversation_id": "...", "message": " answer"} + data: [DONE] + data: {"conversation": {...}, "hits": [...], ...} + +``AsyncSearchStream`` yields the answer pieces as ``MessageChunk`` dicts and keeps +the final event as the search response, returned by ``get_final_response``. +""" + +import sys +from types import TracebackType + +from typesense.exceptions import TypesenseClientError +from typesense.http_backend import ResponseType +from typesense.sse import SSEDecoder, ServerSentEvent, aiter_events +from typesense.types.document import MessageChunk, StreamConfig, StreamConfigBuilder + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +TFinal = typing.TypeVar("TFinal") + +_DONE = "[DONE]" + + +class AsyncSearchStream(typing.Generic[TFinal]): + """ + An open streaming search response. + + Iterate over it for the pieces of the LLM's answer, then call + ``get_final_response`` for the full search response. Use it as a context + manager, or call ``aclose``, to release the connection if you stop early. + + Attributes: + response (httpx.Response | httpx2.Response): The underlying response. + """ + + def __init__( + self, + response: ResponseType, + on_close: typing.Callable[[], None], + ) -> None: + """ + Initialize the stream. + + Args: + response (httpx.Response | httpx2.Response): A successful response + opened with ``stream=True``. + on_close (Callable[[], None]): Called once when the stream is closed, + to release the request's concurrency slot. + """ + self.response = response + self._on_close = on_close + self._closed = False + self._final: typing.Optional[TFinal] = None + self._decoder = SSEDecoder() + self._iterator = self._iter_chunks() + + def __aiter__(self) -> typing.Self: + """Return the stream itself; it can be iterated only once.""" + return self + + async def __anext__(self) -> MessageChunk: + """Return the next piece of the answer.""" + return await self._iterator.__anext__() + + async def __aenter__(self) -> typing.Self: + """Enter the context manager.""" + return self + + async def __aexit__( + self, + exc_type: typing.Optional[typing.Type[BaseException]], + exc_val: typing.Optional[BaseException], + exc_tb: typing.Optional[TracebackType], + ) -> None: + """Close the stream.""" + await self.aclose() + + async def get_final_response(self) -> TFinal: + """ + Read the rest of the stream and return the full search response. + + Returns: + TFinal: The search response sent after the answer. + + Raises: + TypesenseClientError: If the stream was closed or ended before the + search response arrived. + """ + if self._final is None and self._closed: + raise TypesenseClientError( + "The stream was closed before the search response arrived.", + ) + async for _ in self: + pass + if self._final is None: + raise TypesenseClientError(self._missing_final_message()) + return self._final + + async def aclose(self) -> None: + """Close the response and release its connection.""" + await self._iterator.aclose() + await self._close_response() + + async def _close_response(self) -> None: + """Close the response once, then run ``on_close``.""" + if self._closed: + return + self._closed = True + try: + await self.response.aclose() + finally: + self._on_close() + + async def _iter_chunks(self) -> typing.AsyncGenerator[MessageChunk, None]: + """Yield the answer pieces and keep the final search response.""" + try: + if "event-stream" not in self.response.headers.get("Content-Type", ""): + # Typesense answers with plain JSON when the LLM is never called, + # e.g. when every search of a multi-search fails. + await self.response.aread() + self._final = typing.cast(TFinal, self.response.json()) + return + async for event in aiter_events(self.response.aiter_bytes(), self._decoder): + chunk = self._handle_event(event) + if chunk is not None: + yield chunk + if self._final is None: + raise TypesenseClientError(self._missing_final_message()) + finally: + await self._close_response() + + def _handle_event(self, event: ServerSentEvent) -> typing.Optional[MessageChunk]: + """Return the event's answer piece, or keep it as the final response.""" + if event.data == _DONE: + return None + try: + payload = event.json() + except ValueError as json_error: + raise TypesenseClientError( + f"Invalid event in stream: {event.data}", + ) from json_error + if not isinstance(payload, dict): + return None + if any(key in payload for key in ("hits", "grouped_hits", "results")): + self._final = typing.cast(TFinal, payload) + return None + if "conversation_id" in payload and "message" in payload: + return MessageChunk( + conversation_id=payload["conversation_id"], + message=payload["message"], + ) + if "message" in payload: + raise TypesenseClientError(payload["message"]) + return None + + def _missing_final_message(self) -> str: + """Describe a stream that ended without the search response.""" + message = "The stream ended before the search response arrived." + remainder = self._decoder.remainder.strip() + return f"{message} {remainder}" if remainder else message + + +def resolve_stream_config( + stream_config: typing.Union[ + StreamConfig[TFinal], + StreamConfigBuilder[TFinal], + None, + ], +) -> typing.Optional[StreamConfig[TFinal]]: + """Return the callbacks of a ``StreamConfig`` or ``StreamConfigBuilder``.""" + if isinstance(stream_config, StreamConfigBuilder): + return stream_config.build() + return stream_config + + +def notify_error( + stream_config: typing.Optional[StreamConfig[TFinal]], + error: BaseException, +) -> None: + """Run the ``on_error`` callback, if there is one.""" + on_error = (stream_config or {}).get("on_error") + if on_error is not None: + on_error(error) + + +async def consume_stream( + stream: AsyncSearchStream[TFinal], + stream_config: typing.Optional[StreamConfig[TFinal]], +) -> TFinal: + """ + Read a stream to the end, running the ``on_chunk`` and ``on_complete`` callbacks. + + Args: + stream (AsyncSearchStream): The stream to read. + stream_config (StreamConfig | None): The callbacks to run. + + Returns: + TFinal: The full search response. + """ + stream_config = stream_config or {} + on_chunk = stream_config.get("on_chunk") + async with stream: + async for chunk in stream: + if on_chunk is not None: + on_chunk(chunk) + final_response = await stream.get_final_response() + on_complete = stream_config.get("on_complete") + if on_complete is not None: + on_complete(final_response) + return final_response diff --git a/src/typesense/stream_handlers.py b/src/typesense/stream_handlers.py deleted file mode 100644 index 45e38d5..0000000 --- a/src/typesense/stream_handlers.py +++ /dev/null @@ -1,164 +0,0 @@ -""" -SSE stream parsing and chunk combining for conversation search streaming. - -This module contains pure logic for parsing server-sent event lines from -conversation_stream responses and combining message chunks into a final -search response. Used by async API calls. -""" - -import json -import sys - -if sys.version_info >= (3, 11): - import typing -else: - import typing_extensions as typing - -from typesense.types.document import MessageChunk - -JSONPrimitive: typing.TypeAlias = typing.Union[str, int, float, bool, None] -JSONValue: typing.TypeAlias = typing.Union[ - JSONPrimitive, typing.Dict[str, "JSONValue"], typing.List["JSONValue"] -] -JSONDict: typing.TypeAlias = typing.Dict[str, JSONValue] - -_SEARCH_RESPONSE_KEYS = frozenset( - {"results", "found", "hits", "page", "search_time_ms"} -) - -StreamChunk: typing.TypeAlias = typing.Union[MessageChunk, JSONDict] - - -def parse_sse_line(line: str) -> typing.Optional[StreamChunk]: - """ - Parse a single SSE line into a MessageChunk, search response dict, or None. - - Handles: - - Empty lines and "data: [DONE]" -> None - - "data: {...}" -> parse JSON, return MessageChunk or search response - - Raw JSON line starting with "{" -> same - - Plain text -> return chunk with conversation_id="unknown", message=line - - Returns: - MessageChunk for conversation chunks, dict for search responses, or None to skip. - """ - line = line.strip() - if not line or line == "data: [DONE]": - return None - - # SSE format: "data: {...}" - if line.startswith("data:"): - content = line[5:].strip() - return _parse_data_content(content) - - # Raw JSON - if line.startswith("{"): - return _parse_json_content(line) - - return _chunk_from_text(line) - - -def _parse_data_content(content: str) -> typing.Optional[StreamChunk]: - """Parse the content after 'data:' into a MessageChunk, search response, or None.""" - if not content: - return None - if content.startswith("{"): - return _parse_json_content(content) - return _chunk_from_text(content) - - -def _parse_json_content(raw: str) -> StreamChunk: - """Parse a JSON string into a MessageChunk or search response dict.""" - try: - data = json.loads(raw) - except json.JSONDecodeError: - return _chunk_from_text(raw) - if not isinstance(data, dict): - return _chunk_from_text(json.dumps(data)) - - parsed = typing.cast(JSONDict, data) - conversation_id = parsed.get("conversation_id") - message = parsed.get("message") - nested_conversation = parsed.get("conversation") - - if conversation_id is None or message is None: - if isinstance(nested_conversation, dict): - nested_conversation_id = nested_conversation.get("conversation_id") - nested_message = nested_conversation.get("message") - if conversation_id is None and nested_conversation_id is not None: - conversation_id = nested_conversation_id - if message is None and nested_message is not None: - message = nested_message - - if conversation_id is None: - parsed["conversation_id"] = "unknown" - elif not isinstance(conversation_id, str): - parsed["conversation_id"] = str(conversation_id) - else: - parsed["conversation_id"] = conversation_id - - if message is None: - parsed["message"] = "" - elif not isinstance(message, str): - parsed["message"] = str(message) - else: - parsed["message"] = message - - return parsed - - -def _is_search_response_dict(data: typing.Mapping[str, JSONValue]) -> bool: - """Check if a dict is a search response (has found, hits, results, etc.).""" - return bool(set(data.keys()) & _SEARCH_RESPONSE_KEYS) - - -def is_message_chunk(chunk: JSONValue) -> bool: - """Return True if chunk is a conversation message chunk (has conversation_id and message).""" - if not isinstance(chunk, dict): - return False - if "message" not in chunk or "conversation_id" not in chunk: - return False - return not _is_search_response_dict(chunk) - - -def is_complete_search_response(chunk: JSONValue) -> bool: - """Return True if chunk looks like a full search response (has hits, found, etc.).""" - if not isinstance(chunk, dict) or not chunk: - return False - keys = set(chunk.keys()) - return bool(keys & _SEARCH_RESPONSE_KEYS) - - -def combine_stream_chunks( - chunks: typing.Sequence[StreamChunk], -) -> JSONDict: - """ - Combine streamed chunks into a single search response. - - - If no chunks, return empty dict. - - If one chunk, return it. - - If we have message chunks (conversation_id + message), find the metadata - chunk (complete search response) and return it; otherwise return last chunk - if it is complete. - - For regular search streams, return the last chunk if it is a complete response. - """ - if not chunks: - return {} - if len(chunks) == 1: - return typing.cast(JSONDict, chunks[0]) - - message_chunks = [c for c in chunks if is_message_chunk(c)] - if message_chunks: - for chunk in chunks: - if is_complete_search_response(chunk): - return typing.cast(JSONDict, chunk) - return typing.cast(JSONDict, chunks[-1]) - - last = chunks[-1] - if is_complete_search_response(last): - return typing.cast(JSONDict, last) - return typing.cast(JSONDict, last) - - -def _chunk_from_text(text: str) -> MessageChunk: - return {"conversation_id": "unknown", "message": text} diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index 57ff736..beb0e3c 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -32,7 +32,6 @@ """ import time -import json import sys from types import MappingProxyType, TracebackType @@ -58,16 +57,9 @@ backend_errors, verify_option, ) +from .stream import SearchStream from typesense.node_manager import NodeManager -from typesense.request_handler import RequestHandler -from typesense.stream_handlers import ( - JSONDict, - StreamChunk, - combine_stream_chunks, - is_message_chunk, - parse_sse_line, -) -from typesense.types.document import StreamConfig +from typesense.request_handler import RequestHandler, _QueryParams if sys.version_info >= (3, 11): import typing @@ -234,8 +226,6 @@ def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[False], params: typing.Union[TParams, None] = None, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, ) -> str: """ Execute an async GET request to the Typesense API. @@ -257,8 +247,6 @@ def get( entity_type: typing.Type[TEntityDict], as_json: typing.Literal[True] = True, params: typing.Union[TParams, None] = None, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, ) -> TEntityDict: """ Execute an async GET request to the Typesense API. @@ -279,8 +267,6 @@ def get( entity_type: typing.Type[TEntityDict], as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, params: typing.Union[TParams, None] = None, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, ) -> typing.Union[TEntityDict, str]: """ Execute an async GET request to the Typesense API. @@ -300,8 +286,6 @@ def get( entity_type, as_json, params=params, - stream_config=stream_config, - is_streaming_request=is_streaming_request, ) @typing.overload @@ -470,8 +454,6 @@ def _execute_request( as_json: typing.Literal[True], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> TEntityDict: """Execute an async request with retry logic.""" @@ -485,8 +467,6 @@ def _execute_request( as_json: typing.Literal[False], last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> str: """Execute an async request with retry logic.""" @@ -499,8 +479,6 @@ def _execute_request( as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, last_exception: typing.Union[None, Exception] = None, num_retries: int = 0, - stream_config: StreamConfig[TEntityDict] | None = None, - is_streaming_request: bool = False, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: """ @@ -532,10 +510,6 @@ def _execute_request( node, url, request_kwargs = self._prepare_request_params(endpoint, **kwargs) try: - if is_streaming_request and method == "GET": - return self._handle_streaming_get( - url, entity_type, stream_config, **request_kwargs - ) return self._make_request_and_process_response( method, node, @@ -548,13 +522,6 @@ def _execute_request( raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) - if is_streaming_request and stream_config: - on_error = stream_config.get("on_error") - if on_error: - try: - on_error(server_error) - except Exception: - pass if num_retries < self.config.num_retries: time.sleep(self.config.retry_interval_seconds) return self._execute_request( @@ -564,8 +531,6 @@ def _execute_request( as_json, last_exception=server_error, num_retries=num_retries + 1, - stream_config=stream_config, - is_streaming_request=is_streaming_request, **kwargs, ) @@ -595,72 +560,126 @@ def _make_request_and_process_response( else typing.cast(str, request_response) ) - def _handle_streaming_get( + def stream( + self, + method: str, + endpoint: str, + entity_type: typing.Type[TEntityDict], + params: typing.Union[TParams, None] = None, + body: typing.Union[TBody, None] = None, + ) -> SearchStream[TEntityDict]: + """ + Open a streaming request to the Typesense API. + + Failing nodes are retried like any other request until the response + headers arrive. Errors after that are raised while reading the stream and + are not retried, since part of the answer has already been read. + + Args: + method (str): The HTTP method to use. + endpoint (str): The API endpoint to call. + entity_type (Type[TEntityDict]): The type of the final response. + params (Union[TParams, None], optional): Query parameters for the request. + body (Union[TBody, None], optional): The request body. + + Returns: + SearchStream[TEntityDict]: The open stream. + """ + return self._execute_stream_request( + method, + endpoint, + entity_type, + params=params, + data=body, + ) + + def _execute_stream_request( self, + method: str, + endpoint: str, + entity_type: typing.Type[TEntityDict], + last_exception: typing.Union[None, Exception] = None, + num_retries: int = 0, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> SearchStream[TEntityDict]: + """Open a streaming request, failing over to other nodes like ``_execute_request``.""" + if num_retries > self.config.num_retries: + if last_exception: + raise last_exception + raise TypesenseClientError("All nodes are unhealthy") + + node, url, request_kwargs = self._prepare_request_params(endpoint, **kwargs) + + try: + return self._open_stream( + method, node, url, entity_type, **request_kwargs + ) + except _CLIENT_ERRORS: + raise + except _SERVER_ERRORS as server_error: + self.node_manager.set_node_health(node, is_healthy=False) + if num_retries < self.config.num_retries: + time.sleep(self.config.retry_interval_seconds) + return self._execute_stream_request( + method, + endpoint, + entity_type, + last_exception=server_error, + num_retries=num_retries + 1, + **kwargs, + ) + + def _open_stream( + self, + method: str, + node: Node, url: str, entity_type: typing.Type[TEntityDict], - stream_config: StreamConfig[TEntityDict] | None, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], - ) -> TEntityDict: - """Perform an async streaming GET, parse SSE lines, invoke callbacks, return combined result.""" - headers: typing.Dict[str, str] = { - self.request_handler.api_key_header_name: self.config.api_key, - "Accept": "text/event-stream", - } - headers.update(self.config.additional_headers) - extra_headers = kwargs.get("headers") - if extra_headers: - headers.update(extra_headers) - - params = kwargs.get("params") - content: typing.Union[str, bytes, None] = None - if body := kwargs.get("data"): - if isinstance(body, (str, bytes)): - content = body - else: - content = json.dumps(body) - - all_chunks: typing.List[StreamChunk] = [] - with self._client.stream( - "GET", + ) -> SearchStream[TEntityDict]: + """ + Send a streaming request to `node` and return the stream once headers arrive. + + The stream holds a concurrency slot until it is closed. Reads use + ``stream_read_timeout_seconds``, since the first piece of an answer only + arrives once the LLM starts generating it. + """ + request_kwargs = self.request_handler.build_request_kwargs(**kwargs) + headers = request_kwargs.get("headers", {}) + headers["Accept"] = "text/event-stream" + timeout = self._client.timeout + request = self._client.build_request( + method, url, - params=params, - content=content, + params=typing.cast( + typing.Optional[_QueryParams], + request_kwargs.get("params"), + ), + content=request_kwargs.get("content"), headers=headers, - timeout=self.config.connection_timeout_seconds, - ) as response: - if response.status_code < 200 or response.status_code >= 300: - response.read() - error_message = self.request_handler._get_error_message(response) - raise self.request_handler._get_exception(response.status_code)( - response.status_code, - error_message, - ) - for line in response.iter_lines(): - chunk = parse_sse_line(line) - if chunk is not None: - all_chunks.append(chunk) - if stream_config and is_message_chunk(chunk): - on_chunk = stream_config.get("on_chunk") - if on_chunk: - try: - on_chunk(chunk) - except Exception: - pass - - self.node_manager.set_node_health( - self.node_manager.get_node(), - is_healthy=True, + timeout=( + timeout.connect, + self.config.stream_read_timeout_seconds, + timeout.write, + timeout.pool, + ), ) - final: JSONDict = combine_stream_chunks(all_chunks) - if stream_config: - on_complete = stream_config.get("on_complete") - if on_complete: + + self._concurrency_limit.acquire() + try: + response = self._client.send(request, stream=True) + if response.status_code < 200 or response.status_code >= 300: try: - on_complete(typing.cast(TEntityDict, final)) - except Exception: - pass - return typing.cast(TEntityDict, final) + response.read() + finally: + response.close() + self.request_handler.raise_for_status(response) + except BaseException: + self._concurrency_limit.release() + raise + + self.node_manager.set_node_health(node, is_healthy=True) + return SearchStream(response, self._concurrency_limit.release) def _prepare_request_params( self, diff --git a/src/typesense/sync/documents.py b/src/typesense/sync/documents.py index e9225f6..badd527 100644 --- a/src/typesense/sync/documents.py +++ b/src/typesense/sync/documents.py @@ -21,6 +21,12 @@ from .api_call import ApiCall from .document import Document +from .stream import ( + SearchStream, + consume_stream, + notify_error, + resolve_stream_config, +) from typesense.exceptions import TypesenseClientError from typesense.logger import logger from typesense.preprocess import stringify_search_params @@ -43,7 +49,6 @@ ImportResponseWithId, SearchParameters, SearchResponse, - StreamConfigBuilder, UpdateByFilterParameters, UpdateByFilterResponse, ) @@ -361,30 +366,86 @@ def search(self, search_parameters: SearchParameters) -> SearchResponse[TDoc]: """ Search for documents in the collection. + With ``conversation_stream`` enabled, the LLM's answer is streamed and the + callbacks in ``stream_config`` run as it arrives. To iterate over the answer + instead, use ``search_stream``. + Args: search_parameters (SearchParameters): The search parameters. - Use conversation_stream=True and optionally stream_config (on_chunk, - on_complete, on_error) for conversational search streaming. Returns: SearchResponse[TDoc]: The search response containing matching documents. """ - params_for_api = dict(search_parameters) - stream_config = params_for_api.pop("stream_config", None) - if isinstance(stream_config, StreamConfigBuilder): - stream_config = stream_config.build() - conversation_stream = params_for_api.get("conversation_stream") is True - stringified_search_params = stringify_search_params(params_for_api) + if search_parameters.get("conversation_stream"): + stream_config = resolve_stream_config( + search_parameters.get("stream_config"), + ) + try: + search_stream = self._open_search_stream(search_parameters) + streamed_response: SearchResponse[TDoc] = consume_stream( + search_stream, + stream_config, + ) + except Exception as error: + notify_error(stream_config, error) + raise + return streamed_response + + stringified_search_params = stringify_search_params( + { + param: param_value + for param, param_value in search_parameters.items() + if param != "stream_config" + }, + ) response: SearchResponse[TDoc] = self.api_call.get( self._endpoint_path("search"), params=stringified_search_params, entity_type=SearchResponse, as_json=True, - stream_config=stream_config, - is_streaming_request=conversation_stream, ) return response + def search_stream( + self, + search_parameters: SearchParameters, + ) -> SearchStream[SearchResponse[TDoc]]: + """ + Search, streaming the LLM's answer as it is generated. + + Iterate over the returned stream for the pieces of the answer, then call + its ``get_final_response`` for the search response. Use the stream as a + context manager so the connection is released if you stop early. + + ``conversation`` and ``conversation_stream`` are enabled for you; pass the + ``conversation_model_id`` to answer with. ``stream_config`` is ignored. + + Args: + search_parameters (SearchParameters): The search parameters. + + Returns: + SearchStream[SearchResponse[TDoc]]: The open stream. + """ + return self._open_search_stream(search_parameters) + + def _open_search_stream( + self, + search_parameters: SearchParameters, + ) -> SearchStream[SearchResponse[TDoc]]: + """Open the search request as a stream.""" + stream_params: typing.Dict[str, object] = { + "conversation": True, + **search_parameters, + "conversation_stream": True, + } + stream_params.pop("stream_config", None) + return self.api_call.stream( + "GET", + self._endpoint_path("search"), + entity_type=SearchResponse, + params=stringify_search_params(stream_params), + ) + def delete( self, delete_parameters: typing.Union[DeleteQueryParameters, None] = None, diff --git a/src/typesense/sync/multi_search.py b/src/typesense/sync/multi_search.py index 2c81be6..ab14187 100644 --- a/src/typesense/sync/multi_search.py +++ b/src/typesense/sync/multi_search.py @@ -19,6 +19,12 @@ import sys from .api_call import ApiCall +from .stream import ( + SearchStream, + consume_stream, + notify_error, + resolve_stream_config, +) from typesense.preprocess import stringify_search_params from typesense.types.document import MultiSearchCommonParameters from typesense.types.multi_search import MultiSearchRequestSchema, MultiSearchResponse @@ -89,20 +95,93 @@ def perform( ... ], ... } ... ) + + With ``conversation_stream`` enabled in ``common_params``, the LLM's answer + is streamed and the callbacks in ``stream_config`` run as it arrives. To + iterate over the answer instead, use ``perform_stream``. + """ + if common_params and common_params.get("conversation_stream"): + stream_config = resolve_stream_config(common_params.get("stream_config")) + try: + search_stream = self.perform_stream(search_queries, common_params) + streamed_response: MultiSearchResponse = consume_stream( + search_stream, + stream_config, + ) + except Exception as error: + notify_error(stream_config, error) + raise + return streamed_response + + response: MultiSearchResponse = self.api_call.post( + MultiSearch.resource_path, + body=self._search_body(search_queries), + params=_without_stream_config(common_params) if common_params else None, + as_json=True, + entity_type=MultiSearchResponse, + ) + return response + + def perform_stream( + self, + search_queries: MultiSearchRequestSchema, + common_params: typing.Union[MultiSearchCommonParameters, None] = None, + ) -> SearchStream[MultiSearchResponse]: """ + Perform a multi-search, streaming the LLM's answer as it is generated. + + The searches' hits are combined into one context for a single answer, sent + in the response's top-level ``conversation``. Iterate over the returned + stream for the pieces of the answer, then call its ``get_final_response`` + for the multi-search response. Use the stream as a context manager so the + connection is released if you stop early. + + ``conversation`` and ``conversation_stream`` are enabled for you; pass + ``q`` and the ``conversation_model_id`` in ``common_params``, since + Typesense reads them from the query string. ``stream_config`` is ignored. + + Args: + search_queries (MultiSearchRequestSchema): The searches to perform. + common_params (Union[MultiSearchCommonParameters, None], optional): + Parameters for every search, including the conversation parameters. + + Returns: + SearchStream[MultiSearchResponse]: The open stream. + """ + stream_params: typing.Dict[str, object] = { + "conversation": True, + **_without_stream_config(common_params or {}), + "conversation_stream": True, + } + return self.api_call.stream( + "POST", + MultiSearch.resource_path, + entity_type=MultiSearchResponse, + params=stream_params, + body=self._search_body(search_queries), + ) + + @staticmethod + def _search_body( + search_queries: MultiSearchRequestSchema, + ) -> typing.Dict[str, object]: + """Build the request body, with every search's parameters stringified.""" stringified_search_params = [ stringify_search_params(search_params) for search_params in search_queries.get("searches") ] - search_body = { + return { "searches": stringified_search_params, "union": search_queries.get("union", False), } - response: MultiSearchResponse = self.api_call.post( - MultiSearch.resource_path, - body=search_body, - params=common_params, - as_json=True, - entity_type=MultiSearchResponse, - ) - return response + + +def _without_stream_config( + common_params: MultiSearchCommonParameters, +) -> typing.Dict[str, object]: + """Return the parameters to send, leaving out the client-side ``stream_config``.""" + return { + param: param_value + for param, param_value in common_params.items() + if param != "stream_config" + } diff --git a/src/typesense/sync/stream.py b/src/typesense/sync/stream.py new file mode 100644 index 0000000..91e2a4b --- /dev/null +++ b/src/typesense/sync/stream.py @@ -0,0 +1,220 @@ +""" +Streamed conversational search responses. + +With ``conversation_stream`` enabled, Typesense sends the LLM's answer as +server-sent events while it is generated, then the full search response: + + data: {"conversation_id": "...", "message": "The"} + data: {"conversation_id": "...", "message": " answer"} + data: [DONE] + data: {"conversation": {...}, "hits": [...], ...} + +``SearchStream`` yields the answer pieces as ``MessageChunk`` dicts and keeps +the final event as the search response, returned by ``get_final_response``. +""" + +import sys +from types import TracebackType + +from typesense.exceptions import TypesenseClientError +from typesense.http_backend import ResponseType +from typesense.sse import SSEDecoder, ServerSentEvent, iter_events +from typesense.types.document import MessageChunk, StreamConfig, StreamConfigBuilder + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +TFinal = typing.TypeVar("TFinal") + +_DONE = "[DONE]" + + +class SearchStream(typing.Generic[TFinal]): + """ + An open streaming search response. + + Iterate over it for the pieces of the LLM's answer, then call + ``get_final_response`` for the full search response. Use it as a context + manager, or call ``close``, to release the connection if you stop early. + + Attributes: + response (httpx.Response | httpx2.Response): The underlying response. + """ + + def __init__( + self, + response: ResponseType, + on_close: typing.Callable[[], None], + ) -> None: + """ + Initialize the stream. + + Args: + response (httpx.Response | httpx2.Response): A successful response + opened with ``stream=True``. + on_close (Callable[[], None]): Called once when the stream is closed, + to release the request's concurrency slot. + """ + self.response = response + self._on_close = on_close + self._closed = False + self._final: typing.Optional[TFinal] = None + self._decoder = SSEDecoder() + self._iterator = self._iter_chunks() + + def __iter__(self) -> typing.Self: + """Return the stream itself; it can be iterated only once.""" + return self + + def __next__(self) -> MessageChunk: + """Return the next piece of the answer.""" + return self._iterator.__next__() + + def __enter__(self) -> typing.Self: + """Enter the context manager.""" + return self + + def __exit__( + self, + exc_type: typing.Optional[typing.Type[BaseException]], + exc_val: typing.Optional[BaseException], + exc_tb: typing.Optional[TracebackType], + ) -> None: + """Close the stream.""" + self.close() + + def get_final_response(self) -> TFinal: + """ + Read the rest of the stream and return the full search response. + + Returns: + TFinal: The search response sent after the answer. + + Raises: + TypesenseClientError: If the stream was closed or ended before the + search response arrived. + """ + if self._final is None and self._closed: + raise TypesenseClientError( + "The stream was closed before the search response arrived.", + ) + for _ in self: + pass + if self._final is None: + raise TypesenseClientError(self._missing_final_message()) + return self._final + + def close(self) -> None: + """Close the response and release its connection.""" + self._iterator.close() + self._close_response() + + def _close_response(self) -> None: + """Close the response once, then run ``on_close``.""" + if self._closed: + return + self._closed = True + try: + self.response.close() + finally: + self._on_close() + + def _iter_chunks(self) -> typing.Generator[MessageChunk, None, None]: + """Yield the answer pieces and keep the final search response.""" + try: + if "event-stream" not in self.response.headers.get("Content-Type", ""): + # Typesense answers with plain JSON when the LLM is never called, + # e.g. when every search of a multi-search fails. + self.response.read() + self._final = typing.cast(TFinal, self.response.json()) + return + for event in iter_events(self.response.iter_bytes(), self._decoder): + chunk = self._handle_event(event) + if chunk is not None: + yield chunk + if self._final is None: + raise TypesenseClientError(self._missing_final_message()) + finally: + self._close_response() + + def _handle_event(self, event: ServerSentEvent) -> typing.Optional[MessageChunk]: + """Return the event's answer piece, or keep it as the final response.""" + if event.data == _DONE: + return None + try: + payload = event.json() + except ValueError as json_error: + raise TypesenseClientError( + f"Invalid event in stream: {event.data}", + ) from json_error + if not isinstance(payload, dict): + return None + if any(key in payload for key in ("hits", "grouped_hits", "results")): + self._final = typing.cast(TFinal, payload) + return None + if "conversation_id" in payload and "message" in payload: + return MessageChunk( + conversation_id=payload["conversation_id"], + message=payload["message"], + ) + if "message" in payload: + raise TypesenseClientError(payload["message"]) + return None + + def _missing_final_message(self) -> str: + """Describe a stream that ended without the search response.""" + message = "The stream ended before the search response arrived." + remainder = self._decoder.remainder.strip() + return f"{message} {remainder}" if remainder else message + + +def resolve_stream_config( + stream_config: typing.Union[ + StreamConfig[TFinal], + StreamConfigBuilder[TFinal], + None, + ], +) -> typing.Optional[StreamConfig[TFinal]]: + """Return the callbacks of a ``StreamConfig`` or ``StreamConfigBuilder``.""" + if isinstance(stream_config, StreamConfigBuilder): + return stream_config.build() + return stream_config + + +def notify_error( + stream_config: typing.Optional[StreamConfig[TFinal]], + error: BaseException, +) -> None: + """Run the ``on_error`` callback, if there is one.""" + on_error = (stream_config or {}).get("on_error") + if on_error is not None: + on_error(error) + + +def consume_stream( + stream: SearchStream[TFinal], + stream_config: typing.Optional[StreamConfig[TFinal]], +) -> TFinal: + """ + Read a stream to the end, running the ``on_chunk`` and ``on_complete`` callbacks. + + Args: + stream (SearchStream): The stream to read. + stream_config (StreamConfig | None): The callbacks to run. + + Returns: + TFinal: The full search response. + """ + stream_config = stream_config or {} + on_chunk = stream_config.get("on_chunk") + with stream: + for chunk in stream: + if on_chunk is not None: + on_chunk(chunk) + final_response = stream.get_final_response() + on_complete = stream_config.get("on_complete") + if on_complete is not None: + on_complete(final_response) + return final_response diff --git a/src/typesense/types/document.py b/src/typesense/types/document.py index f1307f5..b5cf565 100644 --- a/src/typesense/types/document.py +++ b/src/typesense/types/document.py @@ -586,47 +586,67 @@ class NLLanguageParameters(typing.TypedDict): nl_query_debug: typing.NotRequired[bool] -class MessageChunk(typing.TypedDict): - """ - A single chunk from a conversation stream response. +TFinal = typing.TypeVar("TFinal") - Attributes: - conversation_id (str): ID of the conversation. - message (str): Message content for this chunk. + +class ConversationParameters(typing.TypedDict): """ + Parameters for [conversational search](https://typesense.org/docs/29.0/api/conversational-search-rag.html). - conversation_id: str - message: str + Attributes: + conversation (bool): Whether to answer the query with an LLM. + conversation_model_id (str): The ID of the conversation model to answer with. + conversation_id (str): The ID of an earlier conversation to continue. + conversation_stream (bool): Whether to stream the answer as server-sent + events. Use ``search_stream`` to iterate over the answer as it arrives. + stream_config (StreamConfig | StreamConfigBuilder): Callbacks to run while + a ``conversation_stream`` search streams. Not sent to the server. + """ + + conversation: typing.NotRequired[bool] + conversation_model_id: typing.NotRequired[str] + conversation_id: typing.NotRequired[str] + conversation_stream: typing.NotRequired[bool] + stream_config: typing.NotRequired[ + typing.Union["StreamConfig[typing.Any]", "StreamConfigBuilder[typing.Any]"] + ] -class StreamConfig(typing.Generic[TDoc], typing.TypedDict, total=False): +class MessageChunk(typing.TypedDict): """ - Configuration for streaming conversation search responses. + A piece of a streamed conversation answer. Attributes: - on_chunk: Callback invoked for each streamed chunk (conversation_id, message). - on_complete: Callback invoked when the stream completes with the full search response. - on_error: Callback invoked if an error occurs during streaming. + conversation_id (str): The ID of the conversation. + message (str): The next piece of the answer. """ - on_chunk: typing.Callable[[MessageChunk], None] - on_complete: "OnCompleteCallback[TDoc]" - on_error: typing.Callable[[BaseException], None] + conversation_id: str + message: str OnChunkCallback = typing.Callable[[MessageChunk], None] +OnErrorCallback = typing.Callable[[BaseException], None] -class OnCompleteCallback(typing.Protocol[TDoc]): - def __call__(self, response: "SearchResponse[TDoc]") -> None: ... +class StreamConfig(typing.Generic[TFinal], typing.TypedDict, total=False): + """ + Callbacks for a streamed conversation search. + Attributes: + on_chunk: Called with each piece of the answer. + on_complete: Called with the full search response once the stream ends. + on_error: Called with the error if the search fails; the error is then raised. + """ -OnErrorCallback = typing.Callable[[BaseException], None] + on_chunk: OnChunkCallback + on_complete: typing.Callable[[TFinal], None] + on_error: OnErrorCallback -class StreamConfigBuilder(typing.Generic[TDoc]): +class StreamConfigBuilder(typing.Generic[TFinal]): """ - Builder for StreamConfig using decorators. + Build a ``StreamConfig`` by registering callbacks with decorators. Example: >>> stream = StreamConfigBuilder() @@ -635,111 +655,43 @@ class StreamConfigBuilder(typing.Generic[TDoc]): ... def handle_chunk(chunk: MessageChunk) -> None: ... print(chunk["message"], end="", flush=True) >>> - >>> @stream.on_complete - ... def handle_complete(response: dict) -> None: - ... print(f"Done! Found {response.get('found', 0)}") - >>> - >>> response = await client.collections["docs"].documents.search({ - ... "q": "query", - ... "query_by": "content", - ... "conversation_stream": True, - ... "stream_config": stream, - ... }) + >>> response = client.collections["docs"].documents.search( + ... { + ... "q": "query", + ... "query_by": "content", + ... "conversation": True, + ... "conversation_model_id": "conv-model", + ... "conversation_stream": True, + ... "stream_config": stream, + ... } + ... ) """ def __init__(self) -> None: - """Initialize an empty StreamConfigBuilder.""" - self._on_chunk: OnChunkCallback | None = None - self._on_complete: OnCompleteCallback[TDoc] | None = None - self._on_error: OnErrorCallback | None = None + """Initialize a builder with no callbacks.""" + self._config: StreamConfig[TFinal] = {} def on_chunk(self, func: OnChunkCallback) -> OnChunkCallback: - """ - Decorator to register an on_chunk callback. - - Args: - func: Callback invoked for each streamed message chunk. - - Returns: - The original function (unmodified). - """ - self._on_chunk = func + """Register ``func`` to be called with each piece of the answer.""" + self._config["on_chunk"] = func return func - def on_complete(self, func: OnCompleteCallback[TDoc]) -> OnCompleteCallback[TDoc]: - """ - Decorator to register an on_complete callback. - - Args: - func: Callback invoked when streaming completes with the full response. - - Returns: - The original function (unmodified). - """ - self._on_complete = func + def on_complete( + self, + func: typing.Callable[[TFinal], None], + ) -> typing.Callable[[TFinal], None]: + """Register ``func`` to be called with the full search response.""" + self._config["on_complete"] = func return func def on_error(self, func: OnErrorCallback) -> OnErrorCallback: - """ - Decorator to register an on_error callback. - - Args: - func: Callback invoked if an error occurs during streaming. - - Returns: - The original function (unmodified). - """ - self._on_error = func + """Register ``func`` to be called with the error if the search fails.""" + self._config["on_error"] = func return func - def build(self) -> StreamConfig[TDoc]: - """ - Build the StreamConfig dictionary. - - Returns: - A StreamConfig with the registered callbacks. - """ - config: StreamConfig[TDoc] = {} - if self._on_chunk is not None: - config["on_chunk"] = self._on_chunk - if self._on_complete is not None: - config["on_complete"] = self._on_complete - if self._on_error is not None: - config["on_error"] = self._on_error - return config - - def get( - self, - key: typing.Literal["on_chunk", "on_complete", "on_error"], - default: typing.Callable[..., None] | None = None, - ) -> typing.Callable[..., None] | None: - """ - Get a callback by key (for compatibility with dict-like access). - - Args: - key: The callback name ('on_chunk', 'on_complete', or 'on_error'). - default: Default value if the callback is not set. - - Returns: - The callback function or the default value. - """ - return self.build().get(key, default) - - -class ConversationStreamParameters(typing.Generic[TDoc], typing.TypedDict): - """ - Parameters for conversational search streaming. - - Attributes: - conversation_stream (bool): When true, the search response is streamed (SSE). - stream_config: Callbacks for stream events. Not sent to the API. - Can be a StreamConfig dict or a StreamConfigBuilder instance. - """ - - conversation_stream: typing.NotRequired[bool] - stream_config: typing.NotRequired[ - typing.Union[StreamConfig[TDoc], StreamConfigBuilder[TDoc]] - ] + def build(self) -> StreamConfig[TFinal]: + """Return the registered callbacks as a ``StreamConfig``.""" + return self._config.copy() class SearchParameters( @@ -754,13 +706,12 @@ class SearchParameters( TypoToleranceParameters, CachingParameters, NLLanguageParameters, - ConversationStreamParameters[TDoc], - typing.Generic[TDoc], + ConversationParameters, ): """Parameters for searching documents.""" -class MultiSearchParameters(SearchParameters[TDoc], typing.Generic[TDoc]): +class MultiSearchParameters(SearchParameters): """ Parameters for performing a [Federated/Multi-Search](https://typesense.org/docs/26.0/api/federated-multi-search.html#federated-multi-search). @@ -784,6 +735,7 @@ class MultiSearchCommonParameters( ResultsParameters, TypoToleranceParameters, CachingParameters, + ConversationParameters, ): """ [Query parameters](https://typesense.org/docs/26.0/api/federated-multi-search.html#multi-search-parameters) for multi-search. @@ -1025,7 +977,7 @@ class LLMResponse(typing.TypedDict): model: str -class ParsedNLQuery(typing.Generic[TDoc], typing.TypedDict): +class ParsedNLQuery(typing.TypedDict): """ Schema for a parsed natural language query. @@ -1037,8 +989,8 @@ class ParsedNLQuery(typing.Generic[TDoc], typing.TypedDict): """ parse_time_ms: int - generated_params: SearchParameters[TDoc] - augmented_params: SearchParameters[TDoc] + generated_params: SearchParameters + augmented_params: SearchParameters llm_response: typing.NotRequired[LLMResponse] @@ -1070,7 +1022,7 @@ class SearchResponse(typing.Generic[TDoc], typing.TypedDict): hits: typing.List[Hit[TDoc]] grouped_hits: typing.NotRequired[typing.List[GroupedHit[TDoc]]] conversation: typing.NotRequired[Conversation] - parsed_nl_query: typing.NotRequired[ParsedNLQuery[TDoc]] + parsed_nl_query: typing.NotRequired[ParsedNLQuery] class DeleteSingleDocumentParameters(typing.TypedDict): diff --git a/src/typesense/types/multi_search.py b/src/typesense/types/multi_search.py index 3619c0b..13be8cf 100644 --- a/src/typesense/types/multi_search.py +++ b/src/typesense/types/multi_search.py @@ -2,7 +2,11 @@ import sys -from typesense.types.document import MultiSearchParameters, SearchResponse +from typesense.types.document import ( + Conversation, + MultiSearchParameters, + SearchResponse, +) if sys.version_info >= (3, 11): import typing @@ -16,9 +20,11 @@ class MultiSearchResponse(typing.TypedDict): Attributes: results (list[SearchResponse]): The search results. + conversation (Conversation): The LLM's answer, for a conversational search. """ results: typing.List[SearchResponse[typing.Any]] # noqa: WPS110 + conversation: typing.NotRequired[Conversation] class MultiSearchRequestSchema(typing.TypedDict): diff --git a/tests/fixtures/streaming_fixtures.py b/tests/fixtures/streaming_fixtures.py deleted file mode 100644 index 4df0d52..0000000 --- a/tests/fixtures/streaming_fixtures.py +++ /dev/null @@ -1,180 +0,0 @@ -"""Fixtures for streaming tests.""" - -import json -import os -import sys -from types import TracebackType - -import pytest -import requests - -if sys.version_info >= (3, 11): - import typing -else: - import typing_extensions as typing - - -JSONPrimitive: typing.TypeAlias = typing.Union[str, int, float, bool, None] -JSONValue: typing.TypeAlias = typing.Union[ - JSONPrimitive, typing.Dict[str, "JSONValue"], typing.List["JSONValue"] -] -JSONDict: typing.TypeAlias = typing.Dict[str, JSONValue] - - -class FakeAsyncStreamResponse: - """Minimal async streaming response for httpx.AsyncClient.stream().""" - - def __init__( - self, - *, - lines: typing.Sequence[str], - status_code: int = 200, - headers: typing.Mapping[str, str] | None = None, - text: str = "", - ) -> None: - self.status_code = status_code - self._lines = list(lines) - self.headers = dict(headers or {}) - self.text = text - - async def aiter_lines(self) -> typing.AsyncIterator[str]: - for line in self._lines: - yield line - - async def aread(self) -> bytes: - return self.text.encode() - - def json(self) -> JSONDict: - return typing.cast(JSONDict, json.loads(self.text)) - - -class FakeAsyncStreamContext: - """Async context manager that yields a fake streaming response.""" - - def __init__(self, response: FakeAsyncStreamResponse) -> None: - self._response = response - - async def __aenter__(self) -> FakeAsyncStreamResponse: - return self._response - - async def __aexit__( - self, - exc_type: type[BaseException] | None, - exc_val: BaseException | None, - exc_tb: TracebackType | None, - ) -> None: - return None - - -class FakeStreamResponse: - """Minimal streaming response for httpx.Client.stream().""" - - def __init__( - self, - *, - lines: typing.Sequence[str], - status_code: int = 200, - headers: typing.Mapping[str, str] | None = None, - text: str = "", - ) -> None: - self.status_code = status_code - self._lines = list(lines) - self.headers = dict(headers or {}) - self.text = text - - def iter_lines(self) -> typing.Iterator[str]: - for line in self._lines: - yield line - - def read(self) -> bytes: - return self.text.encode() - - def json(self) -> JSONDict: - return typing.cast(JSONDict, json.loads(self.text)) - - -class FakeStreamContext: - """Sync context manager that yields a fake streaming response.""" - - def __init__(self, response: FakeStreamResponse) -> None: - self._response = response - - def __enter__(self) -> FakeStreamResponse: - return self._response - - def __exit__( - self, - exc_type: type[BaseException] | None, - exc_val: BaseException | None, - exc_tb: TracebackType | None, - ) -> None: - return None - - -@pytest.fixture(name="stream_response_async") -def stream_response_async_fixture() -> type[FakeAsyncStreamResponse]: - return FakeAsyncStreamResponse - - -@pytest.fixture(name="stream_context_async") -def stream_context_async_fixture() -> type[FakeAsyncStreamContext]: - return FakeAsyncStreamContext - - -@pytest.fixture(name="stream_response") -def stream_response_fixture() -> type[FakeStreamResponse]: - return FakeStreamResponse - - -@pytest.fixture(name="stream_context") -def stream_context_fixture() -> type[FakeStreamContext]: - return FakeStreamContext - - -@pytest.fixture(name="create_streaming_collection") -def create_streaming_collection_fixture(delete_all: None) -> str: - """Create a collection for streaming tests with an auto-embedding field.""" - open_ai_key = os.environ.get("OPEN_AI_KEY") - if not open_ai_key: - pytest.skip("OPEN_AI_KEY is required for streaming integration tests.") - url = "http://localhost:8108/collections" - headers = {"X-TYPESENSE-API-KEY": "xyz"} - collection_data = { - "name": "streaming_docs", - "fields": [ - { - "name": "title", - "type": "string", - }, - { - "name": "embedding", - "type": "float[]", - "embed": { - "from": ["title"], - "model_config": { - "model_name": "openai/text-embedding-3-small", - "api_key": open_ai_key, - }, - }, - }, - ], - } - - response = requests.post(url, headers=headers, json=collection_data, timeout=3) - response.raise_for_status() - return "streaming_docs" - - -@pytest.fixture(name="create_streaming_document") -def create_streaming_document_fixture(create_streaming_collection: str) -> str: - """Create a document for streaming tests.""" - url = "http://localhost:8108/collections/streaming_docs/documents" - headers = {"X-TYPESENSE-API-KEY": "xyz"} - document_data = { - "id": "stream-1", - "title": "Company profile", - } - - response = requests.post(url, headers=headers, json=document_data, timeout=3) - response.raise_for_status() - return "stream-1" diff --git a/tests/streaming_async_test.py b/tests/streaming_async_test.py deleted file mode 100644 index 2882c22..0000000 --- a/tests/streaming_async_test.py +++ /dev/null @@ -1,454 +0,0 @@ -"""Async streaming conversation search tests.""" - -import sys - -import pytest - -if sys.version_info >= (3, 11): - import typing -else: - import typing_extensions as typing - -from tests.fixtures.streaming_fixtures import ( - FakeAsyncStreamContext, - FakeAsyncStreamResponse, - JSONValue, -) -from typesense.async_.api_call import AsyncApiCall -from typesense.async_.documents import AsyncDocuments -from typesense.exceptions import ServerError -from typesense.types.document import ( - DocumentSchema, - MessageChunk, - StreamConfig, - StreamConfigBuilder, -) - - -async def test_streaming_search_invokes_on_chunk_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test that streaming search invokes on_chunk for each message chunk.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - stream_config: StreamConfig[DocumentSchema] = {"on_chunk": on_chunk} - - sse_lines = [ - 'data: {"conversation_id":"123","message":"First chunk"}', - 'data: {"conversation_id":"123","message":"Second chunk"}', - '{"found": 2, "hits": [], "page": 1, "search_time_ms": 10}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - result = await fake_async_documents.search( - { - "q": "test query", - "query_by": "title", - "conversation_stream": True, - "stream_config": stream_config, - } - ) - - assert len(chunks_received) == 2 - assert chunks_received[0]["message"] == "First chunk" - assert chunks_received[1]["message"] == "Second chunk" - assert result["found"] == 2 - - -async def test_streaming_search_handles_plain_text_lines_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Plain text lines should be treated as message chunks with unknown id.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - "Hello", - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - assert chunks_received[0]["conversation_id"] == "unknown" - assert chunks_received[0]["message"] == "Hello" - - -async def test_streaming_search_handles_missing_fields_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """JSON lines without conversation_id/message should use defaults.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - 'data: {"foo":"bar"}', - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - assert chunks_received[0]["conversation_id"] == "unknown" - assert chunks_received[0]["message"] == "" - - -async def test_streaming_search_skips_done_marker_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """data: [DONE] lines should be ignored.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - 'data: {"conversation_id":"123","message":"Chunk"}', - "data: [DONE]", - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - - -async def test_streaming_search_handles_json_array_lines_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """JSON arrays should be treated as plain text message chunks.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - 'data: ["a", "b"]', - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - assert chunks_received[0]["message"] == '["a", "b"]' - - -async def test_streaming_search_supports_builder_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test StreamConfigBuilder for streaming callbacks.""" - complete_calls: typing.List[int] = [] - - stream = StreamConfigBuilder() - - @stream.on_complete - def on_complete(response: typing.Mapping[str, JSONValue]) -> None: - found = response.get("found") - if isinstance(found, int): - complete_calls.append(found) - - sse_lines = [ - 'data: {"conversation_id":"123","message":"Hello"}', - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": stream, - } - ) - - assert complete_calls == [1] - - -async def test_stream_config_not_sent_to_api_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test that stream_config is removed from API params.""" - captured_params: typing.Dict[str, str] = {} - - sse_lines = [ - '{"found": 0, "hits": [], "page": 1, "search_time_ms": 1}', - ] - response = stream_response_async(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - if params: - captured_params.update(params) - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - stream_config: StreamConfig[DocumentSchema] = {"on_chunk": lambda _: None} - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": stream_config, - } - ) - - assert "stream_config" not in captured_params - assert captured_params.get("conversation_stream") == "true" - - -async def test_streaming_search_invokes_on_error_async( - fake_async_documents: AsyncDocuments[DocumentSchema], - stream_response_async: type[FakeAsyncStreamResponse], - stream_context_async: type[FakeAsyncStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test that streaming search invokes on_error for request failures.""" - errors: typing.List[BaseException] = [] - - def on_error(error: BaseException) -> None: - errors.append(error) - - fake_async_documents.api_call.config.num_retries = 0 - - response = stream_response_async( - lines=[], - status_code=500, - headers={"Content-Type": "application/json"}, - text='{"message": "Server error"}', - ) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeAsyncStreamContext: - return stream_context_async(response) - - monkeypatch.setattr( - fake_async_documents.api_call._client, - "stream", - fake_stream, - ) - - with pytest.raises(ServerError): - await fake_async_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_error": on_error}, - } - ) - - assert len(errors) == 1 - assert isinstance(errors[0], ServerError) - - -@pytest.mark.open_ai -async def test_actual_streaming_search_async( - actual_async_api_call: AsyncApiCall, - create_streaming_collection: str, - create_streaming_document: str, - create_conversations_model: str, -) -> None: - """Integration test against a real Typesense server with conversation streaming.""" - actual_async_documents = AsyncDocuments( - actual_async_api_call, - create_streaming_collection, - ) - chunks_received: typing.List[MessageChunk] = [] - complete_called: typing.List[bool] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - def on_complete(response: typing.Mapping[str, JSONValue]) -> None: - complete_called.append(True) - - response = await actual_async_documents.search( - { - "q": "What is this document about?", - "query_by": "embedding", - "conversation": True, - "conversation_stream": True, - "conversation_model_id": create_conversations_model, - "prefix": False, - "exclude_fields": "embedding", - "stream_config": {"on_chunk": on_chunk, "on_complete": on_complete}, - } - ) - - assert complete_called == [True] - assert len(chunks_received) > 0 - assert "found" in response or "hits" in response diff --git a/tests/streaming_test.py b/tests/streaming_test.py deleted file mode 100644 index 44b110b..0000000 --- a/tests/streaming_test.py +++ /dev/null @@ -1,414 +0,0 @@ -"""Sync streaming conversation search tests.""" - -import sys - -import pytest - -if sys.version_info >= (3, 11): - import typing -else: - import typing_extensions as typing - -from tests.fixtures.streaming_fixtures import ( - FakeStreamContext, - FakeStreamResponse, - JSONValue, -) -from typesense.exceptions import ServerError -from typesense.sync.documents import Documents -from typesense.types.document import ( - DocumentSchema, - MessageChunk, - StreamConfig, - StreamConfigBuilder, -) - - -def test_streaming_search_invokes_on_chunk( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test that streaming search invokes on_chunk for each message chunk.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - stream_config: StreamConfig[DocumentSchema] = {"on_chunk": on_chunk} - - sse_lines = [ - 'data: {"conversation_id":"123","message":"First chunk"}', - 'data: {"conversation_id":"123","message":"Second chunk"}', - '{"found": 2, "hits": [], "page": 1, "search_time_ms": 10}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - result = fake_documents.search( - { - "q": "test query", - "query_by": "title", - "conversation_stream": True, - "stream_config": stream_config, - } - ) - - assert len(chunks_received) == 2 - assert chunks_received[0]["message"] == "First chunk" - assert chunks_received[1]["message"] == "Second chunk" - assert result["found"] == 2 - - -def test_streaming_search_handles_plain_text_lines( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Plain text lines should be treated as message chunks with unknown id.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - "Hello", - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - assert chunks_received[0]["conversation_id"] == "unknown" - assert chunks_received[0]["message"] == "Hello" - - -def test_streaming_search_handles_missing_fields( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """JSON lines without conversation_id/message should use defaults.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - 'data: {"foo":"bar"}', - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - assert chunks_received[0]["conversation_id"] == "unknown" - assert chunks_received[0]["message"] == "" - - -def test_streaming_search_skips_done_marker( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """data: [DONE] lines should be ignored.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - 'data: {"conversation_id":"123","message":"Chunk"}', - "data: [DONE]", - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - - -def test_streaming_search_handles_json_array_lines( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """JSON arrays should be treated as plain text message chunks.""" - chunks_received: typing.List[MessageChunk] = [] - - def on_chunk(chunk: MessageChunk) -> None: - chunks_received.append(chunk) - - sse_lines = [ - 'data: ["a", "b"]', - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 1}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_chunk": on_chunk}, - } - ) - - assert len(chunks_received) == 1 - assert chunks_received[0]["message"] == '["a", "b"]' - - -def test_streaming_search_supports_builder( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test StreamConfigBuilder for streaming callbacks.""" - complete_calls: typing.List[int] = [] - - stream = StreamConfigBuilder() - - @stream.on_complete - def on_complete(response: typing.Mapping[str, JSONValue]) -> None: - found = response.get("found") - if isinstance(found, int): - complete_calls.append(found) - - sse_lines = [ - 'data: {"conversation_id":"123","message":"Hello"}', - '{"found": 1, "hits": [], "page": 1, "search_time_ms": 5}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": stream, - } - ) - - assert complete_calls == [1] - - -def test_stream_config_not_sent_to_api( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test that stream_config is removed from API params.""" - captured_params: typing.Dict[str, str] = {} - - sse_lines = [ - '{"found": 0, "hits": [], "page": 1, "search_time_ms": 1}', - ] - response = stream_response(lines=sse_lines) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - if params: - captured_params.update(params) - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - stream_config: StreamConfig[DocumentSchema] = {"on_chunk": lambda _: None} - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": stream_config, - } - ) - - assert "stream_config" not in captured_params - assert captured_params.get("conversation_stream") == "true" - - -def test_streaming_search_invokes_on_error( - fake_documents: Documents[DocumentSchema], - stream_response: type[FakeStreamResponse], - stream_context: type[FakeStreamContext], - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Test that streaming search invokes on_error for request failures.""" - errors: typing.List[BaseException] = [] - - def on_error(error: BaseException) -> None: - errors.append(error) - - fake_documents.api_call.config.num_retries = 0 - - response = stream_response( - lines=[], - status_code=500, - headers={"Content-Type": "application/json"}, - text='{"message": "Server error"}', - ) - - def fake_stream( - method: str, - url: str, - params: typing.Mapping[str, str] | None = None, - content: str | bytes | None = None, - headers: typing.Mapping[str, str] | None = None, - timeout: float | None = None, - ) -> FakeStreamContext: - return stream_context(response) - - monkeypatch.setattr( - fake_documents.api_call._client, - "stream", - fake_stream, - ) - - with pytest.raises(ServerError): - fake_documents.search( - { - "q": "test", - "query_by": "title", - "conversation_stream": True, - "stream_config": {"on_error": on_error}, - } - ) - - assert len(errors) == 1 - assert isinstance(errors[0], ServerError) diff --git a/utils/run-unasync.py b/utils/run-unasync.py index 7b11b6c..3c22283 100644 --- a/utils/run-unasync.py +++ b/utils/run-unasync.py @@ -31,8 +31,13 @@ def collect_class_replacements(source_dir: Path) -> dict[str, str]: replacements["AsyncConcurrencyLimit"] = "ConcurrencyLimit" # Defined in the shared ``typesense.http_backend`` module, outside async_. replacements["ASYNC_CLIENT_TYPES"] = "CLIENT_TYPES" - replacements["aiter_lines"] = "iter_lines" replacements["aread"] = "read" + # Defined in the shared ``typesense.sse`` module, outside async_. + replacements["aiter_events"] = "iter_events" + replacements["aiter_bytes"] = "iter_bytes" + # ``AsyncGenerator`` takes two type arguments, but ``Generator`` needs three + # before Python 3.13. + replacements["Generator[MessageChunk, None]"] = "Generator[MessageChunk, None, None]" return replacements From 7095aa6a1dd6089017a5207f34bc6ff6a32a3a8e Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:56:59 +0300 Subject: [PATCH 14/16] test(streaming): cover stream failover, cleanup and callbacks --- tests/stream_async_test.py | 346 +++++++++++++++++++++++++++++++ tests/stream_integration_test.py | 83 ++++++++ tests/stream_test.py | 326 +++++++++++++++++++++++++++++ tests/utils/streaming.py | 84 ++++++++ 4 files changed, 839 insertions(+) create mode 100644 tests/stream_async_test.py create mode 100644 tests/stream_integration_test.py create mode 100644 tests/stream_test.py create mode 100644 tests/utils/streaming.py diff --git a/tests/stream_async_test.py b/tests/stream_async_test.py new file mode 100644 index 0000000..95009ce --- /dev/null +++ b/tests/stream_async_test.py @@ -0,0 +1,346 @@ +"""Tests for streamed conversational search with the async client.""" + +import json +import sys + +import httpx +import pytest +import respx + +from tests.utils.streaming import ( + CHUNKS, + FINAL_RESPONSE, + SEARCH_URL, + MULTI_SEARCH_URL, + sse_body, + sse_response, +) +from typesense.configuration import Configuration +from typesense.exceptions import RequestMalformed, TypesenseClientError +from typesense.async_.api_call import AsyncApiCall +from typesense.async_.documents import AsyncDocuments +from typesense.async_.multi_search import AsyncMultiSearch +from typesense.types.document import MessageChunk, StreamConfigBuilder + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +SEARCH_PARAMS: typing.Final = { + "q": "who wrote it", + "query_by": "title", + "conversation_model_id": "conv-model", +} + + +@pytest.fixture(name="documents") +def documents_fixture(fake_async_api_call: AsyncApiCall) -> AsyncDocuments: + """Return the documents of a collection, sent through the fake API call.""" + return AsyncDocuments(fake_async_api_call, "books") + + +async def test_search_stream_yields_chunks_then_final_response( + documents: AsyncDocuments, +) -> None: + """Test that the stream yields each answer piece and keeps the search response.""" + with respx.mock: + route = respx.get(SEARCH_URL).mock(return_value=sse_response()) + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + chunks = [chunk async for chunk in stream] + final_response = await stream.get_final_response() + + assert chunks == CHUNKS + assert final_response == FINAL_RESPONSE + request = route.calls.last.request + assert request.headers["Accept"] == "text/event-stream" + assert request.url.params["conversation"] == "true" + assert request.url.params["conversation_stream"] == "true" + assert request.url.params["conversation_model_id"] == "conv-model" + + +async def test_get_final_response_reads_the_whole_stream( + documents: AsyncDocuments, +) -> None: + """Test that the search response can be read without iterating first.""" + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + assert await stream.get_final_response() == FINAL_RESPONSE + + +async def test_grouped_search_stream_keeps_final_response( + documents: AsyncDocuments, +) -> None: + """A grouped response has grouped_hits in place of hits.""" + grouped_response = { + **{key: value for key, value in FINAL_RESPONSE.items() if key != "hits"}, + "grouped_hits": [{"group_key": ["fiction"], "hits": FINAL_RESPONSE["hits"]}], + } + with respx.mock: + respx.get(SEARCH_URL).mock( + return_value=sse_response(final_response=grouped_response), + ) + + async with await documents.search_stream( + {**SEARCH_PARAMS, "group_by": "category"}, + ) as stream: + assert [chunk async for chunk in stream] == CHUNKS + assert await stream.get_final_response() == grouped_response + + +async def test_search_stream_uses_the_stream_read_timeout( + documents: AsyncDocuments, +) -> None: + """Test that streaming reads wait for ``stream_read_timeout_seconds``.""" + with respx.mock: + route = respx.get(SEARCH_URL).mock(return_value=sse_response()) + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + await stream.get_final_response() + + timeout = route.calls.last.request.extensions["timeout"] + assert timeout["read"] == 60.0 + assert timeout["connect"] == 0.001 + + +async def test_search_runs_stream_config_callbacks(documents: AsyncDocuments) -> None: + """Test that search runs the callbacks and returns the search response.""" + received: typing.List[object] = [] + + with respx.mock: + route = respx.get(SEARCH_URL).mock(return_value=sse_response()) + + response = await documents.search( + { + **SEARCH_PARAMS, + "conversation": True, + "conversation_stream": True, + "stream_config": { + "on_chunk": received.append, + "on_complete": received.append, + }, + }, + ) + + assert response == FINAL_RESPONSE + assert received == [*CHUNKS, FINAL_RESPONSE] + assert "stream_config" not in route.calls.last.request.url.params + + +async def test_search_accepts_a_stream_config_builder( + documents: AsyncDocuments, +) -> None: + """Test that callbacks registered on a builder run.""" + stream_config: StreamConfigBuilder[typing.Any] = StreamConfigBuilder() + messages: typing.List[str] = [] + + @stream_config.on_chunk + def on_chunk(chunk: MessageChunk) -> None: + messages.append(chunk["message"]) + + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + await documents.search( + { + **SEARCH_PARAMS, + "conversation_stream": True, + "stream_config": stream_config, + }, + ) + + assert "".join(messages) == "The Hobbit was written by Tolkien." + + +async def test_search_without_stream_config_returns_final_response( + documents: AsyncDocuments, +) -> None: + """Test that a streamed search with no callbacks returns the search response.""" + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + response = await documents.search( + {**SEARCH_PARAMS, "conversation_stream": True} + ) + + assert response == FINAL_RESPONSE + + +async def test_search_stream_fails_over_before_the_stream_starts( + fake_async_api_call: AsyncApiCall, + documents: AsyncDocuments, +) -> None: + """Test that a 5xx is retried on the next node, which is marked healthy.""" + node0_search_url = SEARCH_URL.replace("nearest", "node0") + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=httpx.Response(503, text="Down")) + respx.get(node0_search_url).mock(return_value=sse_response()) + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + final_response = await stream.get_final_response() + + assert len(respx.calls) == 2 + + assert final_response == FINAL_RESPONSE + assert fake_async_api_call.config.nearest_node is not None + assert fake_async_api_call.config.nearest_node.healthy is False + assert fake_async_api_call.config.nodes[0].healthy is True + + +async def test_errors_mid_stream_are_raised_without_retrying( + documents: AsyncDocuments, +) -> None: + """Test that a read error after the answer started is raised, not retried.""" + errors: typing.List[BaseException] = [] + received: typing.List[object] = [] + + async def body() -> typing.AsyncIterator[bytes]: + yield sse_body(CHUNKS[:1]) + raise httpx.ReadError("connection reset") + + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response(body())) + + with pytest.raises(httpx.ReadError): + await documents.search( + { + **SEARCH_PARAMS, + "conversation_stream": True, + "stream_config": { + "on_chunk": received.append, + "on_error": errors.append, + }, + }, + ) + + assert len(respx.calls) == 1 + + assert received == CHUNKS[:1] + assert len(errors) == 1 + assert isinstance(errors[0], httpx.ReadError) + + +async def test_stream_ending_without_search_response_raises( + documents: AsyncDocuments, +) -> None: + """Test that an error appended after the answer started is raised.""" + body = sse_body(CHUNKS) + b'{"message": "Conversation history is full."}' + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response(body)) + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + with pytest.raises(TypesenseClientError, match="history is full"): + await stream.get_final_response() + + +async def test_client_errors_are_raised_and_reported_once( + documents: AsyncDocuments, +) -> None: + """Test that a 400 with a plain-text body raises without failing over.""" + errors: typing.List[BaseException] = [] + with respx.mock: + respx.get(SEARCH_URL).mock( + return_value=httpx.Response(400, text="Conversation model not found"), + ) + + with pytest.raises(RequestMalformed, match="Conversation model not found"): + await documents.search( + { + **SEARCH_PARAMS, + "conversation_stream": True, + "stream_config": {"on_error": errors.append}, + }, + ) + + assert len(respx.calls) == 1 + + assert len(errors) == 1 + + +async def test_closing_early_releases_the_connection_and_slot( + fake_config: Configuration, +) -> None: + """Test that leaving the stream early closes the response and frees its slot.""" + fake_config.max_concurrent_requests = 1 + api_call = AsyncApiCall(fake_config) + documents = AsyncDocuments(api_call, "books") + + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + assert await stream.__anext__() == CHUNKS[0] + + async with await documents.search_stream(SEARCH_PARAMS) as second_stream: + assert await second_stream.get_final_response() == FINAL_RESPONSE + + assert stream.response.is_closed + with pytest.raises(TypesenseClientError, match="closed before"): + await stream.get_final_response() + + +async def test_multi_search_stream_sends_conversation_params_in_query( + fake_async_api_call: AsyncApiCall, +) -> None: + """Test that multi-search streams with the conversation in the query string.""" + multi_search_response = {"results": [FINAL_RESPONSE], "conversation": {}} + with respx.mock: + route = respx.post(MULTI_SEARCH_URL).mock( + return_value=sse_response(final_response=multi_search_response), + ) + + async with await AsyncMultiSearch(fake_async_api_call).perform_stream( + {"searches": [{"collection": "books", "query_by": "title"}]}, + {"q": "who wrote it", "conversation_model_id": "conv-model"}, + ) as stream: + chunks = [chunk async for chunk in stream] + final_response = await stream.get_final_response() + + assert chunks == CHUNKS + assert final_response == multi_search_response + request = route.calls.last.request + assert request.url.params["q"] == "who wrote it" + assert request.url.params["conversation_stream"] == "true" + assert json.loads(request.content)["searches"] == [ + {"collection": "books", "query_by": "title"}, + ] + + +async def test_multi_search_stream_accepts_a_json_response( + fake_async_api_call: AsyncApiCall, +) -> None: + """Test the plain JSON Typesense sends when every search fails.""" + multi_search_response = {"results": [{"code": 404, "error": "Not found."}]} + with respx.mock: + respx.post(MULTI_SEARCH_URL).mock( + return_value=httpx.Response(200, json=multi_search_response), + ) + + response = await AsyncMultiSearch(fake_async_api_call).perform( + {"searches": [{"collection": "missing", "query_by": "title"}]}, + {"q": "who", "conversation_model_id": "m", "conversation_stream": True}, + ) + + assert response == multi_search_response + + +async def test_search_stream_with_httpx2_client(fake_config: Configuration) -> None: + """Test streaming through a user-supplied httpx2 client.""" + httpx2 = pytest.importorskip("httpx2") + + def handler(request: typing.Any) -> typing.Any: + return httpx2.Response( + 200, + headers={"Content-Type": "text/event-stream"}, + content=sse_body(CHUNKS, FINAL_RESPONSE), + ) + + http_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) + documents = AsyncDocuments(AsyncApiCall(fake_config, http_client), "books") + + async with await documents.search_stream(SEARCH_PARAMS) as stream: + assert [chunk async for chunk in stream] == CHUNKS + assert await stream.get_final_response() == FINAL_RESPONSE diff --git a/tests/stream_integration_test.py b/tests/stream_integration_test.py new file mode 100644 index 0000000..4e2aab2 --- /dev/null +++ b/tests/stream_integration_test.py @@ -0,0 +1,83 @@ +"""Tests for streamed conversational search against a Typesense server and OpenAI.""" + +import pytest + +from typesense.async_.api_call import AsyncApiCall +from typesense.async_.documents import AsyncDocuments +from typesense.sync.api_call import ApiCall +from typesense.sync.documents import Documents +from typesense.sync.multi_search import MultiSearch + + +@pytest.mark.open_ai +def test_search_stream( + delete_all: None, + delete_all_conversations_models: None, + create_collection: None, + create_document: None, + create_conversations_model: str, + actual_api_call: ApiCall, +) -> None: + """Test that the streamed pieces make up the answer in the search response.""" + documents = Documents(actual_api_call, "companies") + + with documents.search_stream( + { + "q": "company", + "query_by": "company_name", + "conversation_model_id": create_conversations_model, + }, + ) as stream: + messages = [chunk["message"] for chunk in stream] + response = stream.get_final_response() + + assert messages + assert response["found"] == 1 + assert "".join(messages) == response["conversation"]["answer"] + + +@pytest.mark.open_ai +async def test_search_stream_async( + delete_all: None, + delete_all_conversations_models: None, + create_collection: None, + create_document: None, + create_conversations_model: str, + actual_async_api_call: AsyncApiCall, +) -> None: + """Test streaming with the async client.""" + documents = AsyncDocuments(actual_async_api_call, "companies") + + async with await documents.search_stream( + { + "q": "company", + "query_by": "company_name", + "conversation_model_id": create_conversations_model, + }, + ) as stream: + messages = [chunk["message"] async for chunk in stream] + response = await stream.get_final_response() + + assert messages + assert "".join(messages) == response["conversation"]["answer"] + + +@pytest.mark.open_ai +def test_multi_search_stream( + delete_all: None, + delete_all_conversations_models: None, + create_collection: None, + create_document: None, + create_conversations_model: str, + actual_api_call: ApiCall, +) -> None: + """Test that a streamed multi-search answers once, at the top level.""" + with MultiSearch(actual_api_call).perform_stream( + {"searches": [{"collection": "companies", "query_by": "company_name"}]}, + {"q": "company", "conversation_model_id": create_conversations_model}, + ) as stream: + messages = [chunk["message"] for chunk in stream] + response = stream.get_final_response() + + assert len(response["results"]) == 1 + assert "".join(messages) == response["conversation"]["answer"] diff --git a/tests/stream_test.py b/tests/stream_test.py new file mode 100644 index 0000000..212912e --- /dev/null +++ b/tests/stream_test.py @@ -0,0 +1,326 @@ +"""Tests for streamed conversational search with the sync client.""" + +import json +import sys + +import httpx +import pytest +import respx + +from tests.utils.streaming import ( + CHUNKS, + FINAL_RESPONSE, + SEARCH_URL, + MULTI_SEARCH_URL, + sse_body, + sse_response, +) +from typesense.configuration import Configuration +from typesense.exceptions import RequestMalformed, TypesenseClientError +from typesense.sync.api_call import ApiCall +from typesense.sync.documents import Documents +from typesense.sync.multi_search import MultiSearch +from typesense.types.document import MessageChunk, StreamConfigBuilder + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +SEARCH_PARAMS: typing.Final = { + "q": "who wrote it", + "query_by": "title", + "conversation_model_id": "conv-model", +} + + +@pytest.fixture(name="documents") +def documents_fixture(fake_api_call: ApiCall) -> Documents: + """Return the documents of a collection, sent through the fake API call.""" + return Documents(fake_api_call, "books") + + +def test_search_stream_yields_chunks_then_final_response( + documents: Documents, +) -> None: + """Test that the stream yields each answer piece and keeps the search response.""" + with respx.mock: + route = respx.get(SEARCH_URL).mock(return_value=sse_response()) + + with documents.search_stream(SEARCH_PARAMS) as stream: + chunks = list(stream) + final_response = stream.get_final_response() + + assert chunks == CHUNKS + assert final_response == FINAL_RESPONSE + request = route.calls.last.request + assert request.headers["Accept"] == "text/event-stream" + assert request.url.params["conversation"] == "true" + assert request.url.params["conversation_stream"] == "true" + assert request.url.params["conversation_model_id"] == "conv-model" + + +def test_get_final_response_reads_the_whole_stream(documents: Documents) -> None: + """Test that the search response can be read without iterating first.""" + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + with documents.search_stream(SEARCH_PARAMS) as stream: + assert stream.get_final_response() == FINAL_RESPONSE + + +def test_grouped_search_stream_keeps_final_response(documents: Documents) -> None: + """A grouped response has grouped_hits in place of hits.""" + grouped_response = { + **{key: value for key, value in FINAL_RESPONSE.items() if key != "hits"}, + "grouped_hits": [{"group_key": ["fiction"], "hits": FINAL_RESPONSE["hits"]}], + } + with respx.mock: + respx.get(SEARCH_URL).mock( + return_value=sse_response(final_response=grouped_response), + ) + + with documents.search_stream({**SEARCH_PARAMS, "group_by": "category"}) as stream: + assert list(stream) == CHUNKS + assert stream.get_final_response() == grouped_response + + +def test_search_stream_uses_the_stream_read_timeout(documents: Documents) -> None: + """Test that streaming reads wait for ``stream_read_timeout_seconds``.""" + with respx.mock: + route = respx.get(SEARCH_URL).mock(return_value=sse_response()) + + with documents.search_stream(SEARCH_PARAMS) as stream: + stream.get_final_response() + + timeout = route.calls.last.request.extensions["timeout"] + assert timeout["read"] == 60.0 + assert timeout["connect"] == 0.001 + + +def test_search_runs_stream_config_callbacks(documents: Documents) -> None: + """Test that search runs the callbacks and returns the search response.""" + received: typing.List[object] = [] + + with respx.mock: + route = respx.get(SEARCH_URL).mock(return_value=sse_response()) + + response = documents.search( + { + **SEARCH_PARAMS, + "conversation": True, + "conversation_stream": True, + "stream_config": { + "on_chunk": received.append, + "on_complete": received.append, + }, + }, + ) + + assert response == FINAL_RESPONSE + assert received == [*CHUNKS, FINAL_RESPONSE] + assert "stream_config" not in route.calls.last.request.url.params + + +def test_search_accepts_a_stream_config_builder(documents: Documents) -> None: + """Test that callbacks registered on a builder run.""" + stream_config: StreamConfigBuilder[typing.Any] = StreamConfigBuilder() + messages: typing.List[str] = [] + + @stream_config.on_chunk + def on_chunk(chunk: MessageChunk) -> None: + messages.append(chunk["message"]) + + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + documents.search( + { + **SEARCH_PARAMS, + "conversation_stream": True, + "stream_config": stream_config, + }, + ) + + assert "".join(messages) == "The Hobbit was written by Tolkien." + + +def test_search_without_stream_config_returns_final_response( + documents: Documents, +) -> None: + """Test that a streamed search with no callbacks returns the search response.""" + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + response = documents.search({**SEARCH_PARAMS, "conversation_stream": True}) + + assert response == FINAL_RESPONSE + + +def test_search_stream_fails_over_before_the_stream_starts( + fake_api_call: ApiCall, + documents: Documents, +) -> None: + """Test that a 5xx is retried on the next node, which is marked healthy.""" + node0_search_url = SEARCH_URL.replace("nearest", "node0") + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=httpx.Response(503, text="Down")) + respx.get(node0_search_url).mock(return_value=sse_response()) + + with documents.search_stream(SEARCH_PARAMS) as stream: + final_response = stream.get_final_response() + + assert len(respx.calls) == 2 + + assert final_response == FINAL_RESPONSE + assert fake_api_call.config.nearest_node is not None + assert fake_api_call.config.nearest_node.healthy is False + assert fake_api_call.config.nodes[0].healthy is True + + +def test_errors_mid_stream_are_raised_without_retrying(documents: Documents) -> None: + """Test that a read error after the answer started is raised, not retried.""" + errors: typing.List[BaseException] = [] + received: typing.List[object] = [] + + def body() -> typing.Iterator[bytes]: + yield sse_body(CHUNKS[:1]) + raise httpx.ReadError("connection reset") + + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response(body())) + + with pytest.raises(httpx.ReadError): + documents.search( + { + **SEARCH_PARAMS, + "conversation_stream": True, + "stream_config": { + "on_chunk": received.append, + "on_error": errors.append, + }, + }, + ) + + assert len(respx.calls) == 1 + + assert received == CHUNKS[:1] + assert len(errors) == 1 + assert isinstance(errors[0], httpx.ReadError) + + +def test_stream_ending_without_search_response_raises(documents: Documents) -> None: + """Test that an error appended after the answer started is raised.""" + body = sse_body(CHUNKS) + b'{"message": "Conversation history is full."}' + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response(body)) + + with documents.search_stream(SEARCH_PARAMS) as stream: + with pytest.raises(TypesenseClientError, match="history is full"): + stream.get_final_response() + + +def test_client_errors_are_raised_and_reported_once(documents: Documents) -> None: + """Test that a 400 with a plain-text body raises without failing over.""" + errors: typing.List[BaseException] = [] + with respx.mock: + respx.get(SEARCH_URL).mock( + return_value=httpx.Response(400, text="Conversation model not found"), + ) + + with pytest.raises(RequestMalformed, match="Conversation model not found"): + documents.search( + { + **SEARCH_PARAMS, + "conversation_stream": True, + "stream_config": {"on_error": errors.append}, + }, + ) + + assert len(respx.calls) == 1 + + assert len(errors) == 1 + + +def test_closing_early_releases_the_connection_and_slot( + fake_config: Configuration, +) -> None: + """Test that leaving the stream early closes the response and frees its slot.""" + fake_config.max_concurrent_requests = 1 + api_call = ApiCall(fake_config) + documents = Documents(api_call, "books") + + with respx.mock: + respx.get(SEARCH_URL).mock(return_value=sse_response()) + + with documents.search_stream(SEARCH_PARAMS) as stream: + assert next(iter(stream)) == CHUNKS[0] + + with documents.search_stream(SEARCH_PARAMS) as second_stream: + assert second_stream.get_final_response() == FINAL_RESPONSE + + assert stream.response.is_closed + with pytest.raises(TypesenseClientError, match="closed before"): + stream.get_final_response() + + +def test_multi_search_stream_sends_conversation_params_in_query( + fake_api_call: ApiCall, +) -> None: + """Test that multi-search streams with the conversation in the query string.""" + multi_search_response = {"results": [FINAL_RESPONSE], "conversation": {}} + with respx.mock: + route = respx.post(MULTI_SEARCH_URL).mock( + return_value=sse_response(final_response=multi_search_response), + ) + + with MultiSearch(fake_api_call).perform_stream( + {"searches": [{"collection": "books", "query_by": "title"}]}, + {"q": "who wrote it", "conversation_model_id": "conv-model"}, + ) as stream: + chunks = list(stream) + final_response = stream.get_final_response() + + assert chunks == CHUNKS + assert final_response == multi_search_response + request = route.calls.last.request + assert request.url.params["q"] == "who wrote it" + assert request.url.params["conversation_stream"] == "true" + assert json.loads(request.content)["searches"] == [ + {"collection": "books", "query_by": "title"}, + ] + + +def test_multi_search_stream_accepts_a_json_response(fake_api_call: ApiCall) -> None: + """Test the plain JSON Typesense sends when every search fails.""" + multi_search_response = {"results": [{"code": 404, "error": "Not found."}]} + with respx.mock: + respx.post(MULTI_SEARCH_URL).mock( + return_value=httpx.Response(200, json=multi_search_response), + ) + + response = MultiSearch(fake_api_call).perform( + {"searches": [{"collection": "missing", "query_by": "title"}]}, + {"q": "who", "conversation_model_id": "m", "conversation_stream": True}, + ) + + assert response == multi_search_response + + +def test_search_stream_with_httpx2_client(fake_config: Configuration) -> None: + """Test streaming through a user-supplied httpx2 client.""" + httpx2 = pytest.importorskip("httpx2") + + def handler(request: typing.Any) -> typing.Any: + return httpx2.Response( + 200, + headers={"Content-Type": "text/event-stream"}, + content=sse_body(CHUNKS, FINAL_RESPONSE), + ) + + http_client = httpx2.Client(transport=httpx2.MockTransport(handler)) + documents = Documents(ApiCall(fake_config, http_client), "books") + + with documents.search_stream(SEARCH_PARAMS) as stream: + assert list(stream) == CHUNKS + assert stream.get_final_response() == FINAL_RESPONSE diff --git a/tests/utils/streaming.py b/tests/utils/streaming.py new file mode 100644 index 0000000..07da931 --- /dev/null +++ b/tests/utils/streaming.py @@ -0,0 +1,84 @@ +"""Builders for the server-sent event streams Typesense sends.""" + +import json +import sys + +import httpx + +from typesense.types.document import MessageChunk + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + +SEARCH_URL: typing.Final = "http://nearest:8108/collections/books/documents/search" +MULTI_SEARCH_URL: typing.Final = "http://nearest:8108/multi_search" + +CONVERSATION_ID: typing.Final = "6f1c0e5a" + +CHUNKS: typing.Final[typing.List[MessageChunk]] = [ + {"conversation_id": CONVERSATION_ID, "message": "The"}, + {"conversation_id": CONVERSATION_ID, "message": " Hobbit was"}, + {"conversation_id": CONVERSATION_ID, "message": " written by Tolkien."}, +] + +FINAL_RESPONSE: typing.Final[typing.Dict[str, typing.Any]] = { + "conversation": { + "answer": "The Hobbit was written by Tolkien.", + "conversation_history": {"conversation": []}, + "conversation_id": CONVERSATION_ID, + "query": "who wrote it", + }, + "facet_counts": [], + "found": 1, + "hits": [{"document": {"id": "0", "title": "The Hobbit"}}], + "out_of": 1, + "page": 1, + "search_time_ms": 2, +} + + +def sse_body( + chunks: typing.Sequence[MessageChunk], + final_response: typing.Optional[typing.Mapping[str, typing.Any]] = None, +) -> bytes: + """ + Build a stream like Typesense's: the answer pieces, ``[DONE]``, then the response. + + Args: + chunks (Sequence[MessageChunk]): The answer pieces. + final_response (Mapping | None): The search response, or ``None`` to end + the stream after the answer pieces. + + Returns: + bytes: The response body. + """ + events = [json.dumps(chunk) for chunk in chunks] + if final_response is not None: + events.extend(["[DONE]", json.dumps(final_response)]) + return "".join(f"data: {event}\n\n" for event in events).encode() + + +def sse_response( + body: typing.Union[ + bytes, typing.Iterator[bytes], typing.AsyncIterator[bytes], None + ] = None, + final_response: typing.Mapping[str, typing.Any] = FINAL_RESPONSE, +) -> httpx.Response: + """ + Build a ``text/event-stream`` response, by default the full search stream. + + Args: + body (bytes | Iterator[bytes] | AsyncIterator[bytes] | None): The body, + or ``None`` for the answer pieces followed by ``final_response``. + final_response (Mapping): The search response sent after the answer. + + Returns: + httpx.Response: The response. + """ + return httpx.Response( + 200, + headers={"Content-Type": "text/event-stream; charset=utf-8"}, + content=sse_body(CHUNKS, final_response) if body is None else body, + ) From 7fdb0fa9b9b33b17c45da776d05f7c9ab7727d10 Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:56:59 +0300 Subject: [PATCH 15/16] docs(examples): show iterating a search stream --- examples/async_conversation_streaming.py | 34 ++++++++++++++++-------- examples/conversation_streaming.py | 21 +++++++-------- 2 files changed, 33 insertions(+), 22 deletions(-) diff --git a/examples/async_conversation_streaming.py b/examples/async_conversation_streaming.py index 6098c5a..af682df 100644 --- a/examples/async_conversation_streaming.py +++ b/examples/async_conversation_streaming.py @@ -10,7 +10,11 @@ import typesense -from typesense.types.document import MessageChunk, StreamConfigBuilder +from typesense.types.document import ( + MessageChunk, + SearchResponse, + StreamConfigBuilder, +) def require_env(name: str) -> str: @@ -107,24 +111,32 @@ async def main() -> None: } documents = client.collections[documents_collection].documents - @stream.on_chunk + # Iterate over the answer as it is generated, then read the search response. + async with await documents.search_stream(search_parameters) as answer_stream: + async for chunk in answer_stream: + print(chunk["message"], end="", flush=True) + response = await answer_stream.get_final_response() + print("\n---\nFound", response["found"], "documents") + + # Or pass callbacks to search(), which returns the search response at the end. + stream_config: StreamConfigBuilder[SearchResponse[typing.Any]] = ( + StreamConfigBuilder() + ) + + @stream_config.on_chunk def on_chunk(chunk: MessageChunk) -> None: print(chunk["message"], end="", flush=True) - @stream.on_complete - def on_complete(response: dict) -> None: + @stream_config.on_complete + def on_complete(response: SearchResponse[typing.Any]) -> None: print("\n---\nComplete response keys:", response.keys()) - await client.collections["streaming_docs"].documents.search( + await documents.search( { - "q": "What is this document about?", - "query_by": "embedding", - "exclude_fields": "embedding", + **search_parameters, "conversation": True, - "prefix": False, "conversation_stream": True, - "conversation_model_id": conversation_model["id"], - "stream_config": stream, + "stream_config": stream_config, } ) finally: diff --git a/examples/conversation_streaming.py b/examples/conversation_streaming.py index addb9d4..5fb06ec 100644 --- a/examples/conversation_streaming.py +++ b/examples/conversation_streaming.py @@ -1,4 +1,3 @@ -from operator import truediv import os import sys import typing @@ -10,7 +9,11 @@ import typesense -from typesense.types.document import MessageChunk, StreamConfigBuilder +from typesense.types.document import ( + MessageChunk, + SearchResponse, + StreamConfigBuilder, +) def require_env(name: str) -> str: @@ -117,25 +120,21 @@ def require_env(name: str) -> str: stream_config: StreamConfigBuilder[SearchResponse[typing.Any]] = StreamConfigBuilder() -@stream.on_chunk +@stream_config.on_chunk def on_chunk(chunk: MessageChunk) -> None: print(chunk["message"], end="", flush=True) -@stream.on_complete -def on_complete(response: dict) -> None: +@stream_config.on_complete +def on_complete(response: SearchResponse[typing.Any]) -> None: print("\n---\nComplete response keys:", response.keys()) client.collections[documents_collection].documents.search( { - "q": "What is this document about?", - "query_by": "embedding", - "exclude_fields": "embedding", + **search_parameters, "conversation": True, - "prefix": False, "conversation_stream": True, - "conversation_model_id": conversation_model["id"], - "stream_config": stream, + "stream_config": stream_config, } ) From c96ae1262cef9b850c35dcda70276d3739e96ddc Mon Sep 17 00:00:00 2001 From: Fanis Tharropoulos Date: Wed, 7 Oct 2026 13:59:14 +0300 Subject: [PATCH 16/16] fix(streaming): open streams with client.stream so httpx2 type-checks --- src/typesense/async_/api_call.py | 49 ++++++++++++++++++-------------- src/typesense/async_/stream.py | 16 +++++------ src/typesense/sync/api_call.py | 49 ++++++++++++++++++-------------- src/typesense/sync/stream.py | 16 +++++------ utils/run-unasync.py | 3 ++ 5 files changed, 71 insertions(+), 62 deletions(-) diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index f2e27b2..73b7c3f 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -33,6 +33,7 @@ import asyncio import sys +from contextlib import AsyncExitStack from types import MappingProxyType, TracebackType import httpx @@ -54,6 +55,7 @@ from typesense.http_backend import ( ASYNC_CLIENT_TYPES, AsyncClientType, + ResponseType, backend_errors, verify_option, ) @@ -648,38 +650,41 @@ async def _open_stream( headers = request_kwargs.get("headers", {}) headers["Accept"] = "text/event-stream" timeout = self._client.timeout - request = self._client.build_request( - method, - url, - params=typing.cast( - typing.Optional[_QueryParams], - request_kwargs.get("params"), - ), - content=request_kwargs.get("content"), - headers=headers, - timeout=( - timeout.connect, - self.config.stream_read_timeout_seconds, - timeout.write, - timeout.pool, - ), + # Annotated so httpx and httpx2 responses unify as ``ResponseType``. + response_context: typing.AsyncContextManager[ResponseType] = ( + self._client.stream( + method, + url, + params=typing.cast( + typing.Optional[_QueryParams], + request_kwargs.get("params"), + ), + content=request_kwargs.get("content"), + headers=headers, + timeout=( + timeout.connect, + self.config.stream_read_timeout_seconds, + timeout.write, + timeout.pool, + ), + ) ) + # Owns the concurrency slot and the response until the stream is closed. + exit_stack = AsyncExitStack() await self._concurrency_limit.acquire() + exit_stack.callback(self._concurrency_limit.release) try: - response = await self._client.send(request, stream=True) + response = await exit_stack.enter_async_context(response_context) if response.status_code < 200 or response.status_code >= 300: - try: - await response.aread() - finally: - await response.aclose() + await response.aread() self.request_handler.raise_for_status(response) except BaseException: - self._concurrency_limit.release() + await exit_stack.aclose() raise self.node_manager.set_node_health(node, is_healthy=True) - return AsyncSearchStream(response, self._concurrency_limit.release) + return AsyncSearchStream(response, exit_stack) def _prepare_request_params( self, diff --git a/src/typesense/async_/stream.py b/src/typesense/async_/stream.py index e84ec3b..51b9a1a 100644 --- a/src/typesense/async_/stream.py +++ b/src/typesense/async_/stream.py @@ -14,6 +14,7 @@ """ import sys +from contextlib import AsyncExitStack from types import TracebackType from typesense.exceptions import TypesenseClientError @@ -46,7 +47,7 @@ class AsyncSearchStream(typing.Generic[TFinal]): def __init__( self, response: ResponseType, - on_close: typing.Callable[[], None], + exit_stack: AsyncExitStack, ) -> None: """ Initialize the stream. @@ -54,11 +55,11 @@ def __init__( Args: response (httpx.Response | httpx2.Response): A successful response opened with ``stream=True``. - on_close (Callable[[], None]): Called once when the stream is closed, - to release the request's concurrency slot. + exit_stack (AsyncExitStack): Closes the response and releases the + request's concurrency slot when the stream is closed. """ self.response = response - self._on_close = on_close + self._exit_stack = exit_stack self._closed = False self._final: typing.Optional[TFinal] = None self._decoder = SSEDecoder() @@ -112,14 +113,11 @@ async def aclose(self) -> None: await self._close_response() async def _close_response(self) -> None: - """Close the response once, then run ``on_close``.""" + """Close the response and release its concurrency slot, once.""" if self._closed: return self._closed = True - try: - await self.response.aclose() - finally: - self._on_close() + await self._exit_stack.aclose() async def _iter_chunks(self) -> typing.AsyncGenerator[MessageChunk, None]: """Yield the answer pieces and keep the final search response.""" diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index beb0e3c..817c578 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -33,6 +33,7 @@ import time import sys +from contextlib import ExitStack from types import MappingProxyType, TracebackType import httpx @@ -54,6 +55,7 @@ from typesense.http_backend import ( CLIENT_TYPES, SyncClientType, + ResponseType, backend_errors, verify_option, ) @@ -648,38 +650,41 @@ def _open_stream( headers = request_kwargs.get("headers", {}) headers["Accept"] = "text/event-stream" timeout = self._client.timeout - request = self._client.build_request( - method, - url, - params=typing.cast( - typing.Optional[_QueryParams], - request_kwargs.get("params"), - ), - content=request_kwargs.get("content"), - headers=headers, - timeout=( - timeout.connect, - self.config.stream_read_timeout_seconds, - timeout.write, - timeout.pool, - ), + # Annotated so httpx and httpx2 responses unify as ``ResponseType``. + response_context: typing.ContextManager[ResponseType] = ( + self._client.stream( + method, + url, + params=typing.cast( + typing.Optional[_QueryParams], + request_kwargs.get("params"), + ), + content=request_kwargs.get("content"), + headers=headers, + timeout=( + timeout.connect, + self.config.stream_read_timeout_seconds, + timeout.write, + timeout.pool, + ), + ) ) + # Owns the concurrency slot and the response until the stream is closed. + exit_stack = ExitStack() self._concurrency_limit.acquire() + exit_stack.callback(self._concurrency_limit.release) try: - response = self._client.send(request, stream=True) + response = exit_stack.enter_context(response_context) if response.status_code < 200 or response.status_code >= 300: - try: - response.read() - finally: - response.close() + response.read() self.request_handler.raise_for_status(response) except BaseException: - self._concurrency_limit.release() + exit_stack.close() raise self.node_manager.set_node_health(node, is_healthy=True) - return SearchStream(response, self._concurrency_limit.release) + return SearchStream(response, exit_stack) def _prepare_request_params( self, diff --git a/src/typesense/sync/stream.py b/src/typesense/sync/stream.py index 91e2a4b..bbe6f5a 100644 --- a/src/typesense/sync/stream.py +++ b/src/typesense/sync/stream.py @@ -14,6 +14,7 @@ """ import sys +from contextlib import ExitStack from types import TracebackType from typesense.exceptions import TypesenseClientError @@ -46,7 +47,7 @@ class SearchStream(typing.Generic[TFinal]): def __init__( self, response: ResponseType, - on_close: typing.Callable[[], None], + exit_stack: ExitStack, ) -> None: """ Initialize the stream. @@ -54,11 +55,11 @@ def __init__( Args: response (httpx.Response | httpx2.Response): A successful response opened with ``stream=True``. - on_close (Callable[[], None]): Called once when the stream is closed, - to release the request's concurrency slot. + exit_stack (ExitStack): Closes the response and releases the + request's concurrency slot when the stream is closed. """ self.response = response - self._on_close = on_close + self._exit_stack = exit_stack self._closed = False self._final: typing.Optional[TFinal] = None self._decoder = SSEDecoder() @@ -112,14 +113,11 @@ def close(self) -> None: self._close_response() def _close_response(self) -> None: - """Close the response once, then run ``on_close``.""" + """Close the response and release its concurrency slot, once.""" if self._closed: return self._closed = True - try: - self.response.close() - finally: - self._on_close() + self._exit_stack.close() def _iter_chunks(self) -> typing.Generator[MessageChunk, None, None]: """Yield the answer pieces and keep the final search response.""" diff --git a/utils/run-unasync.py b/utils/run-unasync.py index 3c22283..144ebe6 100644 --- a/utils/run-unasync.py +++ b/utils/run-unasync.py @@ -35,6 +35,9 @@ def collect_class_replacements(source_dir: Path) -> dict[str, str]: # Defined in the shared ``typesense.sse`` module, outside async_. replacements["aiter_events"] = "iter_events" replacements["aiter_bytes"] = "iter_bytes" + replacements["AsyncExitStack"] = "ExitStack" + replacements["AsyncContextManager"] = "ContextManager" + replacements["enter_async_context"] = "enter_context" # ``AsyncGenerator`` takes two type arguments, but ``Generator`` needs three # before Python 3.13. replacements["Generator[MessageChunk, None]"] = "Generator[MessageChunk, None, None]"