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..f192464bde 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,32 +1533,38 @@ 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(
- content=resp.content,
- status_code=resp.status_code,
- media_type=resp.headers.get("content-type", "application/json"),
+ 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()
@@ -1573,10 +1574,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 +1599,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 +1668,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 +1687,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"]