diff --git a/.github/workflows/pypi.yml b/.github/workflows/pypi.yml index ad055038..182584e9 100644 --- a/.github/workflows/pypi.yml +++ b/.github/workflows/pypi.yml @@ -23,4 +23,4 @@ jobs: run: hatch build - name: Publish on PyPi - uses: pypa/gh-action-pypi-publish@ed0c53931b1dc9bd32cbe73a98c7f6766f8a527e # v1.13.0 + uses: pypa/gh-action-pypi-publish@dc37677b2e1c63e2034f94d8a5b11f265b73ba33 # v1.14.2 diff --git a/docs/concepts/pipeline-wrapper.md b/docs/concepts/pipeline-wrapper.md index c21cd4ce..60c4375b 100644 --- a/docs/concepts/pipeline-wrapper.md +++ b/docs/concepts/pipeline-wrapper.md @@ -30,6 +30,44 @@ class PipelineWrapper(BasePipelineWrapper): return result["llm"]["replies"][0] ``` +## Request headers + +To access HTTP request headers, explicitly declare `headers: Headers | None = None` +on any of `run_api`, `run_chat_completion`, `run_response`, or their async variants: + +```python +from fastapi import HTTPException +from starlette.datastructures import Headers + +def run_api(self, query: str, *, headers: Headers | None = None) -> str: + authorization = headers.get("authorization") if headers is not None else None + if not authorization: + raise HTTPException(status_code=401, detail="Authorization is required") + # Application code: validate the credential or forward it only to your trusted service. + return self.answer(query, authorization=authorization) +``` + +Import `Headers` at runtime, including when using postponed annotations. The +parameter must accept keyword arguments and have the annotation `Headers | None` +and default `None`; invalid declarations fail during deployment. Headers are +request-local and available throughout streaming. Lookups are case-insensitive. +The OpenAI-compatible endpoints receive a header dictionary from +`fastapi-openai-compat`, so repeated values for a header are not preserved there. + +The injected parameter is excluded from the `/run` JSON or multipart body and from +OpenAPI and MCP input schemas. A body field cannot override it. Existing application +parameters named `headers` with other types (such as `dict[str, str]`) keep their +normal behavior. Declaring only `**kwargs` does not enable injection. + +A2A, MCP and direct calls do not automatically supply HTTP headers, so the default +is `None`. MCP tool arguments cannot populate the injected parameter. Wrappers +that require authentication must handle missing credentials explicitly. Receiving +headers does not itself authenticate the caller. + +Do not store caller credentials on the shared wrapper instance, echo them in +responses, or include them in logs or traces. Hayhooks keeps injected headers out +of its automatic `/run` payload logs and trace tags. + ## Required Methods ### setup() @@ -131,7 +169,8 @@ def run_api(self, urls: list[str], question: str) -> str: **Input argument rules:** -- Arguments must be JSON-serializable +- Body arguments must be JSON-serializable, except [file uploads](#file-upload-support) +- The typed [`headers` parameter](#request-headers) is injected separately from the body - Use proper type hints (`list[str]`, `int | None`, etc.) - Default values are supported - Complex types like `dict[str, Any]` are allowed @@ -185,6 +224,25 @@ async def run_api_async(self, query: str) -> AsyncGenerator: ) ``` +##### Client disconnects + +By default, closing the response stream or cancelling its consumer also cancels the pipeline task. Cancellation still +propagates to the request handler so HTTP cleanup can finish normally. + +Set `shield_pipeline_task=True` when the request should stop but the pipeline must continue running: + +```python +return async_streaming_generator( + pipeline=self.pipeline, + pipeline_run_args={"prompt": {"query": query}}, + shield_pipeline_task=True, +) +``` + +Shielding only changes cleanup of the pipeline task: it does not suppress cancellation in the request handler, and no +more chunks can be delivered to a disconnected client. Hayhooks keeps the detached task alive until it finishes and +logs any eventual exception. + When a generator is detected, Hayhooks automatically wraps it in a FastAPI `StreamingResponse` using the `text/plain` media type. The behaviour is identical for both `run_api()` and `run_api_async()`—the only difference is whether the underlying generator is sync or async. | Method | What you return | Response media type | Notes | diff --git a/docs/features/a2a-support.md b/docs/features/a2a-support.md index e0da2f17..50442d3f 100644 --- a/docs/features/a2a-support.md +++ b/docs/features/a2a-support.md @@ -155,6 +155,7 @@ See [examples/a2a_multi_agent](https://github.com/deepset-ai/hayhooks/tree/main/ ## Current limitations +- **Request headers**: A2A calls leave the optional typed [`headers` parameter](../concepts/pipeline-wrapper.md#request-headers) at `None`; transport headers are not forwarded to the wrapper. - **Request-bound task execution**: Hayhooks currently treats A2A as a chat-shaped bridge. Each task runs inside the request handler by calling `run_chat_completion` / `run_chat_completion_async`, so non-streaming `SendMessage` returns after the task has completed or failed. This means detached task execution via [`returnImmediately`](https://a2a-protocol.org/latest/specification/#322-sendmessageconfiguration), [`input-required`](https://a2a-protocol.org/latest/specification/#63-multi-turn-interaction) pauses, and [push notification delivery](https://a2a-protocol.org/latest/specification/#353-push-notification-delivery) are not supported yet. - **Static agents list**: A2A routes are built from the registry at startup. Pipelines deployed or undeployed at runtime require restarting `hayhooks a2a run`. - **In-memory task store**: task state is kept in memory and lost on restart. diff --git a/docs/features/mcp-support.md b/docs/features/mcp-support.md index 69ad4689..e68046b5 100644 --- a/docs/features/mcp-support.md +++ b/docs/features/mcp-support.md @@ -129,7 +129,11 @@ For each deployed pipeline, Hayhooks will: - Parse **`run_api` method docstring**: - If you use Google-style or reStructuredText-style docstrings, use the first line as MCP Tool `description` and the rest as `parameters` (if present) - Each parameter description will be used as the `description` of the corresponding Pydantic model field (if present) -- Generate a Pydantic model from the `inputSchema` using the **`run_api` method arguments as fields** +- Generate the `inputSchema` from the **`run_api` method arguments**, excluding the optional typed [`headers` parameter](../concepts/pipeline-wrapper.md#request-headers) + +MCP calls leave the typed `headers` parameter at `None`, including over HTTP transports. +Supplying it as a tool argument is rejected. Ordinary application fields named `headers` +(for example, `headers: dict[str, str]`) remain normal tool arguments. **Example:** diff --git a/docs/features/openai-compatibility.md b/docs/features/openai-compatibility.md index 99137c31..65df1c91 100644 --- a/docs/features/openai-compatibility.md +++ b/docs/features/openai-compatibility.md @@ -16,6 +16,10 @@ Hayhooks supports two OpenAI API surfaces: Both APIs are available simultaneously. A pipeline wrapper can implement one or both. +To access caller credentials or other HTTP headers, declare the optional typed +[`headers` parameter](../concepts/pipeline-wrapper.md#request-headers) on your wrapper method. +This requires `fastapi-openai-compat>=1.4.0`, which Hayhooks declares as a dependency. + ## Key Features - **Automatic Endpoint Generation**: OpenAI-compatible endpoints are created automatically diff --git a/pyproject.toml b/pyproject.toml index 7a8e0dad..bad1ba28 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -37,7 +37,7 @@ dependencies = [ "pydantic-settings", "python-dotenv", "docstring-parser", - "fastapi-openai-compat>=1.2.0", + "fastapi-openai-compat>=1.4.0", ] [project.optional-dependencies] @@ -76,8 +76,6 @@ source = "vcs" packages = ["src/hayhooks"] [tool.hatch.build.targets.wheel.force-include] -# .chainlit is a hidden dir excluded by default, so force-include is needed -"src/hayhooks/server/chainlit_app/.chainlit" = "hayhooks/server/chainlit_app/.chainlit" # Bundle dashboard source files so --with-tracing-dashboard can build locally at runtime. "dashboard/index.html" = "hayhooks/dashboard/index.html" "dashboard/components.json" = "hayhooks/dashboard/components.json" @@ -145,6 +143,13 @@ python-version = "3.10" [tool.ty.rules] unused-ignore-comment = "warn" +[[tool.ty.overrides]] +include = ["src/hayhooks/server/tracing.py"] + +[tool.ty.overrides.rules] +# Optional tracing imports resolve only when their extras are installed. +unused-ignore-comment = "ignore" + [tool.ty.src] exclude = ["tests/**/*"] diff --git a/src/hayhooks/server/pipelines/models.py b/src/hayhooks/server/pipelines/models.py index aec68137..14d4a812 100644 --- a/src/hayhooks/server/pipelines/models.py +++ b/src/hayhooks/server/pipelines/models.py @@ -7,54 +7,7 @@ from pydantic import BaseModel, Field, create_model from hayhooks.server.exceptions import PipelineWrapperError -from hayhooks.server.utils.yaml_utils import InputResolution, OutputResolution - - -def get_request_model_from_resolved_io( - pipeline_name: str, declared_inputs: dict[str, InputResolution] -) -> type[BaseModel]: - """ - Create a flat Pydantic request model from declared inputs resolved by yaml_utils. - - Args: - pipeline_name: Name of the pipeline used for model naming. - declared_inputs: Mapping of declared input name to InputResolution. - - Returns: - A Pydantic model with top-level fields matching declared input names. - """ - fields: dict[str, Any] = {} - - for input_name, resolution in declared_inputs.items(): - input_type = resolution.type - default_value = ... if resolution.required else None - fields[input_name] = (input_type, default_value) - - return create_model(f"{pipeline_name.capitalize()}RunRequest", **fields) - - -def get_response_model_from_resolved_io( - pipeline_name: str, declared_outputs: dict[str, OutputResolution] -) -> type[BaseModel]: - """ - Create a flat Pydantic response model from declared outputs resolved by yaml_utils. - - Args: - pipeline_name: Name of the pipeline used for model naming. - declared_outputs: Mapping of declared output name to OutputResolution. - - Returns: - A Pydantic model with top-level fields matching declared output names. - """ - fields: dict[str, Any] = {} - - for output_name, resolution in declared_outputs.items(): - output_type = resolution.type - fields[output_name] = (output_type, ...) - - return create_model( - f"{pipeline_name.capitalize()}RunResponse", result=(dict, Field(..., description="Pipeline result")) - ) +from hayhooks.server.utils.request_headers import accepts_request_headers def create_request_model_from_callable(func: Callable, model_name: str, docstring: Docstring) -> type[BaseModel]: @@ -69,11 +22,14 @@ def create_request_model_from_callable(func: Callable, model_name: str, docstrin Pydantic model class for request """ - params = inspect.signature(func).parameters + params = inspect.signature(func, eval_str=True).parameters + inject_headers = accepts_request_headers(func) param_docs = {p.arg_name: p.description for p in docstring.params} fields: dict[str, Any] = {} for name, param in params.items(): + if name == "headers" and inject_headers: + continue default_value = ... if param.default == param.empty else param.default description = param_docs.get(name) or f"Parameter '{name}'" field_info = Field(default=default_value, description=description) @@ -113,7 +69,7 @@ def create_response_model_from_callable( Pydantic model class for response, or None for streaming/file responses. """ - return_type = inspect.signature(func).return_annotation + return_type = inspect.signature(func, eval_str=True).return_annotation if return_type is inspect.Signature.empty: msg = f"Pipeline wrapper is missing a return type for '{func.__name__}' method" # ty: ignore[unresolved-attribute] @@ -154,7 +110,7 @@ def get_response_class_from_callable(func: Callable) -> type[Response] | None: * ``None`` for normal JSON endpoints (the caller should omit the ``response_class`` kwarg so FastAPI uses its default ``JSONResponse``). """ - return_type = inspect.signature(func).return_annotation + return_type = inspect.signature(func, eval_str=True).return_annotation if return_type is inspect.Signature.empty: return None diff --git a/src/hayhooks/server/pipelines/streaming.py b/src/hayhooks/server/pipelines/streaming.py index fdb6ccb5..42eb1476 100644 --- a/src/hayhooks/server/pipelines/streaming.py +++ b/src/hayhooks/server/pipelines/streaming.py @@ -34,6 +34,7 @@ _ASYNC_STREAMING_LOOP: contextvars.ContextVar[asyncio.AbstractEventLoop] = contextvars.ContextVar( "_hayhooks_async_streaming_loop" ) +_SHIELDED_PIPELINE_TASKS: set[asyncio.Task[Any]] = set() # Streaming callbacks are module-level so Haystack can serialize snapshot inputs. # The active queues live in ContextVars and are set around pipeline execution, so @@ -789,6 +790,25 @@ def _check_pipeline_task_exception(pipeline_task: asyncio.Task) -> None: raise exception +def _handle_shielded_pipeline_task_done(pipeline_task: asyncio.Task) -> None: + """Keep shielded tasks alive until completion and retrieve detached exceptions.""" + _SHIELDED_PIPELINE_TASKS.discard(pipeline_task) + with contextlib.suppress(asyncio.CancelledError): + if error := pipeline_task.exception(): + log.opt(exception=error).error("Error in shielded pipeline task") + + +def _detach_pipeline_task(pipeline_task: asyncio.Task) -> None: + """Keep an unfinished pipeline task alive after its stream closes.""" + if pipeline_task.done(): + with contextlib.suppress(asyncio.CancelledError): + pipeline_task.exception() + return + + _SHIELDED_PIPELINE_TASKS.add(pipeline_task) + pipeline_task.add_done_callback(_handle_shielded_pipeline_task_done) + + async def _stream_chunks_from_queue( queue: asyncio.Queue[StreamingChunk], pipeline_task: asyncio.Task, @@ -826,9 +846,6 @@ async def _stream_chunks_from_queue( if queue.empty(): _check_pipeline_task_exception(pipeline_task) continue - except asyncio.CancelledError: - log.warning("Async streaming generator was cancelled") - break except Exception as e: log.opt(exception=True).error("Unexpected error in async streaming generator: {}", e) raise @@ -837,22 +854,29 @@ async def _stream_chunks_from_queue( _check_pipeline_task_exception(pipeline_task) -async def _cleanup_pipeline_async(pipeline_task: asyncio.Task) -> None: +async def _cleanup_pipeline_async(pipeline_task: asyncio.Task, *, shield_pipeline_task: bool = False) -> None: """ Cleans up the pipeline task if it's still running. Args: pipeline_task: The task to clean up + shield_pipeline_task: Keep the task running after the stream closes instead of cancelling it """ - if not pipeline_task.done(): - pipeline_task.cancel() - try: - await asyncio.wait_for(pipeline_task, timeout=1.0) - except (asyncio.TimeoutError, asyncio.CancelledError): - pass - except Exception as e: - # Don't re-raise - this runs in finally block, so we don't want to mask original errors - log.opt(exception=True).warning("Error during pipeline task cleanup: {}", e) + if shield_pipeline_task: + _detach_pipeline_task(pipeline_task) + return + + if pipeline_task.done(): + return + + pipeline_task.cancel() + try: + await asyncio.wait_for(pipeline_task, timeout=1.0) + except (asyncio.TimeoutError, asyncio.CancelledError): + pass + except Exception as e: + # Don't re-raise - this runs in finally block, so we don't want to mask original errors + log.opt(exception=True).warning("Error during pipeline task cleanup: {}", e) def async_streaming_generator( # noqa: PLR0913, C901 @@ -867,6 +891,7 @@ def async_streaming_generator( # noqa: PLR0913, C901 include_outputs_from: set[str] | None = None, allow_sync_streaming_callbacks: bool = False, external_event_queue: asyncio.Queue[StreamingChunk | PipelineEvent | str | dict[str, Any]] | None = None, + shield_pipeline_task: bool = False, ) -> AsyncGenerator[StreamingChunk | PipelineEvent | str | dict[str, Any], None]: """ Creates an async generator that yields streaming chunks from a pipeline or agent execution. @@ -897,6 +922,8 @@ def async_streaming_generator( # noqa: PLR0913, C901 external_event_queue: Optional external asyncio queue to merge with internal events. Events from this queue will be yielded alongside streaming chunks from the pipeline. Supports StreamingChunk, PipelineEvent, str, or custom dict events. + shield_pipeline_task: Keep the pipeline task running if the stream closes or its consumer is cancelled. + Defaults to False, preserving the existing behavior of cancelling the task. Yields: StreamingChunk: Individual chunks from the streaming execution @@ -966,7 +993,7 @@ async def generator() -> AsyncGenerator[StreamingChunk | PipelineEvent | str | d yield result yield chunk - await pipeline_task + await (asyncio.shield(pipeline_task) if shield_pipeline_task else pipeline_task) final_chunk = _process_pipeline_end(pipeline_task.result(), on_pipeline_end) if final_chunk: yield final_chunk @@ -975,6 +1002,6 @@ async def generator() -> AsyncGenerator[StreamingChunk | PipelineEvent | str | d log.opt(exception=True).error("Unexpected error in async streaming generator: {}", e) raise finally: - await _cleanup_pipeline_async(pipeline_task) + await _cleanup_pipeline_async(pipeline_task, shield_pipeline_task=shield_pipeline_task) return generator() diff --git a/src/hayhooks/server/routers/openai.py b/src/hayhooks/server/routers/openai.py index 8577f3fd..fc4ca229 100644 --- a/src/hayhooks/server/routers/openai.py +++ b/src/hayhooks/server/routers/openai.py @@ -14,6 +14,7 @@ create_responses_router, ) from haystack.dataclasses import StreamingChunk +from starlette.datastructures import Headers from hayhooks.server.logger import log from hayhooks.server.pipelines.registry import registry @@ -27,6 +28,7 @@ trace_sync_stream, ) from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper +from hayhooks.server.utils.request_headers import accepts_request_headers @dataclass(frozen=True) @@ -105,14 +107,21 @@ def _select_execution_mode(wrapper: BasePipelineWrapper, dispatch: _OpenAIDispat async def _invoke_pipeline_method( - wrapper: BasePipelineWrapper, *, mode: str, method_name: str, model: str, call_kwargs: dict[str, Any] + wrapper: BasePipelineWrapper, + *, + mode: str, + method_name: str, + call_kwargs: dict[str, Any], + headers: dict[str, str] | None, ) -> Any: """Invoke the resolved pipeline method in either async or threadpool-sync mode.""" method = getattr(wrapper, method_name) - log.debug("Using {} ({}) for model: {}", method_name, mode, model) + if accepts_request_headers(method): + call_kwargs = {**call_kwargs, "headers": Headers(headers) if headers is not None else None} + log.debug("Using {} ({}) for model: {}", method_name, mode, call_kwargs["model"]) if mode == "async": - return await method(model=model, **call_kwargs) - return await run_in_threadpool(method, model=model, **call_kwargs) + return await method(**call_kwargs) + return await run_in_threadpool(method, **call_kwargs) def _wrap_string_as_streaming(text: str) -> Generator[StreamingChunk, None, None]: @@ -140,9 +149,11 @@ async def _run_pipeline_method( model: str, kwargs: dict[str, Any], body: dict[str, Any], + headers: dict[str, str] | None, ) -> str | Generator | AsyncGenerator: """Shared dispatch logic for chat completions and responses endpoints.""" stream_requested = bool(body.get("stream", False)) + call_kwargs = {"model": model, **kwargs, "body": body} trace_tags = build_trace_tags( { "hayhooks.transport": "openai", @@ -156,7 +167,7 @@ async def _run_pipeline_method( wrapper = _resolve_pipeline_wrapper(model) mode, method_name = _select_execution_mode(wrapper, dispatch) result = await _invoke_pipeline_method( - wrapper, mode=mode, method_name=method_name, model=model, call_kwargs={**kwargs, "body": body} + wrapper, mode=mode, method_name=method_name, call_kwargs=call_kwargs, headers=headers ) normalized_result = await _normalize_result(result, stream_requested=stream_requested) except BaseException: @@ -184,21 +195,25 @@ async def _run_pipeline_method( mode, method_name = _select_execution_mode(wrapper, dispatch) span.set_tag("hayhooks.openai.execution_mode", mode) result = await _invoke_pipeline_method( - wrapper, mode=mode, method_name=method_name, model=model, call_kwargs={**kwargs, "body": body} + wrapper, mode=mode, method_name=method_name, call_kwargs=call_kwargs, headers=headers ) return await _normalize_result(result, stream_requested=stream_requested) async def _run_completion( - model: str, messages: list[dict[str, Any]], body: dict[str, Any] + model: str, messages: list[dict[str, Any]], body: dict[str, Any], headers: dict[str, str] | None = None ) -> str | Generator | AsyncGenerator: - return await _run_pipeline_method(_CHAT_COMPLETION_DISPATCH, model=model, kwargs={"messages": messages}, body=body) + return await _run_pipeline_method( + _CHAT_COMPLETION_DISPATCH, model=model, kwargs={"messages": messages}, body=body, headers=headers + ) async def _run_response( - model: str, input_items: list[dict[str, Any]], body: dict[str, Any] + model: str, input_items: list[dict[str, Any]], body: dict[str, Any], headers: dict[str, str] | None = None ) -> str | Generator | AsyncGenerator: - return await _run_pipeline_method(_RESPONSE_DISPATCH, model=model, kwargs={"input_items": input_items}, body=body) + return await _run_pipeline_method( + _RESPONSE_DISPATCH, model=model, kwargs={"input_items": input_items}, body=body, headers=headers + ) def _find_file_upload_wrapper() -> BasePipelineWrapper | None: diff --git a/src/hayhooks/server/tracing.py b/src/hayhooks/server/tracing.py index 9a883ab5..e4ce4972 100644 --- a/src/hayhooks/server/tracing.py +++ b/src/hayhooks/server/tracing.py @@ -640,10 +640,11 @@ def finish(self, exc: BaseException | None = None) -> None: if self._finished: return + live_tags: dict[str, Any] = {} try: if exc is None: _mark_success(span) - live_tags: dict[str, Any] = {_TAG_SUCCESS: True} + live_tags = {_TAG_SUCCESS: True} elif isinstance(exc, HTTPException): _mark_http_exception(span, exc) live_tags = { diff --git a/src/hayhooks/server/utils/base_pipeline_wrapper.py b/src/hayhooks/server/utils/base_pipeline_wrapper.py index 2e0ae372..8134b362 100644 --- a/src/hayhooks/server/utils/base_pipeline_wrapper.py +++ b/src/hayhooks/server/utils/base_pipeline_wrapper.py @@ -4,6 +4,16 @@ class BasePipelineWrapper(ABC): + """ + Base class for deployed pipelines. + + The run_api, run_chat_completion and run_response methods, including their async + variants, can declare ``headers: Headers | None = None`` to receive HTTP request + headers. Import ``Headers`` from ``starlette.datastructures`` at runtime. Header + lookup is case-insensitive; A2A, MCP and direct calls use the default of None. + Ordinary dict parameters and **kwargs do not opt in to header injection. + """ + # Class attribute to skip MCP listing of the pipeline # If True, the pipeline will not be listed as an MCP tool # Even if it has a description and a request model diff --git a/src/hayhooks/server/utils/deploy_utils.py b/src/hayhooks/server/utils/deploy_utils.py index fc90918e..8115b918 100644 --- a/src/hayhooks/server/utils/deploy_utils.py +++ b/src/hayhooks/server/utils/deploy_utils.py @@ -12,11 +12,12 @@ from typing import Any, cast import docstring_parser -from fastapi import FastAPI, Form, HTTPException +from fastapi import FastAPI, Form, HTTPException, Request from fastapi.concurrency import run_in_threadpool from fastapi.responses import Response, StreamingResponse from fastapi.routing import APIRoute from pydantic import BaseModel +from starlette.datastructures import Headers from hayhooks.server.exceptions import PipelineAlreadyExistsError, PipelineFilesError from hayhooks.server.logger import log, log_elapsed @@ -46,6 +47,7 @@ load_pipeline_module, unload_pipeline_modules, ) +from hayhooks.server.utils.request_headers import accepts_request_headers from hayhooks.server.utils.streaming_response_utils import _streaming_response_from_result from hayhooks.server.utils.yaml_pipeline_wrapper import YAMLPipelineWrapper from hayhooks.settings import DeployConcurrencyPolicy, settings @@ -233,10 +235,20 @@ async def wrapper(*args, **kwargs): async def _execute_pipeline_run( pipeline_wrapper: BasePipelineWrapper, payload: dict[str, Any], + *, + headers: Headers | None = None, ) -> Any: + method = ( + pipeline_wrapper.run_api_async if pipeline_wrapper._is_run_api_async_implemented else pipeline_wrapper.run_api + ) + if accepts_request_headers(method): + if "headers" in payload: + msg = "Request headers cannot be supplied as pipeline arguments" + raise ValueError(msg) + payload = {**payload, "headers": headers} if pipeline_wrapper._is_run_api_async_implemented: - return await pipeline_wrapper.run_api_async(**payload) - return await run_in_threadpool(pipeline_wrapper.run_api, **payload) + return await method(**payload) + return await run_in_threadpool(method, **payload) _SENSITIVE_KEY_PATTERNS = { @@ -330,18 +342,19 @@ async def _execute_pipeline_run_with_tracing( pipeline_wrapper: BasePipelineWrapper, payload: dict[str, Any], *, + headers: Headers, trace_tags: dict[str, Any], is_streaming_response: bool, ) -> Any: if is_streaming_response: try: - return await _execute_pipeline_run(pipeline_wrapper, payload) + return await _execute_pipeline_run(pipeline_wrapper, payload, headers=headers) except BaseException: with trace_operation(SPAN_PIPELINE_RUN, tags=trace_tags): raise with trace_operation(SPAN_PIPELINE_RUN, tags=trace_tags): - return await _execute_pipeline_run(pipeline_wrapper, payload) + return await _execute_pipeline_run(pipeline_wrapper, payload, headers=headers) def _trace_streaming_run_result(result: Any, trace_tags: dict[str, Any]) -> Any: @@ -405,7 +418,7 @@ def create_run_endpoint_handler( ) is_streaming_response = get_response_class_from_callable(run_method) is StreamingResponse - async def _handle_request(run_req: BaseModel) -> Response | BaseModel: + async def _handle_request(run_req: BaseModel, request: Request) -> Response | BaseModel: payload = run_req.model_dump() trace_tags = _build_run_trace_tags(pipeline_name, payload) @@ -414,6 +427,7 @@ async def _handle_request(run_req: BaseModel) -> Response | BaseModel: result = await _execute_pipeline_run_with_tracing( pipeline_wrapper, payload, + headers=request.headers, trace_tags=trace_tags, is_streaming_response=is_streaming_response, ) @@ -449,13 +463,17 @@ async def _handle_request(run_req: BaseModel) -> Response | BaseModel: @handle_pipeline_exceptions() async def run_endpoint_with_files( + request: Request, run_req: request_model = Form(..., media_type="multipart/form-data"), # ty: ignore[invalid-type-form] # noqa: B008 ) -> response_model: # ty: ignore[invalid-type-form] - return await _handle_request(run_req) + return await _handle_request(run_req, request) @handle_pipeline_exceptions() - async def run_endpoint_without_files(run_req: request_model) -> response_model: # ty: ignore[invalid-type-form] - return await _handle_request(run_req) + async def run_endpoint_without_files( + request: Request, + run_req: request_model, # ty: ignore[invalid-type-form] + ) -> response_model: # ty: ignore[invalid-type-form] + return await _handle_request(run_req, request) return run_endpoint_with_files if requires_files else run_endpoint_without_files diff --git a/src/hayhooks/server/utils/mcp_utils.py b/src/hayhooks/server/utils/mcp_utils.py index 382e8441..6b505850 100644 --- a/src/hayhooks/server/utils/mcp_utils.py +++ b/src/hayhooks/server/utils/mcp_utils.py @@ -5,7 +5,6 @@ from enum import Enum from typing import Any -from fastapi.concurrency import run_in_threadpool from haystack.lazy_imports import LazyImport from starlette.applications import Starlette from starlette.responses import JSONResponse @@ -26,6 +25,7 @@ ) from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper from hayhooks.server.utils.deploy_utils import ( + _execute_pipeline_run, deploy_pipeline_files_async, deploy_pipelines, # noqa: F401 (re-exported; historically lived in this module) undeploy_pipeline_async, @@ -140,10 +140,7 @@ async def run_pipeline_as_tool(name: str, arguments: dict[str, Any]) -> list["Te msg = f"Pipeline '{name}' not found" raise ValueError(msg) - if pipeline_wrapper._is_run_api_async_implemented: - result = await pipeline_wrapper.run_api_async(**arguments) - else: - result = await run_in_threadpool(pipeline_wrapper.run_api, **arguments) + result = await _execute_pipeline_run(pipeline_wrapper, arguments) log.trace("Pipeline '{}' returned result: {}", name, result) diff --git a/src/hayhooks/server/utils/module_loader.py b/src/hayhooks/server/utils/module_loader.py index 9883ca18..140d3efc 100644 --- a/src/hayhooks/server/utils/module_loader.py +++ b/src/hayhooks/server/utils/module_loader.py @@ -15,6 +15,7 @@ from hayhooks.server.exceptions import PipelineModuleLoadError, PipelineWrapperError from hayhooks.server.logger import log from hayhooks.server.utils.base_pipeline_wrapper import BasePipelineWrapper +from hayhooks.server.utils.request_headers import accepts_request_headers from hayhooks.settings import settings @@ -241,7 +242,7 @@ def _raise_load_error(self, original_error: Exception) -> NoReturn: def _set_method_implementation_flags(pipeline_wrapper: BasePipelineWrapper) -> None: """ - Set implementation flags for all supported run methods. + Set implementation flags and validate request-header declarations on supported run methods. Args: pipeline_wrapper: The wrapper instance to annotate with flags. @@ -258,6 +259,8 @@ def _set_method_implementation_flags(pipeline_wrapper: BasePipelineWrapper) -> N for attr_name, method_name in methods_to_check: is_implemented = _is_method_overridden(pipeline_wrapper, method_name) + if is_implemented and method_name != "run_file_upload": + accepts_request_headers(getattr(pipeline_wrapper, method_name)) setattr(pipeline_wrapper, attr_name, is_implemented) log.debug("pipeline_wrapper.{}: {}", attr_name, is_implemented) diff --git a/src/hayhooks/server/utils/request_headers.py b/src/hayhooks/server/utils/request_headers.py new file mode 100644 index 00000000..2588b3a0 --- /dev/null +++ b/src/hayhooks/server/utils/request_headers.py @@ -0,0 +1,42 @@ +import inspect +from collections.abc import Callable +from types import SimpleNamespace, UnionType +from typing import Any, Union, get_args, get_origin, get_type_hints + +from starlette.datastructures import Headers + +from hayhooks.server.exceptions import PipelineWrapperError + + +def accepts_request_headers(method: Callable[..., Any]) -> bool: + """Recognize and validate the explicit ``headers: Headers | None = None`` opt-in.""" + try: + parameter = inspect.signature(method).parameters.get("headers") + except (TypeError, ValueError): + return False + if parameter is None or parameter.annotation is inspect.Parameter.empty: + return False + + # Resolve only this annotation; unrelated forward references need not be importable at runtime. + try: + annotation = get_type_hints( + SimpleNamespace(__annotations__={"headers": parameter.annotation}), + globalns=getattr(inspect.unwrap(method), "__globals__", {}), + )["headers"] + except NameError: + # An unresolved application type does not opt in to HTTP metadata. + return False + if annotation is not Headers and not ( + get_origin(annotation) in (Union, UnionType) and Headers in get_args(annotation) + ): + return False + + if ( + annotation != Headers | None + or parameter.default is not None + or parameter.kind not in (inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.KEYWORD_ONLY) + ): + name = getattr(method, "__name__", type(method).__name__) + msg = f"{name}: request headers must be declared as a keyword parameter 'headers: Headers | None = None'" + raise PipelineWrapperError(msg) + return True diff --git a/tests/test_deploy_utils.py b/tests/test_deploy_utils.py index 1869b946..942f58ce 100644 --- a/tests/test_deploy_utils.py +++ b/tests/test_deploy_utils.py @@ -448,37 +448,6 @@ def sample_func_no_doc() -> int: assert "result" in schema["required"] -@pytest.mark.parametrize( - "return_type", - [ - Response, - FileResponse, - StreamingResponse, - Generator, - AsyncGenerator, - Generator[str, None, None], - AsyncGenerator[str, None], - ], - ids=[ - "Response", - "FileResponse", - "StreamingResponse", - "Generator", - "AsyncGenerator", - "Generator[str, None, None]", - "AsyncGenerator[str, None]", - ], -) -def test_create_response_model_returns_none_for_non_json_types(return_type): - func = lambda: None # noqa: E731 - func.__annotations__["return"] = return_type - - docstring = docstring_parser.parse("") - result = create_response_model_from_callable(func, "Test", docstring) - - assert result is None - - @pytest.mark.parametrize( ("return_type", "expected_class"), [ @@ -489,24 +458,21 @@ def test_create_response_model_returns_none_for_non_json_types(return_type): (AsyncGenerator, StreamingResponse), (Generator[str, None, None], StreamingResponse), (AsyncGenerator[str, None], StreamingResponse), - ], - ids=[ - "Response", - "FileResponse", - "StreamingResponse", - "Generator", - "AsyncGenerator", - "Generator[str, None, None]", - "AsyncGenerator[str, None]", + ("Response", Response), + ("FileResponse", FileResponse), + ("StreamingResponse", StreamingResponse), + ("Generator", StreamingResponse), + ("AsyncGenerator", StreamingResponse), + ("Generator[str, None, None]", StreamingResponse), + ("AsyncGenerator[str, None]", StreamingResponse), ], ) -def test_get_response_class_for_non_json_types(return_type, expected_class): +def test_non_json_return_types_skip_response_models_and_select_response_class(return_type, expected_class): func = lambda: None # noqa: E731 func.__annotations__["return"] = return_type - result = get_response_class_from_callable(func) - - assert result is expected_class + assert create_response_model_from_callable(func, "Test", docstring_parser.parse("")) is None + assert get_response_class_from_callable(func) is expected_class def test_get_response_class_returns_none_for_json_types(): @@ -586,7 +552,8 @@ def setup(self): with pytest.raises( PipelineWrapperError, match=re.escape( - "At least one of run_api, run_api_async, run_chat_completion, run_chat_completion_async, run_response, or run_response_async must be implemented" + "At least one of run_api, run_api_async, run_chat_completion, run_chat_completion_async, " + "run_response, or run_response_async must be implemented" ), ): create_pipeline_wrapper_instance(module) diff --git a/tests/test_it_streaming_disconnect.py b/tests/test_it_streaming_disconnect.py new file mode 100644 index 00000000..b6ef8e94 --- /dev/null +++ b/tests/test_it_streaming_disconnect.py @@ -0,0 +1,103 @@ +import asyncio +import contextlib +from typing import Any + +import pytest +from fastapi import FastAPI +from haystack import component +from haystack.dataclasses import StreamingChunk + +from hayhooks.server.pipelines.streaming import _SHIELDED_PIPELINE_TASKS, async_streaming_generator +from hayhooks.server.utils.haystack_compat import AsyncPipeline +from hayhooks.server.utils.streaming_response_utils import _streaming_response_from_async_gen + + +@component +class _BlockingStreamingComponent: + def __init__(self) -> None: + self.release = asyncio.Event() + self.finished = asyncio.Event() + + @component.output_types(result=str) + def run(self, streaming_callback: Any | None = None) -> dict[str, str]: + msg = "The async pipeline must call run_async" + raise AssertionError(msg) + + @component.output_types(result=str) + async def run_async(self, streaming_callback: Any | None = None) -> dict[str, str]: + try: + await streaming_callback(StreamingChunk(content="first", index=0)) + await self.release.wait() + return {"result": "done"} + finally: + self.finished.set() + + +async def _disconnect_after_first_chunk(app: FastAPI) -> None: + incoming = asyncio.Queue() + first_chunk = asyncio.Event() + await incoming.put({"type": "http.request", "body": b"", "more_body": False}) + + async def receive(): + return await incoming.get() + + async def send(message): + if message["type"] == "http.response.body" and message.get("body"): + first_chunk.set() + + scope = { + "type": "http", + "asgi": {"version": "3.0", "spec_version": "2.3"}, + "http_version": "1.1", + "method": "GET", + "scheme": "http", + "path": "/stream", + "raw_path": b"/stream", + "query_string": b"", + "headers": [], + "client": ("test", 1), + "server": ("test", 80), + "root_path": "", + } + request_task = asyncio.create_task(app(scope, receive, send)) + + try: + await asyncio.wait_for(first_chunk.wait(), timeout=1.0) + await incoming.put({"type": "http.disconnect"}) + await asyncio.wait_for(request_task, timeout=1.0) + finally: + if not request_task.done(): + request_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await request_task + + +@pytest.mark.integration +@pytest.mark.parametrize("shield_pipeline_task", [False, True], ids=["cancel", "shield"]) +async def test_http_disconnect_pipeline_task(shield_pipeline_task): + component = _BlockingStreamingComponent() + pipeline = AsyncPipeline() + pipeline.add_component("blocking", component) + app = FastAPI() + + @app.get("/stream") + async def stream(): + generator = async_streaming_generator(pipeline, shield_pipeline_task=shield_pipeline_task) + return _streaming_response_from_async_gen(generator) + + tasks_before = set(_SHIELDED_PIPELINE_TASKS) + await _disconnect_after_first_chunk(app) + detached_tasks = _SHIELDED_PIPELINE_TASKS - tasks_before + + assert bool(detached_tasks) is shield_pipeline_task + assert all(not task.done() for task in detached_tasks) + if shield_pipeline_task: + assert not component.finished.is_set() + + component.release.set() + # Haystack may either complete the component or cancel it with the pipeline task. + await asyncio.wait_for(component.finished.wait(), timeout=1.0) + if detached_tasks: + await asyncio.wait_for(asyncio.gather(*detached_tasks), timeout=1.0) + await asyncio.sleep(0) + assert detached_tasks.isdisjoint(_SHIELDED_PIPELINE_TASKS) diff --git a/tests/test_request_headers.py b/tests/test_request_headers.py new file mode 100644 index 00000000..de3b34b5 --- /dev/null +++ b/tests/test_request_headers.py @@ -0,0 +1,298 @@ +import asyncio +import importlib.util +from types import SimpleNamespace + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from httpx import ASGITransport, AsyncClient + +from hayhooks.server.pipelines.registry import registry +from hayhooks.server.routers.deploy import router as deploy_router +from hayhooks.server.routers.openai import _run_completion, _run_response +from hayhooks.server.routers.openai import router as openai_router +from hayhooks.server.utils.a2a_utils import _run_chat_completion +from hayhooks.server.utils.mcp_utils import list_pipelines_as_tools, run_pipeline_as_tool +from hayhooks.settings import settings + +METHODS = ["run_api", "run_chat_completion", "run_response"] +METHODS += [method + "_async" for method in METHODS] +BASE_SOURCE = """\ +import asyncio +import time +from collections.abc import Generator, AsyncGenerator +from fastapi import UploadFile +from starlette.datastructures import Headers +from hayhooks import BasePipelineWrapper + +class PipelineWrapper(BasePipelineWrapper): + def setup(self): + pass +""" + + +@pytest.fixture +def headers_client(): + registry.clear() + app = FastAPI() + app.include_router(deploy_router) + app.include_router(openai_router) + with TestClient(app) as client: + yield client + registry.clear() + + +def deploy(client, source): + response = client.post( + "/deploy_files", + json={ + "name": "headers_test", + "files": {"pipeline_wrapper.py": source}, + "save_files": False, + }, + ) + assert response.status_code == 200, response.text + + +def wrapper_source(method, stream=False, files=False): + is_async = method.endswith("_async") + arguments = ( + "query: str" + if method.startswith("run_api") + else ( + "model: str, messages: list[dict], body: dict" + if "chat" in method + else "model: str, input_items: list[dict], body: dict" + ) + ) + if files: + arguments += ", files: list[UploadFile]" + prefix = "async " if is_async else "" + delay = "await asyncio.sleep(0.01)" if is_async else "time.sleep(0.01)" + return_type = ("AsyncGenerator" if is_async else "Generator") if stream else "str" + header_value = "headers.get('Authorization', 'missing') if headers is not None else 'no-context'" + result = ( + f" {prefix}def chunks():\n {delay}\n" + f" yield {header_value}\n return chunks()\n" + if stream + else f" return {header_value}\n" + ) + return BASE_SOURCE + ( + f" {prefix}def {method}(self, {arguments}, *, headers: Headers | None = None) -> {return_type}:\n" + f" {delay}\n" + result + ) + + +def request_for(method, stream=False, prefix="/v1"): + if method.startswith("run_api"): + return "/headers_test/run", {"query": "hi", "headers": {"authorization": "body-spoof"}} + body = {"model": "headers_test", "stream": stream, "headers": {"authorization": "body-spoof"}} + if "chat" in method: + return prefix + "/chat/completions", {**body, "messages": [{"role": "user", "content": "hi"}]} + return prefix + "/responses", {**body, "input": "hi"} + + +@pytest.mark.parametrize("method", METHODS) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("postponed", [False, True]) +async def test_headers_are_request_local_through_execution_and_streaming(headers_client, method, stream, postponed): + source = wrapper_source(method, stream) + deploy(headers_client, "from __future__ import annotations\n" + source if postponed else source) + if method.startswith("run_api"): + schema = headers_client.get("/openapi.json").json()["components"]["schemas"]["headers_testRunRequest"] + assert set(schema["properties"]) == {"query"} + url, body = request_for(method, stream) + tokens = ["Bearer alice-test", "Bearer bob-test", None] + async with AsyncClient(transport=ASGITransport(app=headers_client.app), base_url="http://test") as client: + responses = await asyncio.gather( + *[ + client.post(url, json=body, headers={"Authorization": token} if token is not None else {}) + for token in tokens + ] + ) + for token, response in zip(tokens, responses, strict=True): + assert response.status_code == 200, response.text + assert (token or "missing") in response.text + assert "body-spoof" not in response.text + assert all((other or "missing") not in response.text for other in tokens if other != token) + + +@pytest.mark.parametrize("method", [m for m in METHODS if not m.startswith("run_api")]) +def test_openai_aliases_forward_headers(headers_client, method): + deploy(headers_client, wrapper_source(method)) + url, body = request_for(method, prefix="") + response = headers_client.post(url, json=body, headers={"Authorization": "alias-token"}) + assert response.status_code == 200, response.text + assert "alias-token" in response.text + + +@pytest.mark.parametrize("method", ["run_api", "run_api_async"]) +@pytest.mark.parametrize("postponed", [False, True]) +def test_multipart_headers_are_injected_outside_the_form(headers_client, method, postponed): + source = wrapper_source(method, files=True) + deploy(headers_client, "from __future__ import annotations\n" + source if postponed else source) + response = headers_client.post( + "/headers_test/run", + data={"query": "hi", "headers": "body-spoof"}, + files={"files": ("test.txt", b"test", "text/plain")}, + headers={"Authorization": "upload-token"}, + ) + assert response.status_code == 200, response.text + assert response.json() == {"result": "upload-token"} + + +@pytest.mark.parametrize("method", ["run_api", "run_api_async"]) +@pytest.mark.skipif(importlib.util.find_spec("mcp") is None, reason="MCP is not installed") +@pytest.mark.mcp +async def test_injected_headers_stay_out_of_mcp_schema_and_arguments(headers_client, method): + deploy(headers_client, wrapper_source(method)) + tools = await list_pipelines_as_tools() + assert set(tools[0].inputSchema["properties"]) == {"query"} + result = await run_pipeline_as_tool("headers_test", {"query": "hi"}) + assert result[0].text == "no-context" + with pytest.raises(ValueError, match="cannot be supplied as pipeline arguments"): + await run_pipeline_as_tool("headers_test", {"query": "hi", "headers": {"authorization": "tool-spoof"}}) + + +@pytest.mark.parametrize("method", [m for m in METHODS if not m.startswith("run_api")]) +async def test_non_http_openai_calls_use_the_default(headers_client, method): + deploy(headers_client, wrapper_source(method)) + if "chat" in method: + assert await _run_completion("headers_test", [], {}) == "no-context" + else: + assert await _run_response("headers_test", [], {}) == "no-context" + + +@pytest.mark.parametrize("method", ["run_chat_completion", "run_chat_completion_async"]) +@pytest.mark.skipif(importlib.util.find_spec("a2a") is None, reason="A2A is not installed") +@pytest.mark.a2a +async def test_a2a_calls_use_the_default(headers_client, method): + deploy(headers_client, wrapper_source(method)) + assert await _run_chat_completion("headers_test", SimpleNamespace(message=None, current_task=None)) == "no-context" + + +@pytest.mark.parametrize( + "extra,assertion", + [ + ("", "True"), + (", **kwargs", "kwargs == {}"), + (", *headers", "headers == ()"), + (", headers: dict[str, str] | None = None", "headers is None"), + (", headers: 'NotImportedAtRuntime' = None", "headers is None"), + (", headers: dict[str, Headers] = None", "headers is None"), + ], +) +@pytest.mark.parametrize("method", [m for m in METHODS if not m.startswith("run_api")]) +def test_legacy_openai_signatures_do_not_receive_headers(headers_client, extra, assertion, method): + args = ( + "model: str, messages: list[dict], body: dict" + if "chat" in method + else "model: str, input_items: list[dict], body: dict" + ) + prefix = "async " if method.endswith("_async") else "" + deploy( + headers_client, + BASE_SOURCE + + ( + f" {prefix}def {method}(self, {args}{extra}) -> str:\n" + f" assert {assertion}\n return 'legacy'\n" + ), + ) + url, body = request_for(method) + response = headers_client.post(url, json=body, headers={"Authorization": "must-not-inject"}) + assert response.status_code == 200, response.text + assert "legacy" in response.text + + +@pytest.mark.parametrize("method", ["run_api", "run_api_async"]) +@pytest.mark.parametrize( + "transport", + [ + "http", + pytest.param( + "mcp", + marks=[ + pytest.mark.mcp, + pytest.mark.skipif(importlib.util.find_spec("mcp") is None, reason="MCP is not installed"), + ], + ), + ], +) +async def test_regular_headers_body_field_keeps_its_schema_and_value(headers_client, method, transport): + prefix = "async " if method.endswith("_async") else "" + deploy( + headers_client, + BASE_SOURCE + + ( + f" {prefix}def {method}(self, headers: dict[str, str]) -> str:\n" + " return headers['authorization']\n" + ), + ) + if transport == "http": + schema = headers_client.get("/openapi.json").json()["components"]["schemas"]["headers_testRunRequest"] + assert schema["required"] == ["headers"] + response = headers_client.post( + "/headers_test/run", + json={"headers": {"authorization": "body-value"}}, + headers={"Authorization": "transport-value"}, + ) + assert response.status_code == 200, response.text + assert response.json() == {"result": "body-value"} + else: + tools = await list_pipelines_as_tools() + assert tools[0].inputSchema["required"] == ["headers"] + result = await run_pipeline_as_tool("headers_test", {"headers": {"authorization": "tool-value"}}) + assert result[0].text == "tool-value" + + +@pytest.mark.parametrize( + "declaration", + [ + "headers: Headers", + "headers: Headers | None", + "headers: Headers = None", + "headers: Headers | str | None = None", + "headers: Headers | None = None, /", + "*headers: Headers | None", + "**headers: Headers | None", + ], +) +def test_invalid_opt_in_is_rejected_at_deployment(headers_client, declaration): + source = BASE_SOURCE + f" def run_api(self, {declaration}) -> str:\n return 'unused'\n" + response = headers_client.post( + "/deploy_files", + json={ + "name": "invalid_headers", + "files": {"pipeline_wrapper.py": source}, + "save_files": False, + }, + ) + assert response.status_code == 422, response.text + assert "headers: Headers | None = None" in response.json()["detail"] + assert registry.get("invalid_headers") is None + + +def test_postponed_header_annotation_is_resolved_without_resolving_unrelated_types(headers_client): + source = "from __future__ import annotations\n" + wrapper_source("run_chat_completion") + source = source.replace("messages: list[dict]", "messages: NotImportedAtRuntime") + deploy(headers_client, source) + response = headers_client.post( + "/chat/completions", + json={"model": "headers_test", "messages": []}, + headers={"Authorization": "resolved-token"}, + ) + assert response.status_code == 200, response.text + assert "resolved-token" in response.text + + +def test_injected_headers_are_not_payload_logs_or_trace_tags(headers_client, recording_tracer, caplog, monkeypatch): + monkeypatch.setattr(settings, "dashboard_trace_include_payload_values", True) + deploy(headers_client, wrapper_source("run_api")) + response = headers_client.post( + "/headers_test/run", + json={"query": "hi"}, + headers={"Authorization": "secret-transport-token"}, + ) + assert response.status_code == 200 + assert "secret-transport-token" not in caplog.text + assert "secret-transport-token" not in repr([span.tags for span in recording_tracer.spans]) diff --git a/tests/test_streaming.py b/tests/test_streaming.py index 0e419f3e..d91cdac1 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -1,3 +1,4 @@ +import asyncio import contextvars import os import time @@ -178,6 +179,28 @@ async def mock_run_async(messages=None, streaming_callback=None, **kwargs): return _factory +@pytest.fixture +def blocking_mock_agent(mocker): + release = asyncio.Event() + cancelled = asyncio.Event() + completed = asyncio.Event() + mock_agent = mocker.Mock(spec=Agent) + assert isinstance(mock_agent, Agent) + + async def mock_run_async(messages=None, streaming_callback=None, **kwargs): + await streaming_callback(StreamingChunk(content="First chunk", index=0)) + try: + await release.wait() + except asyncio.CancelledError: + cancelled.set() + raise + completed.set() + return {"messages": ["Done"]} + + mock_agent.run_async = mocker.AsyncMock(side_effect=mock_run_async) + return mock_agent, release, cancelled, completed + + def test_streaming_generator_with_sync_only_generator(pipeline_with_sync_only_generator): generator = streaming_generator(pipeline_with_sync_only_generator, pipeline_run_args={}) @@ -415,6 +438,62 @@ async def consume_with_cancel(): assert len(chunks) >= 1 +@pytest.mark.parametrize("shield_pipeline_task", [False, True], ids=["cancel", "shield"]) +async def test_async_streaming_generator_close_pipeline_task(blocking_mock_agent, shield_pipeline_task): + mock_agent, release, cancelled, completed = blocking_mock_agent + gen = async_streaming_generator( + mock_agent, + pipeline_run_args={"messages": []}, + shield_pipeline_task=shield_pipeline_task, + ) + + assert (await anext(gen)).content == "First chunk" + await gen.aclose() + + if shield_pipeline_task: + assert not cancelled.is_set() + release.set() + await asyncio.wait_for(completed.wait(), timeout=1.0) + else: + assert cancelled.is_set() + assert not completed.is_set() + + mock_agent.run_async.assert_awaited_once() + + +@pytest.mark.parametrize("shield_pipeline_task", [False, True], ids=["cancel", "shield"]) +async def test_async_streaming_consumer_cancellation_pipeline_task(blocking_mock_agent, shield_pipeline_task): + mock_agent, release, cancelled, completed = blocking_mock_agent + gen = async_streaming_generator( + mock_agent, + pipeline_run_args={"messages": []}, + shield_pipeline_task=shield_pipeline_task, + ) + first_chunk_received = asyncio.Event() + + async def consume(): + assert (await anext(gen)).content == "First chunk" + first_chunk_received.set() + await anext(gen) + + consumer_task = asyncio.create_task(consume()) + await first_chunk_received.wait() + consumer_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await consumer_task + + if shield_pipeline_task: + assert not cancelled.is_set() + release.set() + await asyncio.wait_for(completed.wait(), timeout=1.0) + else: + assert cancelled.is_set() + assert not completed.is_set() + + mock_agent.run_async.assert_awaited_once() + + # ContextVar used to verify context propagation into the streaming thread _test_context_var: contextvars.ContextVar[str] = contextvars.ContextVar("_test_context_var")