From 284b5035002560b96fadbb63632f875682f70d93 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 22 Jul 2026 21:06:34 +0000 Subject: [PATCH 1/2] feat: add SSE streaming for service-mode answer generation Introduce POST /v1/answer/stream to forward LLM tokens to clients as Server-Sent Events after VectorDB retrieval, with TTFT metrics and incremental think-tag filtering. Refactor shared answer helpers and add RetrieverServiceClient.aanswer_stream() for async consumption. Co-authored-by: Jeremy Dyer --- .../models/llm/clients/litellm.py | 146 ++++++++++- .../nemo_retriever/models/llm/text_utils.py | 52 ++++ .../src/nemo_retriever/service/client.py | 52 ++++ .../nemo_retriever/service/routers/ingest.py | 223 ++++++++++++++--- nemo_retriever/tests/test_llm_params.py | 23 ++ .../tests/test_service_answer_stream.py | 230 ++++++++++++++++++ 6 files changed, 687 insertions(+), 39 deletions(-) create mode 100644 nemo_retriever/tests/test_service_answer_stream.py diff --git a/nemo_retriever/src/nemo_retriever/models/llm/clients/litellm.py b/nemo_retriever/src/nemo_retriever/models/llm/clients/litellm.py index e498fe4648..e63199b251 100644 --- a/nemo_retriever/src/nemo_retriever/models/llm/clients/litellm.py +++ b/nemo_retriever/src/nemo_retriever/models/llm/clients/litellm.py @@ -14,6 +14,7 @@ import logging import time +from collections.abc import AsyncIterator from typing import Any, Optional from nemo_retriever.common.params.models import ( @@ -27,6 +28,7 @@ _build_rag_prompt as _task_build_rag_prompt, _deep_merge_dicts, ) +from nemo_retriever.models.llm.text_utils import ThinkTagStreamFilter from nemo_retriever.models.llm.types import ( GenerationResult, UnsupportedTextResponseError, @@ -129,16 +131,15 @@ def from_kwargs( ) return cls(transport=transport, sampling=sampling) - def complete( + def _build_call_kwargs( self, messages: list[dict], max_tokens: Optional[int] = None, extra_params: Optional[dict[str, Any]] = None, - ) -> tuple[str, float]: - """Raw litellm completion call. Returns (content_text, latency_s).""" + ) -> dict[str, Any]: + """Build provider-neutral kwargs shared by sync and streaming completion.""" validate_llm_extra_params(self.transport.extra_params, source="LLMRemoteClientParams.extra_params") validate_llm_extra_params(extra_params, source="GenerationRequest.extra_params") - import litellm sampling_kwargs = self.sampling.to_sampling_kwargs() if max_tokens is not None: @@ -160,6 +161,39 @@ def complete( if resolved_api_key: call_kwargs["api_key"] = resolved_api_key call_kwargs.update(_deep_merge_dicts(self.transport.extra_params, extra_params or {})) + return call_kwargs + + @staticmethod + def _delta_from_stream_chunk(chunk: object) -> str | None: + """Extract a text delta from one LiteLLM/OpenAI-compatible stream chunk.""" + choices = _field(chunk, "choices") + if not isinstance(choices, (list, tuple)) or not choices: + return None + choice = choices[0] + delta = _field(choice, "delta") + if delta is None: + message = _field(choice, "message") + if message is not None: + delta = message + if delta is None: + return None + content = _field(delta, "content") + if content is None: + return None + if not isinstance(content, str) or not content: + return None + return content + + def complete( + self, + messages: list[dict], + max_tokens: Optional[int] = None, + extra_params: Optional[dict[str, Any]] = None, + ) -> tuple[str, float]: + """Raw litellm completion call. Returns (content_text, latency_s).""" + import litellm + + call_kwargs = self._build_call_kwargs(messages, max_tokens, extra_params) t0 = time.monotonic() try: @@ -201,6 +235,110 @@ def complete( content = content.strip() return content, latency + async def stream_complete( + self, + messages: list[dict], + max_tokens: Optional[int] = None, + extra_params: Optional[dict[str, Any]] = None, + ) -> AsyncIterator[str]: + """Yield raw text deltas from ``litellm.acompletion(..., stream=True)``.""" + import litellm + + call_kwargs = self._build_call_kwargs(messages, max_tokens, extra_params) + call_kwargs["stream"] = True + + try: + response = await litellm.acompletion(**call_kwargs) + except Exception as exc: + err = str(exc) + if "temperature" in err and "top_p" in err: + logger.error( + "Model %s rejected the request because both `temperature` " + "and `top_p` were specified. Some providers (e.g. Bedrock) " + "only accept one. Either remove `top_p` from the model " + "config or set `temperature` to null. Sent: " + "temperature=%s, top_p=%s", + self.transport.model, + call_kwargs.get("temperature"), + call_kwargs.get("top_p"), + ) + raise + + async for chunk in response: + delta = self._delta_from_stream_chunk(chunk) + if delta is not None: + yield delta + + async def stream_generate( + self, + query: str, + chunks: list[str], + *, + reasoning_enabled: Optional[bool] = None, + ) -> AsyncIterator[tuple[str, dict[str, Any]]]: + """Stream visible RAG answer tokens and final generation metadata. + + Yields ``("token", {"delta": str, "index": int})`` events followed by + ``("complete", {"answer": str, "latency_s": float, "model": str, "error": str | None})``. + """ + effective_reasoning_enabled = ( + self.transport.reasoning_enabled if reasoning_enabled is None else reasoning_enabled + ) + task = RagAnswerTask( + system_prompt=self.transport.rag_system_prompt, + system_prompt_prefix=self.transport.rag_system_prompt_prefix, + reasoning_enabled=effective_reasoning_enabled, + ) + request = task.build_request( + query=query, + chunks=chunks, + reasoning_enabled=reasoning_enabled, + ) + think_filter = ThinkTagStreamFilter() if effective_reasoning_enabled is not False else None + + t0 = time.monotonic() + ttft_s: float | None = None + raw_parts: list[str] = [] + token_index = 0 + + try: + async for delta in self.stream_complete( + request.messages, + max_tokens=request.max_tokens, + extra_params=request.extra_params, + ): + raw_parts.append(delta) + visible_deltas = [delta] if think_filter is None else think_filter.feed(delta) + for visible in visible_deltas: + if ttft_s is None: + ttft_s = time.monotonic() - t0 + yield "metrics", {"ttft_s": ttft_s} + yield "token", {"delta": visible, "index": token_index} + token_index += 1 + except Exception as exc: + yield "complete", { + "answer": "", + "latency_s": time.monotonic() - t0, + "model": self.model, + "error": str(exc), + "ttft_s": ttft_s, + } + return + + raw_text = "".join(raw_parts) + parsed = task.parse(raw_text) + latency_s = time.monotonic() - t0 + error = None if parsed else task.empty_output_error + if ttft_s is not None: + yield "metrics", {"ttft_s": ttft_s, "generation_latency_s": latency_s} + yield "complete", { + "answer": parsed, + "latency_s": latency_s, + "model": self.model, + "error": error, + "ttft_s": ttft_s, + } + def generate( self, query: str, diff --git a/nemo_retriever/src/nemo_retriever/models/llm/text_utils.py b/nemo_retriever/src/nemo_retriever/models/llm/text_utils.py index e022a42677..c3ef78cf57 100644 --- a/nemo_retriever/src/nemo_retriever/models/llm/text_utils.py +++ b/nemo_retriever/src/nemo_retriever/models/llm/text_utils.py @@ -14,6 +14,58 @@ import re +_THINK_OPEN = "" +_THINK_CLOSE = "" + + +class ThinkTagStreamFilter: + """Incrementally suppress ```` blocks from a token stream.""" + + def __init__(self) -> None: + self._pending = "" + self._in_thinking = False + + def feed(self, chunk: str) -> list[str]: + """Return zero or more visible answer deltas from one streamed chunk.""" + if not chunk: + return [] + + self._pending += chunk + emitted: list[str] = [] + + while self._pending: + if self._in_thinking: + close_idx = self._pending.find(_THINK_CLOSE) + if close_idx == -1: + self._pending = "" + break + self._pending = self._pending[close_idx + len(_THINK_CLOSE) :] + self._in_thinking = False + continue + + open_idx = self._pending.find(_THINK_OPEN) + if open_idx == -1: + safe, self._pending = _split_safe_suffix(self._pending, _THINK_OPEN) + if safe: + emitted.append(safe) + break + + if open_idx > 0: + emitted.append(self._pending[:open_idx]) + self._pending = self._pending[open_idx + len(_THINK_OPEN) :] + self._in_thinking = True + + return emitted + + +def _split_safe_suffix(text: str, sentinel: str) -> tuple[str, str]: + """Split *text* into (safe_to_emit, suffix_that_may_prefix_sentinel).""" + max_prefix = min(len(text), len(sentinel) - 1) + for prefix_len in range(max_prefix, 0, -1): + if sentinel.startswith(text[-prefix_len:]): + return text[:-prefix_len], text[-prefix_len:] + return text, "" + def strip_think_tags(text: str) -> str: """Remove ``...`` reasoning blocks from model output. diff --git a/nemo_retriever/src/nemo_retriever/service/client.py b/nemo_retriever/src/nemo_retriever/service/client.py index 52bb763149..2df0059436 100644 --- a/nemo_retriever/src/nemo_retriever/service/client.py +++ b/nemo_retriever/src/nemo_retriever/service/client.py @@ -245,6 +245,58 @@ def query(self, query: str | list[str], *, top_k: int) -> list[list[dict[str, An except (ValidationError, ValueError) as exc: raise RuntimeError(f"Service query returned invalid response: {exc}") from exc + async def aanswer_stream( + self, + query: str, + *, + top_k: int = 5, + include_chunks: bool = False, + include_metadata: bool = False, + reasoning_enabled: bool | None = None, + reference: str | None = None, + judge: bool = False, + ) -> AsyncIterator[dict[str, Any]]: + """Stream ``POST /v1/answer/stream`` SSE events for one answer request.""" + url = f"{self._base_url}/v1/answer/stream" + body: dict[str, Any] = { + "query": query, + "top_k": top_k, + "include_chunks": include_chunks, + "include_metadata": include_metadata, + "judge": judge, + } + if reasoning_enabled is not None: + body["reasoning_enabled"] = reasoning_enabled + if reference is not None: + body["reference"] = reference + + async with httpx.AsyncClient( + timeout=httpx.Timeout(None, connect=30.0), + headers=self._auth_headers, + ) as client: + async with client.stream("POST", url, json=body) as response: + if response.status_code >= 400: + detail = (await response.aread()).decode(errors="replace")[:500] + raise RuntimeError(f"Service answer stream failed: HTTP {response.status_code}: {detail}") + + event_type = "" + data_buf = "" + async for line in response.aiter_lines(): + if line.startswith("event:"): + event_type = line[6:].strip() + elif line.startswith("data:"): + data_buf = line[5:].strip() + elif line == "" and data_buf: + try: + payload = json.loads(data_buf) + except json.JSONDecodeError: + data_buf = "" + event_type = "" + continue + yield {"event": event_type or "message", **payload} + data_buf = "" + event_type = "" + # ------------------------------------------------------------------ # Job lifecycle # ------------------------------------------------------------------ diff --git a/nemo_retriever/src/nemo_retriever/service/routers/ingest.py b/nemo_retriever/src/nemo_retriever/service/routers/ingest.py index 347414bc0b..0f6b832546 100644 --- a/nemo_retriever/src/nemo_retriever/service/routers/ingest.py +++ b/nemo_retriever/src/nemo_retriever/service/routers/ingest.py @@ -21,6 +21,7 @@ import ipaddress import json import logging +import time import uuid from datetime import datetime, timezone from typing import Any @@ -49,6 +50,7 @@ from nemo_retriever.models.llm.types import ( AnswerRequest as CoreAnswerRequest, AnswerResult, + GenerationResult, RetrievalResult, build_answer_result, ) @@ -1508,15 +1510,8 @@ def _metadata_from_hit(hit: dict[str, Any]) -> dict[str, Any]: return {k: v for k, v in hit.items() if k not in {"text", "content", "chunk", "page_content", "vector"}} -@router.post( - "/answer", - response_model=AnswerResult, - summary="Search ingested documents and generate an answer", -) -async def answer(req: ServiceAnswerRequest, request: Request) -> Response | AnswerResult: - """Retrieve context from VectorDB and answer with the configured LLM.""" - import httpx - +def _validate_answer_request(request: Request) -> None: + """Raise HTTP 404 when query/answer prerequisites are missing on this pod.""" config = request.app.state.config if not config.vectordb.enabled: @@ -1538,33 +1533,35 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ detail="Answer endpoint is not available on worker pods. Use the gateway.", ) - answer_req = CoreAnswerRequest( - query=req.query, - top_k=req.top_k, - reasoning_enabled=req.reasoning_enabled, - reference=req.reference, - judge_enabled=req.judge, - ) +async def _fetch_retrieval_for_answer( + request: Request, + *, + query: str, + top_k: int, +) -> tuple[RetrievalResult | None, Response | None, float]: + """Query VectorDB and return retrieval context plus retrieval latency.""" + config = request.app.state.config vectordb_url = config.vectordb.vectordb_url.rstrip("/") target = f"{vectordb_url}/v1/query" + started_at = time.monotonic() try: async with httpx.AsyncClient(timeout=60.0) as client: - resp = await client.post(target, json={"query": answer_req.query, "top_k": answer_req.top_k}) + resp = await client.post(target, json={"query": query, "top_k": top_k}) except Exception as exc: logger.exception("Failed to query vectordb at %s for answer generation", target) raise HTTPException( status_code=502, detail=f"Failed to reach VectorDB service: {type(exc).__name__}: {exc}", - ) + ) from exc if resp.status_code != 200: - return Response( + return None, Response( content=resp.content, status_code=resp.status_code, media_type=resp.headers.get("content-type", "application/json"), - ) + ), 0.0 payload = resp.json() result_sets = payload.get("results") or [] @@ -1573,10 +1570,14 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ chunks=[_text_from_hit(hit) for hit in hits], metadata=[_metadata_from_hit(hit) for hit in hits], ) + return retrieval, None, time.monotonic() - started_at + - from nemo_retriever.models.llm.clients import LLMJudge, LiteLLMClient +def _resolve_answer_llm(request: Request): + """Return the cached answer LLM client, creating it from config when needed.""" + from nemo_retriever.models.llm.clients import LiteLLMClient - llm_cfg = config.llm + llm_cfg = request.app.state.config.llm llm = getattr(request.app.state, "answer_llm_client", None) if llm is None: llm = LiteLLMClient.from_kwargs( @@ -1594,6 +1595,58 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ reasoning_enabled=llm_cfg.reasoning_enabled, ) request.app.state.answer_llm_client = llm + return llm + + +def _resolve_answer_judge(request: Request): + """Return the cached judge client when answer scoring is requested.""" + from nemo_retriever.models.llm.clients import LLMJudge + + llm_cfg = request.app.state.config.llm + judge = getattr(request.app.state, "answer_judge_client", None) + if judge is None: + judge = LLMJudge.from_kwargs( + model=llm_cfg.model, + api_base=llm_cfg.api_base, + api_key=llm_cfg.api_key, + extra_params=dict(llm_cfg.extra_params), + num_retries=llm_cfg.num_retries, + timeout=llm_cfg.timeout, + ) + request.app.state.answer_judge_client = judge + return judge + + +def _format_answer_sse_event(event_type: str, payload: dict[str, Any]) -> str: + return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n" + + +@router.post( + "/answer", + response_model=AnswerResult, + summary="Search ingested documents and generate an answer", +) +async def answer(req: ServiceAnswerRequest, request: Request) -> Response | AnswerResult: + """Retrieve context from VectorDB and answer with the configured LLM.""" + _validate_answer_request(request) + + answer_req = CoreAnswerRequest( + query=req.query, + top_k=req.top_k, + reasoning_enabled=req.reasoning_enabled, + reference=req.reference, + judge_enabled=req.judge, + ) + + retrieval, error_response, _retrieval_latency = await _fetch_retrieval_for_answer( + request, + query=answer_req.query, + top_k=answer_req.top_k, + ) + if error_response is not None: + return error_response + + llm = _resolve_answer_llm(request) generate_kwargs: dict[str, Any] = {} if answer_req.reasoning_enabled is not None: @@ -1611,19 +1664,7 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ detail=f"LLM answer generation failed: {gen.error}", ) - judge = None - if answer_req.judge_enabled: - judge = getattr(request.app.state, "answer_judge_client", None) - if judge is None: - judge = LLMJudge.from_kwargs( - model=llm_cfg.model, - api_base=llm_cfg.api_base, - api_key=llm_cfg.api_key, - extra_params=dict(llm_cfg.extra_params), - num_retries=llm_cfg.num_retries, - timeout=llm_cfg.timeout, - ) - request.app.state.answer_judge_client = judge + judge = _resolve_answer_judge(request) if answer_req.judge_enabled else None result = await asyncio.to_thread( build_answer_result, @@ -1642,6 +1683,118 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ ) +@router.post( + "/answer/stream", + summary="Search ingested documents and stream an answer over SSE", + response_model=None, +) +async def answer_stream(req: ServiceAnswerRequest, request: Request) -> Response | StreamingResponse: + """Stream retrieval and LLM answer tokens to the client as Server-Sent Events.""" + _validate_answer_request(request) + + answer_req = CoreAnswerRequest( + query=req.query, + top_k=req.top_k, + reasoning_enabled=req.reasoning_enabled, + reference=req.reference, + judge_enabled=req.judge, + ) + + retrieval, error_response, retrieval_latency_s = await _fetch_retrieval_for_answer( + request, + query=answer_req.query, + top_k=answer_req.top_k, + ) + if error_response is not None: + return error_response + + llm = _resolve_answer_llm(request) + generate_kwargs: dict[str, Any] = {} + if answer_req.reasoning_enabled is not None: + generate_kwargs["reasoning_enabled"] = answer_req.reasoning_enabled + + async def event_generator(): + retrieval_payload: dict[str, Any] = { + "query": answer_req.query, + "chunk_count": len(retrieval.chunks), + "retrieval_latency_s": retrieval_latency_s, + } + if req.include_chunks: + retrieval_payload["chunks"] = retrieval.chunks + if req.include_metadata: + retrieval_payload["metadata"] = retrieval.metadata + yield _format_answer_sse_event("retrieval_done", retrieval_payload) + + generation: GenerationResult | None = None + try: + async for event_type, payload in llm.stream_generate( + answer_req.query, + retrieval.chunks, + **generate_kwargs, + ): + if await request.is_disconnected(): + logger.info("SSE answer stream client disconnected (query=%r)", answer_req.query) + break + + if event_type == "metrics": + yield _format_answer_sse_event("metrics", payload) + continue + + if event_type == "token": + yield _format_answer_sse_event("token", payload) + continue + + if event_type == "complete": + if payload.get("error"): + logger.error( + "LLM answer streaming failed for model %s: %s", + payload.get("model"), + payload.get("error"), + ) + yield _format_answer_sse_event( + "error", + {"detail": f"LLM answer generation failed: {payload['error']}"}, + ) + return + + generation = GenerationResult( + answer=str(payload.get("answer") or ""), + latency_s=float(payload.get("latency_s") or 0.0), + model=str(payload.get("model") or llm.model), + error=payload.get("error"), + ) + except Exception as exc: + logger.exception("Unexpected failure while streaming answer for query=%r", answer_req.query) + yield _format_answer_sse_event( + "error", + {"detail": f"LLM answer generation failed: {type(exc).__name__}: {exc}"}, + ) + return + + if generation is None or await request.is_disconnected(): + return + + judge = _resolve_answer_judge(request) if answer_req.judge_enabled else None + result = await asyncio.to_thread( + build_answer_result, + query=answer_req.query, + retrieval=retrieval, + generation=generation, + reference=answer_req.reference, + judge=judge, + ) + done_payload = result.model_dump() + done_payload["chunks"] = result.chunks if req.include_chunks else None + done_payload["metadata"] = result.metadata if req.include_metadata else None + yield _format_answer_sse_event("done", done_payload) + + return StreamingResponse( + event_generator(), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, + ) + + # ------------------------------------------------------------------ # POST /v1/query — vector search (proxied to vectordb pod) # ------------------------------------------------------------------ diff --git a/nemo_retriever/tests/test_llm_params.py b/nemo_retriever/tests/test_llm_params.py index b3893e7586..14915dc1d9 100644 --- a/nemo_retriever/tests/test_llm_params.py +++ b/nemo_retriever/tests/test_llm_params.py @@ -255,6 +255,29 @@ def test_allowed_nested_extra_params_merge_recursively(self, mock_completion): assert kwargs["stop"] == ["END"] +class TestLiteLLMStreaming: + @patch("litellm.acompletion") + def test_stream_complete_forwards_stream_flag(self, mock_acompletion): + import asyncio + + from nemo_retriever.models.llm.clients import LiteLLMClient + + async def fake_stream(): + yield SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(content="Hi"))]) + yield SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(content=" there"))]) + + mock_acompletion.return_value = fake_stream() + client = LiteLLMClient.from_kwargs(model="openai/gpt-4o-mini") + + async def collect() -> list[str]: + return [delta async for delta in client.stream_complete([{"role": "user", "content": "hi"}])] + + deltas = asyncio.run(collect()) + + assert deltas == ["Hi", " there"] + assert mock_acompletion.call_args.kwargs["stream"] is True + + class TestLiteLLMHardening: """Credential, protected-field, and text-only response contracts.""" diff --git a/nemo_retriever/tests/test_service_answer_stream.py b/nemo_retriever/tests/test_service_answer_stream.py new file mode 100644 index 0000000000..caa5b40271 --- /dev/null +++ b/nemo_retriever/tests/test_service_answer_stream.py @@ -0,0 +1,230 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. +# All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Service-mode SSE streaming for POST /v1/answer/stream.""" + +from __future__ import annotations + +import json +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Any +from unittest.mock import patch + +import pytest +from fastapi.testclient import TestClient + +from nemo_retriever.models.llm.text_utils import ThinkTagStreamFilter, strip_think_tags +from nemo_retriever.service.app import create_app +from nemo_retriever.service.config import LLMConfig, LoggingConfig, PipelinePoolConfig, ServiceConfig, VectorDbConfig + + +def _parse_sse_response(raw: str) -> list[dict[str, Any]]: + events: list[dict[str, Any]] = [] + event_type = "" + data_buf = "" + for line in raw.splitlines(): + if line.startswith("event:"): + event_type = line[6:].strip() + elif line.startswith("data:"): + data_buf = line[5:].strip() + elif line == "" and data_buf: + payload = json.loads(data_buf) + events.append({"event": event_type, **payload}) + data_buf = "" + event_type = "" + return events + + +@pytest.fixture +def app_with_answer_config(monkeypatch: pytest.MonkeyPatch, tmp_path): + async def _stub_work(_item): + return 0, [] + + monkeypatch.setattr( + "nemo_retriever.service.services.pipeline_executor.create_realtime_work_fn", + lambda _config: _stub_work, + ) + monkeypatch.setattr( + "nemo_retriever.service.services.pipeline_executor.create_batch_work_fn", + lambda _config: _stub_work, + ) + + cfg = ServiceConfig( + mode="standalone", + logging=LoggingConfig(file=str(tmp_path / "service.log")), + pipeline=PipelinePoolConfig(realtime_workers=1, batch_workers=1), + vectordb=VectorDbConfig(enabled=True, vectordb_url="http://vectordb:7671"), + llm=LLMConfig( + enabled=True, + model="openai/nvidia/llama-3.3-nemotron-super-49b-v1.5", + api_base="http://llama-3-3-nemotron-super-49b-v1-5:8000/v1", + api_key="not-needed", + max_tokens=128, + timeout=180.0, + reasoning_enabled=False, + ), + ) + app = create_app(cfg) + with TestClient(app) as client: + yield client + + +def _install_fake_vectordb(monkeypatch: pytest.MonkeyPatch) -> None: + class _FakeResponse: + status_code = 200 + content = json.dumps( + { + "results": [ + { + "hits": [ + {"text": "Super-49B is the answer generator.", "source": "doc.pdf"}, + ] + } + ] + } + ).encode() + + def json(self) -> dict[str, Any]: + return json.loads(self.content.decode()) + + class _FakeAsyncClient: + def __init__(self, *args, **kwargs) -> None: + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + async def post(self, url: str, **kwargs) -> _FakeResponse: + return _FakeResponse() + + monkeypatch.setattr("httpx.AsyncClient", _FakeAsyncClient) + + +def test_think_tag_stream_filter_matches_batch_strip() -> None: + raw = "secretVisible answer" + filt = ThinkTagStreamFilter() + emitted: list[str] = [] + for chunk in ("secretVisible ", "answer"): + emitted.extend(filt.feed(chunk)) + assert "".join(emitted) == strip_think_tags(raw) + + +def test_answer_stream_emits_retrieval_tokens_and_done( + app_with_answer_config: TestClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fake_vectordb(monkeypatch) + + async def fake_stream_generate( + query: str, + chunks: list[str], + *, + reasoning_enabled: bool | None = None, + ) -> AsyncIterator[tuple[str, dict[str, Any]]]: + assert query == "What generates answers?" + assert chunks == ["Super-49B is the answer generator."] + yield "metrics", {"ttft_s": 0.05} + yield "token", {"delta": "Super-49B", "index": 0} + yield "token", {"delta": " does.", "index": 1} + yield "metrics", {"ttft_s": 0.05, "generation_latency_s": 0.2} + yield "complete", { + "answer": "Super-49B does.", + "latency_s": 0.2, + "model": "openai/nvidia/llama-3.3-nemotron-super-49b-v1.5", + "error": None, + "ttft_s": 0.05, + } + + fake_llm = SimpleNamespace(stream_generate=fake_stream_generate) + + with patch("nemo_retriever.models.llm.clients.LiteLLMClient.from_kwargs", return_value=fake_llm): + with app_with_answer_config.stream( + "POST", + "/v1/answer/stream", + json={"query": "What generates answers?", "include_chunks": True}, + ) as resp: + assert resp.status_code == 200 + assert resp.headers["content-type"].startswith("text/event-stream") + body = "".join(resp.iter_text()) + + events = _parse_sse_response(body) + assert [event["event"] for event in events] == [ + "retrieval_done", + "metrics", + "token", + "token", + "metrics", + "done", + ] + assert events[0]["chunk_count"] == 1 + assert events[0]["chunks"] == ["Super-49B is the answer generator."] + assert events[2]["delta"] == "Super-49B" + assert events[-1]["answer"] == "Super-49B does." + assert events[-1]["chunks"] == ["Super-49B is the answer generator."] + + +def test_answer_stream_emits_error_event_when_generation_fails( + app_with_answer_config: TestClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fake_vectordb(monkeypatch) + + async def fake_stream_generate( + query: str, + chunks: list[str], + *, + reasoning_enabled: bool | None = None, + ) -> AsyncIterator[tuple[str, dict[str, Any]]]: + yield "complete", { + "answer": "", + "latency_s": 0.0, + "model": "m", + "error": "connection refused", + "ttft_s": None, + } + + fake_llm = SimpleNamespace(stream_generate=fake_stream_generate) + + with patch("nemo_retriever.models.llm.clients.LiteLLMClient.from_kwargs", return_value=fake_llm): + with app_with_answer_config.stream("POST", "/v1/answer/stream", json={"query": "q"}) as resp: + body = "".join(resp.iter_text()) + + events = _parse_sse_response(body) + assert events[0]["event"] == "retrieval_done" + assert events[-1]["event"] == "error" + assert "connection refused" in events[-1]["detail"] + + +def test_answer_stream_returns_404_when_llm_disabled(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + async def _stub_work(_item): + return 0, [] + + monkeypatch.setattr( + "nemo_retriever.service.services.pipeline_executor.create_realtime_work_fn", + lambda _config: _stub_work, + ) + monkeypatch.setattr( + "nemo_retriever.service.services.pipeline_executor.create_batch_work_fn", + lambda _config: _stub_work, + ) + + app = create_app( + ServiceConfig( + mode="standalone", + logging=LoggingConfig(file=str(tmp_path / "service.log")), + pipeline=PipelinePoolConfig(realtime_workers=1, batch_workers=1), + vectordb=VectorDbConfig(enabled=True, vectordb_url="http://vectordb:7671"), + llm=LLMConfig(enabled=False), + ) + ) + + with TestClient(app) as client: + resp = client.post("/v1/answer/stream", json={"query": "q"}) + + assert resp.status_code == 404 + assert "LLM answer generation is not enabled" in resp.json()["detail"] From 008c485e4056d7460e19a0f5825e48c6911a1fa1 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 23 Jul 2026 02:23:12 +0000 Subject: [PATCH 2/2] chore: apply black formatting to ingest router Fixes pre-commit failure on the answer stream refactor return tuple. Co-authored-by: Jeremy Dyer --- .../src/nemo_retriever/service/routers/ingest.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/nemo_retriever/src/nemo_retriever/service/routers/ingest.py b/nemo_retriever/src/nemo_retriever/service/routers/ingest.py index 0f6b832546..f192464bde 100644 --- a/nemo_retriever/src/nemo_retriever/service/routers/ingest.py +++ b/nemo_retriever/src/nemo_retriever/service/routers/ingest.py @@ -1557,11 +1557,15 @@ async def _fetch_retrieval_for_answer( ) from exc if resp.status_code != 200: - return None, Response( - content=resp.content, - status_code=resp.status_code, - media_type=resp.headers.get("content-type", "application/json"), - ), 0.0 + return ( + None, + Response( + content=resp.content, + status_code=resp.status_code, + media_type=resp.headers.get("content-type", "application/json"), + ), + 0.0, + ) payload = resp.json() result_sets = payload.get("results") or []