diff --git a/examples/async_conversation_streaming.py b/examples/async_conversation_streaming.py new file mode 100644 index 0000000..af682df --- /dev/null +++ b/examples/async_conversation_streaming.py @@ -0,0 +1,147 @@ +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, + SearchResponse, + 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 + + # 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_config.on_complete + def on_complete(response: SearchResponse[typing.Any]) -> None: + print("\n---\nComplete response keys:", response.keys()) + + await documents.search( + { + **search_parameters, + "conversation": True, + "conversation_stream": True, + "stream_config": stream_config, + } + ) + 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..5fb06ec --- /dev/null +++ b/examples/conversation_streaming.py @@ -0,0 +1,140 @@ +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, + SearchResponse, + 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_config.on_chunk +def on_chunk(chunk: MessageChunk) -> None: + print(chunk["message"], end="", flush=True) + + +@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( + { + **search_parameters, + "conversation": True, + "conversation_stream": True, + "stream_config": stream_config, + } +) diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index bfaf6de..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,11 +55,13 @@ from typesense.http_backend import ( ASYNC_CLIENT_TYPES, AsyncClientType, + ResponseType, backend_errors, verify_option, ) +from .stream import AsyncSearchStream from typesense.node_manager import NodeManager -from typesense.request_handler import RequestHandler +from typesense.request_handler import RequestHandler, _QueryParams if sys.version_info >= (3, 11): import typing @@ -559,6 +562,130 @@ async def _make_request_and_process_response( else typing.cast(str, request_response) ) + 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], + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> 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 + # 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 exit_stack.enter_async_context(response_context) + if response.status_code < 200 or response.status_code >= 300: + await response.aread() + self.request_handler.raise_for_status(response) + except BaseException: + await exit_stack.aclose() + raise + + self.node_manager.set_node_health(node, is_healthy=True) + return AsyncSearchStream(response, exit_stack) + def _prepare_request_params( self, endpoint: str, diff --git a/src/typesense/async_/documents.py b/src/typesense/async_/documents.py index 399c82d..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 @@ -360,13 +366,38 @@ 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. Returns: SearchResponse[TDoc]: The search response containing matching documents. """ - stringified_search_params = stringify_search_params(search_parameters) + 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, @@ -375,6 +406,46 @@ async def search(self, search_parameters: SearchParameters) -> SearchResponse[TD ) 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..51b9a1a --- /dev/null +++ b/src/typesense/async_/stream.py @@ -0,0 +1,218 @@ +""" +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 contextlib import AsyncExitStack +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, + exit_stack: AsyncExitStack, + ) -> None: + """ + Initialize the stream. + + Args: + response (httpx.Response | httpx2.Response): A successful response + opened with ``stream=True``. + exit_stack (AsyncExitStack): Closes the response and releases the + request's concurrency slot when the stream is closed. + """ + self.response = response + self._exit_stack = exit_stack + 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 and release its concurrency slot, once.""" + if self._closed: + return + self._closed = True + await self._exit_stack.aclose() + + 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/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() 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/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: """ 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/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index 0ffd4bc..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,11 +55,13 @@ from typesense.http_backend import ( CLIENT_TYPES, SyncClientType, + ResponseType, backend_errors, verify_option, ) +from .stream import SearchStream from typesense.node_manager import NodeManager -from typesense.request_handler import RequestHandler +from typesense.request_handler import RequestHandler, _QueryParams if sys.version_info >= (3, 11): import typing @@ -559,6 +562,130 @@ def _make_request_and_process_response( else typing.cast(str, request_response) ) + 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], + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> 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 + # 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 = exit_stack.enter_context(response_context) + if response.status_code < 200 or response.status_code >= 300: + response.read() + self.request_handler.raise_for_status(response) + except BaseException: + exit_stack.close() + raise + + self.node_manager.set_node_health(node, is_healthy=True) + return SearchStream(response, exit_stack) + def _prepare_request_params( self, endpoint: str, diff --git a/src/typesense/sync/documents.py b/src/typesense/sync/documents.py index 0c7d7f7..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 @@ -360,13 +366,38 @@ 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. Returns: SearchResponse[TDoc]: The search response containing matching documents. """ - stringified_search_params = stringify_search_params(search_parameters) + 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, @@ -375,6 +406,46 @@ def search(self, search_parameters: SearchParameters) -> SearchResponse[TDoc]: ) 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..bbe6f5a --- /dev/null +++ b/src/typesense/sync/stream.py @@ -0,0 +1,218 @@ +""" +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 contextlib import ExitStack +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, + exit_stack: ExitStack, + ) -> None: + """ + Initialize the stream. + + Args: + response (httpx.Response | httpx2.Response): A successful response + opened with ``stream=True``. + exit_stack (ExitStack): Closes the response and releases the + request's concurrency slot when the stream is closed. + """ + self.response = response + self._exit_stack = exit_stack + 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 and release its concurrency slot, once.""" + if self._closed: + return + self._closed = True + self._exit_stack.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 ee44b04..b5cf565 100644 --- a/src/typesense/types/document.py +++ b/src/typesense/types/document.py @@ -586,6 +586,114 @@ class NLLanguageParameters(typing.TypedDict): nl_query_debug: typing.NotRequired[bool] +TFinal = typing.TypeVar("TFinal") + + +class ConversationParameters(typing.TypedDict): + """ + Parameters for [conversational search](https://typesense.org/docs/29.0/api/conversational-search-rag.html). + + 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 MessageChunk(typing.TypedDict): + """ + A piece of a streamed conversation answer. + + Attributes: + conversation_id (str): The ID of the conversation. + message (str): The next piece of the answer. + """ + + conversation_id: str + message: str + + +OnChunkCallback = typing.Callable[[MessageChunk], None] +OnErrorCallback = typing.Callable[[BaseException], 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. + """ + + on_chunk: OnChunkCallback + on_complete: typing.Callable[[TFinal], None] + on_error: OnErrorCallback + + +class StreamConfigBuilder(typing.Generic[TFinal]): + """ + Build a ``StreamConfig`` by registering callbacks with decorators. + + Example: + >>> stream = StreamConfigBuilder() + >>> + >>> @stream.on_chunk + ... def handle_chunk(chunk: MessageChunk) -> None: + ... print(chunk["message"], end="", flush=True) + >>> + >>> 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 a builder with no callbacks.""" + self._config: StreamConfig[TFinal] = {} + + def on_chunk(self, func: OnChunkCallback) -> OnChunkCallback: + """Register ``func`` to be called with each piece of the answer.""" + self._config["on_chunk"] = func + return 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: + """Register ``func`` to be called with the error if the search fails.""" + self._config["on_error"] = func + return func + + def build(self) -> StreamConfig[TFinal]: + """Return the registered callbacks as a ``StreamConfig``.""" + return self._config.copy() + + class SearchParameters( RequiredSearchParameters, QueryParameters, @@ -598,6 +706,7 @@ class SearchParameters( TypoToleranceParameters, CachingParameters, NLLanguageParameters, + ConversationParameters, ): """Parameters for searching documents.""" @@ -626,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. 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/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, 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"] 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, + ) diff --git a/utils/run-unasync.py b/utils/run-unasync.py index 5dd8816..144ebe6 100644 --- a/utils/run-unasync.py +++ b/utils/run-unasync.py @@ -31,6 +31,16 @@ 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["aread"] = "read" + # 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]" return replacements