diff --git a/nemo_retriever/helm/templates/configmap.yaml b/nemo_retriever/helm/templates/configmap.yaml index 3ed7a77ed6..afeb11fb0c 100644 --- a/nemo_retriever/helm/templates/configmap.yaml +++ b/nemo_retriever/helm/templates/configmap.yaml @@ -35,6 +35,20 @@ inherits the NIMService resource name, so the mapping is fixed: {{- $captionModelName = "nvidia/nemotron-3-nano-omni-30b-a3b-reasoning" -}} {{- end -}} {{- $audioGrpcEndpoint := $ctx.Values.serviceConfig.nimEndpoints.audioGrpcEndpoint | default "" -}} +{{- $rerankURL := include "nemo-retriever.nim.endpointURL" (dict "context" $ctx "key" "rerankqa" "serviceName" "llama-nemotron-rerank-vl-1b-v2" "configKey" "rerankInvokeUrl" "invokePath" "/v1/ranking") -}} +{{- /* + Model name resolution for the remote rerank endpoint: + 1. Explicit `serviceConfig.nimEndpoints.rerankModelName` always wins. + 2. Else, when a rerank URL was resolved (explicit or operator), fall back + to the canonical VL rerank model id (matches + `nemo_retriever.models.__init__.VL_RERANK_MODEL`). + 3. Else, empty string — leaves `rerank_model_name: null`. +*/}} +{{- $rerankModelName := $ctx.Values.serviceConfig.nimEndpoints.rerankModelName | default "" -}} +{{- if and (not $rerankModelName) $rerankURL -}} +{{- $rerankModelName = "nvidia/llama-nemotron-rerank-vl-1b-v2" -}} +{{- end -}} +{{- $rerankEnabled := or $ctx.Values.serviceConfig.rerank.enabled $rerankURL -}} {{- $llmAPIBase := $ctx.Values.serviceConfig.llm.apiBase | default "" -}} {{- $llmOperatorAPIBase := "" -}} {{- if not $llmAPIBase -}} @@ -76,6 +90,8 @@ nim_endpoints: embed_invoke_url: {{ .embedURL | quote }} embed_model_name: {{ .Values.serviceConfig.vectordb.embedModel | quote }} embed_model_provider_prefix: {{ if .Values.serviceConfig.vectordb.embedModelProviderPrefix }}{{ .Values.serviceConfig.vectordb.embedModelProviderPrefix | quote }}{{ else }}null{{ end }} + rerank_invoke_url: {{ if .rerankURL }}{{ .rerankURL | quote }}{{ else }}null{{ end }} + rerank_model_name: {{ if .rerankModelName }}{{ .rerankModelName | quote }}{{ else }}null{{ end }} caption_invoke_url: {{ if .captionURL }}{{ .captionURL | quote }}{{ else }}null{{ end }} caption_model_name: {{ if .captionModelName }}{{ .captionModelName | quote }}{{ else }}null{{ end }} audio_grpc_endpoint: {{ if .audioGrpcEndpoint }}{{ .audioGrpcEndpoint | quote }}{{ else }}null{{ end }} @@ -126,6 +142,11 @@ llm: rag_system_prompt_prefix: {{ if .llmRagSystemPromptPrefix }}{{ .llmRagSystemPromptPrefix | quote }}{{ else }}null{{ end }} reasoning_enabled: {{ .Values.serviceConfig.llm.reasoningEnabled }} +rerank: + enabled: {{ if .rerankEnabled }}true{{ else }}false{{ end }} + refine_factor: {{ .Values.serviceConfig.rerank.refineFactor }} + max_length: {{ .Values.serviceConfig.rerank.maxLength }} + pipeline: realtime_workers: {{ .Values.serviceConfig.pipeline.realtimeWorkers }} realtime_queue_size: {{ .Values.serviceConfig.pipeline.realtimeQueueSize }} @@ -184,7 +205,7 @@ metadata: data: retriever-service.yaml: | mode: standalone -{{ include "nemo-retriever.configBody" (dict "Values" .Values "gatewaySvc" $gatewaySvc "pageElementsURL" $pageElementsURL "tableStructureURL" $tableStructureURL "ocrURL" $ocrURL "embedURL" $embedURL "captionURL" $captionURL "captionModelName" $captionModelName "audioGrpcEndpoint" $audioGrpcEndpoint "llmEnabled" $llmEnabled "llmModel" $llmModel "llmAPIBase" $llmAPIBase "llmRagSystemPromptPrefix" $llmRagSystemPromptPrefix "vectordbSvc" $vectordbSvc "vectordbPort" $vectordbPort) | indent 4 }} +{{ include "nemo-retriever.configBody" (dict "Values" .Values "gatewaySvc" $gatewaySvc "pageElementsURL" $pageElementsURL "tableStructureURL" $tableStructureURL "ocrURL" $ocrURL "embedURL" $embedURL "captionURL" $captionURL "captionModelName" $captionModelName "audioGrpcEndpoint" $audioGrpcEndpoint "rerankURL" $rerankURL "rerankModelName" $rerankModelName "rerankEnabled" $rerankEnabled "llmEnabled" $llmEnabled "llmModel" $llmModel "llmAPIBase" $llmAPIBase "llmRagSystemPromptPrefix" $llmRagSystemPromptPrefix "vectordbSvc" $vectordbSvc "vectordbPort" $vectordbPort) | indent 4 }} {{- else }} # ========================================================================= # Split mode — one ConfigMap per role with the appropriate mode + gateway @@ -212,6 +233,6 @@ data: timeout_s: 300.0 max_connections: 100 {{- end }} -{{ include "nemo-retriever.configBody" (dict "Values" $.Values "gatewaySvc" $gatewaySvc "pageElementsURL" $pageElementsURL "tableStructureURL" $tableStructureURL "ocrURL" $ocrURL "embedURL" $embedURL "captionURL" $captionURL "captionModelName" $captionModelName "audioGrpcEndpoint" $audioGrpcEndpoint "llmEnabled" $llmEnabled "llmModel" $llmModel "llmAPIBase" $llmAPIBase "llmRagSystemPromptPrefix" $llmRagSystemPromptPrefix "vectordbSvc" $vectordbSvc "vectordbPort" $vectordbPort) | indent 4 }} +{{ include "nemo-retriever.configBody" (dict "Values" $.Values "gatewaySvc" $gatewaySvc "pageElementsURL" $pageElementsURL "tableStructureURL" $tableStructureURL "ocrURL" $ocrURL "embedURL" $embedURL "captionURL" $captionURL "captionModelName" $captionModelName "audioGrpcEndpoint" $audioGrpcEndpoint "rerankURL" $rerankURL "rerankModelName" $rerankModelName "rerankEnabled" $rerankEnabled "llmEnabled" $llmEnabled "llmModel" $llmModel "llmAPIBase" $llmAPIBase "llmRagSystemPromptPrefix" $llmRagSystemPromptPrefix "vectordbSvc" $vectordbSvc "vectordbPort" $vectordbPort) | indent 4 }} {{- end }} {{- end }} diff --git a/nemo_retriever/helm/values.yaml b/nemo_retriever/helm/values.yaml index 188ebd851e..3b51973cfb 100644 --- a/nemo_retriever/helm/values.yaml +++ b/nemo_retriever/helm/values.yaml @@ -545,6 +545,15 @@ serviceConfig: tableStructureInvokeUrl: "" ocrInvokeUrl: "" embedInvokeUrl: "" + # Optional remote reranking endpoint. Auto-wired from the in-cluster + # Service when `nimOperator.rerankqa.enabled=true` (and the NIM Operator + # CRDs are present). Set explicitly to point at a hosted endpoint. When a + # rerank URL is resolved, POST /v1/answer reranks retrieval hits (see the + # `rerank` section below). + rerankInvokeUrl: "" + # Model identifier passed to the remote reranking endpoint. Auto-set to + # the VL rerank model id when the operator-managed rerank NIM is enabled. + rerankModelName: "" # Optional remote VLM endpoint for image captioning (Nemotron 3 Nano # Omni). Auto-wired from the in-cluster Service when # `nimOperator.nemotron_3_nano_omni_30b_a3b_reasoning.enabled=true` @@ -634,6 +643,16 @@ serviceConfig: # portable no-reasoning controls. reasoningEnabled: true + # Reranking of retrieval hits on POST /v1/answer. The rerank endpoint / + # model live under nimEndpoints (rerankInvokeUrl / rerankModelName). This + # section only governs behaviour. When left disabled, reranking still turns + # on automatically if a rerank endpoint is resolved (operator NIM or an + # explicit rerankInvokeUrl). + rerank: + enabled: false + refineFactor: 4 + maxLength: 8192 + # Pipeline worker pools. Workers are abstract dispatchers — sizing # depends on whether they do local GPU work or fan out to remote NIMs. # For CPU-only NIM-forwarding nodes, higher worker counts are fine. diff --git a/nemo_retriever/src/nemo_retriever/common/policy.py b/nemo_retriever/src/nemo_retriever/common/policy.py index db5cf8051e..657e27aff3 100644 --- a/nemo_retriever/src/nemo_retriever/common/policy.py +++ b/nemo_retriever/src/nemo_retriever/common/policy.py @@ -349,6 +349,97 @@ def _scheme_of(uri: str) -> str: return uri.split("://", 1)[0].lower() + "://" +class EndpointOverridePolicy: + """Operator opt-in for per-request model-endpoint overrides. + + Mirrors ``pipeline_overrides.endpoint_overrides`` in + ``retriever-service.yaml``. Every stage flag defaults to ``False`` so + the secure baseline (endpoints are server-owned) is preserved unless the + operator explicitly opts in. ``allowed_url_prefixes`` optionally pins + client-supplied URLs to trusted destinations; an empty list imposes no + URL restriction once the stage flag is enabled. + """ + + def __init__( + self, + *, + embed: bool = False, + caption: bool = False, + llm: bool = False, + rerank: bool = False, + allowed_url_prefixes: list[str] | None = None, + ) -> None: + self.embed = embed + self.caption = caption + self.llm = llm + self.rerank = rerank + self.allowed_url_prefixes = list(allowed_url_prefixes or []) + + def _check_url_allowed(self, url: str | None, *, field: str) -> None: + if not url: + return + if not self.allowed_url_prefixes: + return + if not any(url.startswith(prefix) for prefix in self.allowed_url_prefixes): + raise PolicyError( + f"endpoint_overrides.{field} {url!r} does not match any allowed " + f"prefix in {self.allowed_url_prefixes!r}.", + status_code=403, + ) + + def check_embed(self, *, url: str | None) -> None: + if not self.embed: + raise PolicyError( + "endpoint_overrides: per-request embedding endpoint overrides are " + "disabled on this service. Ask the operator to set " + "pipeline_overrides.endpoint_overrides.embed: true in retriever-service.yaml.", + status_code=403, + ) + self._check_url_allowed(url, field="embed_invoke_url") + + def check_caption(self, *, url: str | None) -> None: + if not self.caption: + raise PolicyError( + "endpoint_overrides: per-request caption (VLM) endpoint overrides are " + "disabled on this service. Ask the operator to set " + "pipeline_overrides.endpoint_overrides.caption: true in retriever-service.yaml.", + status_code=403, + ) + self._check_url_allowed(url, field="caption_invoke_url") + + def check_llm(self, *, url: str | None) -> None: + if not self.llm: + raise PolicyError( + "endpoint_overrides: per-request LLM endpoint overrides are disabled " + "on this service. Ask the operator to set " + "pipeline_overrides.endpoint_overrides.llm: true in retriever-service.yaml.", + status_code=403, + ) + self._check_url_allowed(url, field="llm_api_base") + + def check_rerank(self, *, url: str | None) -> None: + if not self.rerank: + raise PolicyError( + "endpoint_overrides: per-request rerank endpoint overrides are disabled " + "on this service. Ask the operator to set " + "pipeline_overrides.endpoint_overrides.rerank: true in retriever-service.yaml.", + status_code=403, + ) + self._check_url_allowed(url, field="rerank_invoke_url") + + def describe(self) -> dict[str, Any]: + return { + "embed": self.embed, + "caption": self.caption, + "llm": self.llm, + "rerank": self.rerank, + "allowed_url_prefixes": self.allowed_url_prefixes, + } + + def any_enabled(self) -> bool: + return self.embed or self.caption or self.llm or self.rerank + + class PolicyError(ValueError): """Raised by :func:`validate_pipeline_spec` when a client spec is rejected. @@ -393,6 +484,7 @@ def __init__( extra_caption_keys: frozenset[str] = frozenset(), sinks: SinkUrlAllowlist | None = None, caption_enabled: bool = False, + endpoint_overrides: EndpointOverridePolicy | None = None, ) -> None: if mode not in {"reject", "allow_list", "allow_all"}: raise ValueError( @@ -401,6 +493,7 @@ def __init__( self.mode = mode self.sinks = sinks or SinkUrlAllowlist() self.caption_enabled = caption_enabled + self.endpoint_overrides = endpoint_overrides or EndpointOverridePolicy() # The base stage set grows incrementally: sinks open up # ``store``/``webhook``; a configured caption endpoint opens up # ``caption``. @@ -439,6 +532,7 @@ def describe(self) -> dict[str, Any]: "caption_enabled": self.caption_enabled, "denied_key_substrings": sorted(_DENYLIST_KEY_SUBSTRINGS), "sinks": self.sinks.describe(), + "endpoint_overrides": self.endpoint_overrides.describe(), } @@ -590,6 +684,72 @@ def _enforce_allowlist( ) +def _validate_endpoint_overrides( + overrides: Any, + policy: EndpointOverridePolicy, +) -> None: + """Gate per-request model-endpoint overrides against the operator policy. + + ``overrides`` is a + :class:`~nemo_retriever.common.schemas.pipeline_spec.EndpointOverrides` + or ``None``. Each stage group (embed / caption) is admitted only when the + operator enabled it; the URL is additionally checked against the + configured prefix allowlist. Raises :class:`PolicyError` (403) otherwise. + """ + if overrides is None or overrides.is_empty(): + return + + embed_requested = any( + ( + overrides.embed_invoke_url, + overrides.embed_model_name, + overrides.embed_model_provider_prefix, + ) + ) + caption_requested = any((overrides.caption_invoke_url, overrides.caption_model_name)) + + if embed_requested: + policy.check_embed(url=overrides.embed_invoke_url) + if caption_requested: + policy.check_caption(url=overrides.caption_invoke_url) + + # A bare credential with no endpoint/model is meaningless — reject so the + # client gets a clear error instead of a silently ignored credential. + if (overrides.api_key or overrides.embed_api_key or overrides.caption_api_key) and not ( + embed_requested or caption_requested + ): + raise PolicyError( + "endpoint_overrides embed/caption API key was set without any " + "endpoint or model override to apply it to.", + status_code=400, + ) + + +def _spec_has_only_endpoint_overrides(spec: PipelineSpec) -> bool: + """``True`` when ``endpoint_overrides`` is the sole client override. + + Used to admit an endpoint-only job independently of ``mode`` (the mode + gate governs the shape-params / stage-order surface, not the separately + opted-in endpoint channel). + """ + if spec.endpoint_overrides is None or spec.endpoint_overrides.is_empty(): + return False + return ( + spec.extract_params is None + and spec.embed_params is None + and spec.dedup_params is None + and spec.caption_params is None + and spec.store_params is None + and spec.vdb_upload_params is None + and spec.webhook_params is None + and spec.split_config is None + and spec.pdf_split is None + and not spec.stage_order + and not spec.return_embeddings + and not spec.return_images + ) + + def validate_pipeline_spec( spec: PipelineSpec | None, policy: PipelineOverridesPolicy, @@ -614,6 +774,7 @@ def validate_pipeline_spec( and spec.webhook_params is None and spec.split_config is None and spec.pdf_split is None + and (spec.endpoint_overrides is None or spec.endpoint_overrides.is_empty()) and not spec.stage_order and not spec.return_embeddings and not spec.return_images @@ -621,6 +782,15 @@ def validate_pipeline_spec( if result_schema_only: return spec + # Model-endpoint overrides are an independent, explicitly gated channel. + # They are validated regardless of ``mode`` (which governs the shape + # params / stage_order surface). When they are the *only* thing the + # client supplied, a submitted job with endpoint overrides is admissible + # even under ``mode='reject'`` because the operator opted in separately. + _validate_endpoint_overrides(spec.endpoint_overrides, policy.endpoint_overrides) + if _spec_has_only_endpoint_overrides(spec): + return spec + if policy.mode == "reject": raise PolicyError( "Per-request pipeline overrides are disabled on this service " @@ -629,23 +799,39 @@ def validate_pipeline_spec( status_code=403, ) + # A client that supplied its own caption (VLM) endpoint — and whose + # operator allows caption endpoint overrides — unlocks the caption stage + # even when the cluster itself has no caption NIM configured. + client_caption_endpoint = bool( + spec.endpoint_overrides is not None + and spec.endpoint_overrides.caption_invoke_url + and policy.endpoint_overrides.caption + ) + caption_enabled = policy.caption_enabled or client_caption_endpoint + allowed_stages = policy.allowed_stages + if client_caption_endpoint and "caption" not in allowed_stages: + allowed_stages = allowed_stages | {"caption"} + for stage_name in spec.stage_order: - if stage_name not in policy.allowed_stages: + if stage_name not in allowed_stages: raise PolicyError( f"stage {stage_name!r} is not in pipeline_overrides.allowed_stages. " - f"Allowed in this phase: {sorted(policy.allowed_stages)}.", + f"Allowed in this phase: {sorted(allowed_stages)}.", status_code=403, ) # caption_params is admitted only when the operator has configured a - # remote VLM endpoint. We also reject local-execution keys outright - # because the CPU worker pod cannot honor them — surfacing the - # mismatch immediately is friendlier than silently ignoring them. + # remote VLM endpoint (or the client supplied an allowed one). We also + # reject local-execution keys outright because the CPU worker pod cannot + # honor them — surfacing the mismatch immediately is friendlier than + # silently ignoring them. if spec.caption_params is not None: - if not policy.caption_enabled: + if not caption_enabled: raise PolicyError( "caption_params overrides require an operator-configured caption " - "endpoint. Set caption.endpoint_url in retriever-service.yaml first.", + "endpoint. Set caption.endpoint_url in retriever-service.yaml first, " + "or supply endpoint_overrides.caption_invoke_url (when the operator " + "enables pipeline_overrides.endpoint_overrides.caption).", status_code=403, ) forbidden = [k for k in spec.caption_params if k in _CAPTION_FORBIDDEN_LOCAL_EXECUTION_KEYS] diff --git a/nemo_retriever/src/nemo_retriever/common/schemas/pipeline_spec.py b/nemo_retriever/src/nemo_retriever/common/schemas/pipeline_spec.py index 4a1db96b83..46a555c126 100644 --- a/nemo_retriever/src/nemo_retriever/common/schemas/pipeline_spec.py +++ b/nemo_retriever/src/nemo_retriever/common/schemas/pipeline_spec.py @@ -52,6 +52,87 @@ class PdfSplitSpec(RichModel): pages_per_chunk: int = Field(default=32, ge=1, le=4096) +class EndpointOverrides(RichModel): + """Per-request model-endpoint overrides shipped from client → server. + + Endpoint URLs, model names, and API keys are normally *server-owned*: + they are baked into the worker pipeline from ``ServiceConfig.nim_endpoints`` + at startup and the :mod:`nemo_retriever.common.policy` denylist rejects + any attempt to smuggle them in through the ordinary ``*_params`` blocks. + + This model is the **explicit, audited channel** for a client to point a + submitted job at a *different* model deployment than the cluster default — + for example a purpose-built VLM for captioning or an alternative embedding + NIM. It is honored **only** when the operator has opted in via + ``pipeline_overrides.endpoint_overrides`` in ``retriever-service.yaml``; + otherwise the server rejects the request with HTTP 403. Because it is a + dedicated field rather than free-form params, the denylist that protects + ``embed_params`` / ``caption_params`` stays fully intact. + + Fields left ``None`` fall back to the server-configured default for that + stage. ``embed_*`` retarget the embedding NIM used by the ``embed`` stage; + ``caption_*`` retarget the VLM used by the ``caption`` stage. + ``embed_api_key`` / ``caption_api_key`` are optional credentials for + those stages; ``api_key`` remains as a legacy fallback when only one + stage is overridden (never replaces the server key for stages left at + their defaults). + """ + + model_config = ConfigDict(extra="forbid") + + embed_invoke_url: Optional[str] = Field( + default=None, + description="Remote embedding NIM URL to use for the embed stage instead of the cluster default.", + ) + embed_model_name: Optional[str] = Field( + default=None, + description="Model identifier passed to the overridden embedding endpoint.", + ) + embed_model_provider_prefix: Optional[str] = Field( + default=None, + description="Optional LiteLLM provider prefix prepended to embed_model_name.", + ) + caption_invoke_url: Optional[str] = Field( + default=None, + description="Remote VLM (caption) endpoint URL to use for the caption stage.", + ) + caption_model_name: Optional[str] = Field( + default=None, + description="Model identifier passed to the overridden caption (VLM) endpoint.", + ) + embed_api_key: Optional[str] = Field( + default=None, + description="Optional API key for the client-overridden embedding endpoint only.", + ) + caption_api_key: Optional[str] = Field( + default=None, + description="Optional API key for the client-overridden caption (VLM) endpoint only.", + ) + api_key: Optional[str] = Field( + default=None, + description=( + "Legacy fallback API key when a single overridden stage supplies " + "no stage-specific key. Prefer embed_api_key / caption_api_key " + "when both stages need credentials." + ), + ) + + def is_empty(self) -> bool: + """``True`` when the client did not set any override field.""" + return not any( + ( + self.embed_invoke_url, + self.embed_model_name, + self.embed_model_provider_prefix, + self.caption_invoke_url, + self.caption_model_name, + self.embed_api_key, + self.caption_api_key, + self.api_key, + ) + ) + + class PipelineSpec(RichModel): """Wire-format representation of fluent pipeline state. @@ -81,6 +162,11 @@ class PipelineSpec(RichModel): split_config: Optional[dict[str, Any]] = None pdf_split: Optional[PdfSplitSpec] = None + # Per-request model-endpoint overrides. Honored only when the operator + # opted in via ``pipeline_overrides.endpoint_overrides``; otherwise the + # policy layer rejects any non-empty value with HTTP 403. + endpoint_overrides: Optional[EndpointOverrides] = None + stage_order: list[StageName] = Field(default_factory=list) result_schema: Literal["legacy", "compact"] = Field( default="legacy", @@ -116,6 +202,7 @@ def is_empty(self) -> bool: and self.webhook_params is None and self.split_config is None and self.pdf_split is None + and (self.endpoint_overrides is None or self.endpoint_overrides.is_empty()) and not self.stage_order and self.result_schema == "legacy" and not self.return_embeddings diff --git a/nemo_retriever/src/nemo_retriever/service/config.py b/nemo_retriever/src/nemo_retriever/service/config.py index 4d729de66d..ae2b145646 100644 --- a/nemo_retriever/src/nemo_retriever/service/config.py +++ b/nemo_retriever/src/nemo_retriever/service/config.py @@ -146,7 +146,21 @@ class NimEndpointsConfig(RichModel): "remote embedding endpoints that require namespaced model IDs." ), ) - rerank_invoke_url: str | None = None + rerank_invoke_url: str | None = Field( + default=None, + description=( + "Remote reranking NIM endpoint used to re-order retrieval hits on " + "POST /v1/answer (subject to rerank.enabled). When set, clients may " + "override it per-request only via pipeline_overrides.endpoint_overrides.rerank." + ), + ) + rerank_model_name: str | None = Field( + default=None, + description=( + "Model identifier passed to the remote reranking endpoint. " + "Server-owned — clients cannot override the deployed rerank NIM SKU." + ), + ) audio_grpc_endpoint: str | None = Field( default=None, description=( @@ -200,6 +214,44 @@ def _validate_enabled_model(self) -> "LLMConfig": return self +class RerankConfig(RichModel): + """Reranking behaviour for the service-mode ``/v1/answer`` retrieval step. + + The rerank endpoint / model / API key live on + :class:`NimEndpointsConfig` (``rerank_invoke_url`` / ``rerank_model_name`` + / ``api_key``) so they share the server-owned trust boundary with the + other NIM endpoints. This section only governs *behaviour*: whether the + answer path reranks by default and how many candidates to over-fetch + before reranking. + """ + + model_config = ConfigDict(extra="forbid") + + enabled: bool = Field( + default=False, + description=( + "Rerank retrieval hits on POST /v1/answer before answer generation. " + "Requires a rerank endpoint (nim_endpoints.rerank_invoke_url or a " + "per-request override). Clients may toggle this per-request via the " + "'rerank' field." + ), + ) + refine_factor: int = Field( + default=4, + ge=1, + le=100, + description=( + "Over-fetch multiplier: retrieve top_k * refine_factor candidates " + "from the vector DB, rerank them, then keep the top_k best." + ), + ) + max_length: int = Field( + default=8192, + ge=1, + description="Tokenizer truncation length forwarded to the rerank endpoint.", + ) + + class ResourceLimitsConfig(RichModel): model_config = ConfigDict(extra="forbid") @@ -370,6 +422,49 @@ class SinksConfig(RichModel): ) +class EndpointOverridesConfig(RichModel): + """Opt-in policy for per-request model-endpoint overrides. + + By default every flag is ``False``, preserving the secure baseline + where endpoint URLs, model names, and API keys are exclusively + server-owned. When a flag is enabled, clients submitting a job may + ship a :class:`~nemo_retriever.common.schemas.pipeline_spec.EndpointOverrides` + that retargets that stage's model deployment. + + ``allowed_url_prefixes`` optionally restricts which URLs a client may + point at. An empty list means "no additional URL restriction" (any URL + is accepted once the corresponding stage flag is enabled). Populate it + (e.g. ``["https://"]`` or a specific host) to pin overrides to trusted + destinations and reduce SSRF exposure in shared clusters. + """ + + model_config = ConfigDict(extra="forbid") + + embed: bool = Field( + default=False, + description="Allow clients to override the embedding NIM endpoint / model per request.", + ) + caption: bool = Field( + default=False, + description="Allow clients to override the caption (VLM) endpoint / model per request.", + ) + llm: bool = Field( + default=False, + description="Allow clients to override the answer LLM api_base / model on POST /v1/answer.", + ) + rerank: bool = Field( + default=False, + description="Allow clients to override the rerank endpoint / model on POST /v1/answer.", + ) + allowed_url_prefixes: list[str] = Field( + default_factory=list, + description=( + "URL prefixes the client-supplied endpoints must start with. " + "Empty means any URL is allowed once the stage flag is enabled." + ), + ) + + class PipelineOverridesConfig(RichModel): """How permissively to accept per-request ``PipelineSpec`` overrides. @@ -397,6 +492,7 @@ class PipelineOverridesConfig(RichModel): extra_vdb_kwargs_keys: list[str] = Field(default_factory=list) extra_caption_keys: list[str] = Field(default_factory=list) sinks: SinksConfig = Field(default_factory=SinksConfig) + endpoint_overrides: EndpointOverridesConfig = Field(default_factory=EndpointOverridesConfig) def to_policy(self, *, caption_enabled: bool = False) -> "PipelineOverridesPolicy": # noqa: F821 """Return a :class:`PipelineOverridesPolicy` configured from this section. @@ -406,6 +502,7 @@ def to_policy(self, *, caption_enabled: bool = False) -> "PipelineOverridesPolic operator has actually wired up a VLM endpoint. """ from nemo_retriever.common.policy import ( + EndpointOverridePolicy, PipelineOverridesPolicy, SinkUrlAllowlist, ) @@ -428,6 +525,13 @@ def to_policy(self, *, caption_enabled: bool = False) -> "PipelineOverridesPolic vdb_uri_schemes=list(self.sinks.vdb_uri_schemes), ), caption_enabled=caption_enabled, + endpoint_overrides=EndpointOverridePolicy( + embed=self.endpoint_overrides.embed, + caption=self.endpoint_overrides.caption, + llm=self.endpoint_overrides.llm, + rerank=self.endpoint_overrides.rerank, + allowed_url_prefixes=list(self.endpoint_overrides.allowed_url_prefixes), + ), ) @@ -453,6 +557,7 @@ class ServiceConfig(RichModel): nim_endpoints: NimEndpointsConfig = Field(default_factory=NimEndpointsConfig) local_models: LocalModelsConfig = Field(default_factory=LocalModelsConfig) llm: LLMConfig = Field(default_factory=LLMConfig) + rerank: RerankConfig = Field(default_factory=RerankConfig) resources: ResourceLimitsConfig = Field(default_factory=ResourceLimitsConfig) auth: AuthConfig = Field(default_factory=AuthConfig) mcp: MCPConfig = Field(default_factory=MCPConfig) diff --git a/nemo_retriever/src/nemo_retriever/service/retriever-service.yaml b/nemo_retriever/src/nemo_retriever/service/retriever-service.yaml index c464c6aff6..d6b097dc1c 100644 --- a/nemo_retriever/src/nemo_retriever/service/retriever-service.yaml +++ b/nemo_retriever/src/nemo_retriever/service/retriever-service.yaml @@ -39,6 +39,12 @@ nim_endpoints: # Optional LiteLLM provider prefix prepended to embed_model_name for # proxies that require provider/model IDs. embed_model_provider_prefix: null + # Remote reranking NIM endpoint. When set (and rerank.enabled), POST + # /v1/answer re-orders retrieval hits before answer generation. The URL, + # model name, and API key stay server-owned unless a client override is + # explicitly enabled via pipeline_overrides.endpoint_overrides.rerank. + rerank_invoke_url: null + rerank_model_name: null # gRPC endpoint for the Parakeet ASR NIM (e.g. parakeet-nim:50051). # When set, audio/video pipelines use remote ASR instead of loading # the local Parakeet model (which requires torch + GPU). @@ -92,6 +98,14 @@ llm: rag_system_prompt_prefix: null reasoning_enabled: true +# Reranking of retrieval hits on POST /v1/answer. The rerank endpoint / +# model / API key live under nim_endpoints (rerank_invoke_url, +# rerank_model_name, api_key); this section only controls behaviour. +rerank: + enabled: false # rerank hits before answer generation (needs an endpoint) + refine_factor: 4 # over-fetch top_k * refine_factor candidates, then keep top_k + max_length: 8192 # tokenizer truncation length forwarded to the rerank endpoint + # Pipeline worker pools. Workers are abstract dispatchers — sizing # depends on whether they do local GPU work or fan out to remote NIMs. # For CPU-only NIM-forwarding nodes, higher worker counts are fine. @@ -202,3 +216,16 @@ pipeline_overrides: storage_uri_schemes: [] # e.g. ["s3://", "gs://", "azure://"] webhook_url_prefixes: [] # e.g. ["https://hooks.example.com/"] vdb_uri_schemes: [] # e.g. ["s3://", "gs://"] + # Per-request model-endpoint overrides. Endpoint URLs, model names, and API + # keys are server-owned by default; a client submitting a job can only point + # a stage at a different model deployment (a custom VLM, an alternative embed + # NIM, or a different answer LLM) when the operator opts in below. Each flag + # defaults to false — the secure baseline. allowed_url_prefixes optionally + # pins client-supplied URLs to trusted destinations (empty = no URL + # restriction once the corresponding flag is enabled). + endpoint_overrides: + embed: false # allow overriding embed_invoke_url / embed_model_name + caption: false # allow overriding the caption (VLM) endpoint / model + llm: false # allow overriding the /v1/answer LLM api_base / model + rerank: false # allow overriding the /v1/answer rerank endpoint / model + allowed_url_prefixes: [] # e.g. ["https://", "http://nim.internal."] diff --git a/nemo_retriever/src/nemo_retriever/service/routers/ingest.py b/nemo_retriever/src/nemo_retriever/service/routers/ingest.py index 347414bc0b..989bd79873 100644 --- a/nemo_retriever/src/nemo_retriever/service/routers/ingest.py +++ b/nemo_retriever/src/nemo_retriever/service/routers/ingest.py @@ -102,6 +102,43 @@ class ServiceAnswerRequest(BaseModel): reasoning_enabled: bool | None = None reference: str | None = None judge: bool = False + # Per-request LLM endpoint override. Honored only when the operator sets + # pipeline_overrides.endpoint_overrides.llm: true; otherwise HTTP 403. + llm_api_base: str | None = Field( + default=None, + description="Override the answer LLM api_base (requires operator opt-in).", + ) + llm_model: str | None = Field( + default=None, + description="Override the answer LLM model identifier (requires operator opt-in).", + ) + llm_api_key: str | None = Field( + default=None, + description="Override the answer LLM API key for the per-request endpoint.", + ) + # Per-request reranking. `rerank` toggles the retrieval-rerank step + # (None = use the server default rerank.enabled). The rerank_* fields + # retarget the rerank endpoint and are honored only when the operator sets + # pipeline_overrides.endpoint_overrides.rerank: true; otherwise HTTP 403. + rerank: bool | None = Field( + default=None, + description=( + "Rerank retrieval hits before answer generation. Defaults to the " + "server's rerank.enabled. Requires a rerank endpoint to be available." + ), + ) + rerank_invoke_url: str | None = Field( + default=None, + description="Override the rerank endpoint URL (requires operator opt-in).", + ) + rerank_model_name: str | None = Field( + default=None, + description="Override the rerank model identifier (requires operator opt-in).", + ) + rerank_api_key: str | None = Field( + default=None, + description="Override the rerank API key for the per-request endpoint.", + ) @model_validator(mode="after") def _validate_judge_reference(self) -> "ServiceAnswerRequest": @@ -109,6 +146,12 @@ def _validate_judge_reference(self) -> "ServiceAnswerRequest": raise ValueError("judge requires reference") return self + def has_llm_override(self) -> bool: + return bool(self.llm_api_base or self.llm_model or self.llm_api_key) + + def has_rerank_override(self) -> bool: + return bool(self.rerank_invoke_url or self.rerank_model_name or self.rerank_api_key) + logger = logging.getLogger(__name__) @@ -1508,6 +1551,38 @@ 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"}} +def _endpoint_override_policy(config: Any): + """Return the runtime :class:`EndpointOverridePolicy` for *config*.""" + return config.pipeline_overrides.to_policy().endpoint_overrides + + +def _resolve_rerank_settings(req: "ServiceAnswerRequest", config: Any) -> tuple[str | None, str | None, str | None]: + """Resolve the (url, model, api_key) used to rerank this request. + + Falls back to the server-configured rerank endpoint + (``nim_endpoints.rerank_invoke_url`` / ``rerank_model_name`` / ``api_key``). + A per-request override is honored only when the operator enabled + ``pipeline_overrides.endpoint_overrides.rerank`` and, when configured, the + URL matches ``allowed_url_prefixes``. Raises :class:`HTTPException` on a + policy denial so the caller surfaces a 403. + """ + nim = config.nim_endpoints + if not req.has_rerank_override(): + return nim.rerank_invoke_url, nim.rerank_model_name, nim.api_key + + try: + _endpoint_override_policy(config).check_rerank(url=req.rerank_invoke_url) + except PolicyError as exc: + raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc + rerank_url = req.rerank_invoke_url or nim.rerank_invoke_url + rerank_model = req.rerank_model_name or nim.rerank_model_name + if req.rerank_invoke_url: + rerank_api_key = req.rerank_api_key + else: + rerank_api_key = req.rerank_api_key or nim.api_key + return (rerank_url, rerank_model, rerank_api_key) + + @router.post( "/answer", response_model=AnswerResult, @@ -1546,12 +1621,35 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ judge_enabled=req.judge, ) + # Resolve reranking for this request: server default (rerank.enabled + + # nim_endpoints.rerank_*) plus an optional, operator-gated per-request + # override. When reranking, over-fetch top_k * refine_factor candidates so + # the reranker has a meaningful pool to re-order before we trim to top_k. + rerank_url, rerank_model, rerank_api_key = _resolve_rerank_settings(req, config) + # Honor an explicit rerank flag; otherwise rerank when the operator enabled + # it by default or the client supplied a per-request rerank endpoint. + want_rerank = req.rerank if req.rerank is not None else (config.rerank.enabled or req.has_rerank_override()) + if want_rerank and not rerank_url: + if req.rerank: + raise HTTPException( + status_code=400, + detail=( + "Reranking was requested but no rerank endpoint is available. " + "Configure nim_endpoints.rerank_invoke_url or pass rerank_invoke_url " + "with pipeline_overrides.endpoint_overrides.rerank enabled." + ), + ) + logger.warning("rerank.enabled is true but no rerank endpoint is configured; skipping rerank.") + do_rerank = bool(want_rerank and rerank_url) + # Over-fetch candidates for reranking, capped at the vectordb top_k ceiling. + fetch_k = min(answer_req.top_k * config.rerank.refine_factor, 1000) if do_rerank else answer_req.top_k + vectordb_url = config.vectordb.vectordb_url.rstrip("/") target = f"{vectordb_url}/v1/query" 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": answer_req.query, "top_k": fetch_k}) except Exception as exc: logger.exception("Failed to query vectordb at %s for answer generation", target) raise HTTPException( @@ -1569,6 +1667,29 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ payload = resp.json() result_sets = payload.get("results") or [] hits = result_sets[0].get("hits", []) if result_sets else [] + + if do_rerank and hits: + from nemo_retriever.operators.rerank import rerank_hits + + rerank_kwargs: dict[str, Any] = { + "rerank_invoke_url": rerank_url, + "api_key": rerank_api_key or "", + "max_length": config.rerank.max_length, + "top_n": answer_req.top_k, + } + if rerank_model: + rerank_kwargs["model_name"] = rerank_model + try: + hits = await asyncio.to_thread(rerank_hits, answer_req.query, hits, **rerank_kwargs) + except Exception as exc: + logger.exception("Reranking failed against %s", rerank_url) + raise HTTPException( + status_code=502, + detail=f"Reranking failed: {type(exc).__name__}: {exc}", + ) + else: + hits = hits[: answer_req.top_k] + retrieval = RetrievalResult( chunks=[_text_from_hit(hit) for hit in hits], metadata=[_metadata_from_hit(hit) for hit in hits], @@ -1577,12 +1698,20 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ from nemo_retriever.models.llm.clients import LLMJudge, LiteLLMClient llm_cfg = config.llm - llm = getattr(request.app.state, "answer_llm_client", None) - if llm is None: + + # Per-request LLM endpoint override — gated by the operator opt-in. When a + # client supplies a different model / api_base, build a one-off client + # instead of reusing (or poisoning) the cached server-default client. + if req.has_llm_override(): + try: + _endpoint_override_policy(config).check_llm(url=req.llm_api_base) + except PolicyError as exc: + raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc + llm_api_key = req.llm_api_key if req.llm_api_base else (req.llm_api_key or llm_cfg.api_key) llm = LiteLLMClient.from_kwargs( - model=llm_cfg.model, - api_base=llm_cfg.api_base, - api_key=llm_cfg.api_key, + model=req.llm_model or llm_cfg.model, + api_base=req.llm_api_base or llm_cfg.api_base, + api_key=llm_api_key, temperature=llm_cfg.temperature, top_p=llm_cfg.top_p, max_tokens=llm_cfg.max_tokens, @@ -1593,7 +1722,24 @@ async def answer(req: ServiceAnswerRequest, request: Request) -> Response | Answ rag_system_prompt_prefix=llm_cfg.rag_system_prompt_prefix, reasoning_enabled=llm_cfg.reasoning_enabled, ) - request.app.state.answer_llm_client = llm + else: + llm = getattr(request.app.state, "answer_llm_client", None) + if llm is None: + llm = LiteLLMClient.from_kwargs( + model=llm_cfg.model, + api_base=llm_cfg.api_base, + api_key=llm_cfg.api_key, + temperature=llm_cfg.temperature, + top_p=llm_cfg.top_p, + max_tokens=llm_cfg.max_tokens, + extra_params=dict(llm_cfg.extra_params), + num_retries=llm_cfg.num_retries, + timeout=llm_cfg.timeout, + rag_system_prompt=llm_cfg.rag_system_prompt, + rag_system_prompt_prefix=llm_cfg.rag_system_prompt_prefix, + reasoning_enabled=llm_cfg.reasoning_enabled, + ) + request.app.state.answer_llm_client = llm generate_kwargs: dict[str, Any] = {} if answer_req.reasoning_enabled is not None: diff --git a/nemo_retriever/src/nemo_retriever/service/service_ingestor.py b/nemo_retriever/src/nemo_retriever/service/service_ingestor.py index c962f245f6..a5677bd619 100644 --- a/nemo_retriever/src/nemo_retriever/service/service_ingestor.py +++ b/nemo_retriever/src/nemo_retriever/service/service_ingestor.py @@ -52,8 +52,29 @@ * ``.vdb_upload(...)`` — vector-DB sink, including sidecar metadata via the dedicated ``POST /v1/ingest/sidecar`` upload endpoint * ``.caption(...)`` — remote VLM captioning when the operator has wired - ``nim_endpoints.caption_invoke_url``; trust-sensitive fields like - endpoint_url / api_key / model_name stay server-owned + ``nim_endpoints.caption_invoke_url``. Behavioural knobs (prompt, + system_prompt, batch_size, …) are honored; passing ``endpoint_url`` / + ``model_name`` retargets the caption/VLM deployment for this job (see + endpoint overrides below). + +Per-request model-endpoint overrides +------------------------------------ +Endpoint URLs, model names, and API keys are normally server-owned. A +client can, however, point a stage at a *different* model deployment than +the cluster default by passing the relevant endpoint fields directly to +the stage method that owns that model: + +* ``.embed(embed_invoke_url=..., embed_model_name=...)`` — retarget the + embedding NIM. +* ``.caption(endpoint_url=..., model_name=...)`` — retarget the caption + (VLM) NIM (and unlock the caption stage even when the cluster has no + VLM configured). + +These are collected into a dedicated, audited ``PipelineSpec.endpoint_overrides`` +field on the wire so the ordinary shape-params denylist stays intact. They +are honored only when the operator opted in via +``pipeline_overrides.endpoint_overrides`` in ``retriever-service.yaml``; +otherwise the server responds with HTTP 403. * ``.save_to_disk(output_directory="...")`` — client-side persistence: fetches ``result_data`` from ``/v1/ingest/status/{id}`` for each completed document and writes JSON or gzipped JSON locally @@ -446,6 +467,22 @@ def _record_stage(self, name: str) -> None: if name not in order: order.append(name) + def _record_endpoint_overrides(self, **fields: Any) -> None: + """Merge non-``None`` model-endpoint override fields onto the spec. + + Stage methods (``.embed(...)`` / ``.caption(...)``) call this to route + the endpoint URL / model name they were handed into the dedicated, + audited ``endpoint_overrides`` channel — keeping those trust-sensitive + values out of the ordinary ``*_params`` blocks (which the server's + denylist rejects). + """ + overrides = {k: v for k, v in fields.items() if v is not None} + if not overrides: + return + existing = dict(self._pipeline_spec.get("endpoint_overrides") or {}) + existing.update(overrides) + self._pipeline_spec["endpoint_overrides"] = existing + def _fetch_document_result_data(self, document_id: str) -> list[dict[str, Any]]: """Fetch ``result_data`` for *document_id* from the status endpoint. @@ -539,6 +576,7 @@ def _pipeline_payload( "webhook_params", "split_config", "pdf_split", + "endpoint_overrides", ) ) and spec.get("result_schema", "legacy") == "legacy" @@ -617,20 +655,41 @@ def dedup(self, params: Any = None, **kwargs: Any) -> "ServiceIngestor": def embed(self, params: Any = None, **kwargs: Any) -> "ServiceIngestor": """Record an embed stage with optional :class:`EmbedParams` overrides. - Embedding endpoint URL and API key are server-owned and will be - rejected if set here. + Behavioural / shape knobs (batch size, input type, dimensions, …) are + bounded by the operator's allow-list. Passing an embedding endpoint + directly — ``embed_invoke_url`` (or ``embedding_endpoint``), + ``embed_model_name`` (or ``model_name``), + ``embed_model_provider_prefix`` — retargets the embedding NIM for this + job. Those fields are routed to the audited ``endpoint_overrides`` + channel and honored only when the operator enabled + ``pipeline_overrides.endpoint_overrides.embed``. """ if params is not None or kwargs: from nemo_retriever.common.policy import _DEFAULT_ALLOWED_EMBED_KEYS merged = _merge_params(params, kwargs) - _wire_client_stage_params( - self._pipeline_spec, - "embed_params", - merged, - method="embed", - allowed=_DEFAULT_ALLOWED_EMBED_KEYS, + params_dict = _params_to_dict(merged) + + # Peel off model-endpoint fields and route them to the endpoint + # override channel rather than the (denylisted) embed_params block. + embed_url = params_dict.pop("embed_invoke_url", None) or params_dict.pop("embedding_endpoint", None) + embed_model = params_dict.pop("embed_model_name", None) or params_dict.pop("model_name", None) + embed_prefix = params_dict.pop("embed_model_provider_prefix", None) + api_key = params_dict.pop("api_key", None) + has_embed_override = bool(embed_url or embed_model or embed_prefix) + self._record_endpoint_overrides( + embed_invoke_url=embed_url, + embed_model_name=embed_model, + embed_model_provider_prefix=embed_prefix, + embed_api_key=api_key if has_embed_override else None, ) + + # Remaining shape params still pass through the denylist + allowlist. + params_dict = _filter_policy_allowed( + _strip_server_owned(params_dict, "embed"), + _DEFAULT_ALLOWED_EMBED_KEYS, + ) + _set_stage_params(self._pipeline_spec, "embed_params", params_dict) self._record_stage("embed") return self @@ -953,23 +1012,28 @@ def save_to_disk( return self def caption(self, params: Any = None, **kwargs: Any) -> "ServiceIngestor": - """Record a caption stage backed by the server's remote VLM endpoint. + """Record a caption stage backed by a remote VLM endpoint. Behavioural knobs — ``prompt``, ``system_prompt``, ``batch_size``, ``context_text_max_chars``, ``caption_infographics``, and generic sampling params (``temperature``, ``max_tokens``, ``top_p``, - ``top_k``) — are honored. Trust-sensitive fields - (``endpoint_url``, ``api_key``, ``model_name``) and - local-execution fields (``device``, ``hf_cache_dir``, - ``tensor_parallel_size``, ``gpu_memory_utilization``) are - rejected on the client; the operator-configured remote endpoint - is the only path to a caption NIM. - - We use Pydantic's ``model_fields_set`` to distinguish fields - the caller *explicitly* set from fields carrying their - ``CaptionParams`` default — only the former are rejected. + ``top_k``) — are honored. + + Passing ``endpoint_url`` and/or ``model_name`` retargets the caption + (VLM) deployment for this job: those fields (plus an accompanying + ``api_key``) are routed to the audited ``endpoint_overrides`` channel + and honored only when the operator enabled + ``pipeline_overrides.endpoint_overrides.caption``. Supplying + ``endpoint_url`` also unlocks the caption stage even when the cluster + has no VLM configured. Without an override, the stage runs against the + operator-configured caption endpoint. + + Local-execution fields (``device``, ``hf_cache_dir``, + ``tensor_parallel_size``, ``gpu_memory_utilization``) have no effect + against a remote endpoint and are rejected. We use Pydantic's + ``model_fields_set`` to distinguish fields the caller *explicitly* + set from fields carrying their ``CaptionParams`` default. """ - trust_sensitive = {"endpoint_url", "api_key", "model_name"} local_only = { "device", "hf_cache_dir", @@ -977,47 +1041,40 @@ def caption(self, params: Any = None, **kwargs: Any) -> "ServiceIngestor": "gpu_memory_utilization", } - # Identify which keys the caller actually meant to pass. The - # signal for kwargs is unambiguous (any key in **kwargs is by - # definition caller-provided); for a passed-in CaptionParams - # instance we compare against class defaults, with one wrinkle: - # ``api_key`` is auto-populated by the model validator from the - # NVIDIA_API_KEY env var, so we cannot distinguish "caller set - # this" from "validator set this" — we conservatively strip the - # value either way and only raise when the caller used kwargs. explicit_keys: set[str] = set(kwargs.keys()) if isinstance(params, CaptionParams): class_defaults = {name: field.default for name, field in CaptionParams.model_fields.items()} - for k in trust_sensitive | local_only: - if k == "api_key": - continue # see comment above; the env-var auto-fill is ambiguous. + for k in local_only: val = getattr(params, k, None) if val is not None and val != class_defaults.get(k): explicit_keys.add(k) - bad_trust = sorted(explicit_keys & trust_sensitive) - if bad_trust: - raise ValueError( - f"ServiceIngestor.caption(): keys {bad_trust!r} are server-owned in " - "run_mode='service'. The operator configures the caption " - "endpoint via retriever-service.yaml (nim_endpoints.caption_invoke_url)." - ) bad_local = sorted(explicit_keys & local_only) if bad_local: raise ValueError( f"ServiceIngestor.caption(): keys {bad_local!r} configure local " "in-process GPU execution and have no effect against a remote " - "caption endpoint. Remove them and rely on the server-owned " - "model / endpoint." + "caption endpoint. Remove them and rely on the caption endpoint." ) merged = _merge_params(params, kwargs) if (params or kwargs) else CaptionParams() params_dict = _params_to_dict(merged) - # Drop both classes of keys before the spec leaves the client — - # the server's allowlist rejects them anyway, but failing fast - # at the boundary keeps the network message small and the policy - # error rare in practice. - scrubbed = {k: v for k, v in params_dict.items() if k not in trust_sensitive | local_only} + + # Peel off model-endpoint fields and route them to the endpoint + # override channel; behavioural knobs stay in caption_params. + caption_url = params_dict.pop("endpoint_url", None) + caption_model = params_dict.pop("model_name", None) + api_key = params_dict.pop("api_key", None) + has_caption_override = bool(caption_url or caption_model) + self._record_endpoint_overrides( + caption_invoke_url=caption_url, + caption_model_name=caption_model, + caption_api_key=api_key if has_caption_override else None, + ) + + # Drop local-execution keys before the spec leaves the client — the + # server's allowlist rejects them anyway. + scrubbed = {k: v for k, v in params_dict.items() if k not in local_only} self._pipeline_spec["caption_params"] = scrubbed self._record_stage("caption") return self diff --git a/nemo_retriever/src/nemo_retriever/service/services/pipeline_executor.py b/nemo_retriever/src/nemo_retriever/service/services/pipeline_executor.py index 3031de1d2a..a658303ce1 100644 --- a/nemo_retriever/src/nemo_retriever/service/services/pipeline_executor.py +++ b/nemo_retriever/src/nemo_retriever/service/services/pipeline_executor.py @@ -360,6 +360,71 @@ def _merge_server_owned( return merged +def _apply_embed_endpoint_override( + base_embed: dict[str, Any] | None, + endpoint_overrides: dict[str, Any], +) -> dict[str, Any] | None: + """Fold a validated per-request embed endpoint override onto the base dict. + + The override is applied to the *base* (server-owned) params so it survives + the belt-and-suspenders :func:`_merge_server_owned` pass — the policy layer + has already authorized it. Returns the (possibly new) base dict, or the + original when no embed override was supplied. + """ + embed_url = endpoint_overrides.get("embed_invoke_url") + embed_model = endpoint_overrides.get("embed_model_name") + embed_prefix = endpoint_overrides.get("embed_model_provider_prefix") + if not (embed_url or embed_model or embed_prefix): + return base_embed + + effective = dict(base_embed or {}) + if embed_url: + effective["embed_invoke_url"] = embed_url + # Keep the two URL fields consistent; the operator normalizes + # embed_invoke_url → embedding_endpoint downstream. + effective["embedding_endpoint"] = embed_url + if embed_model: + # Remote HTTP embedding reads model_name; local/GPU paths read embed_model_name. + effective["model_name"] = embed_model + effective["embed_model_name"] = embed_model + if embed_prefix: + effective["embed_model_provider_prefix"] = embed_prefix + api_key = endpoint_overrides.get("embed_api_key") or endpoint_overrides.get("api_key") + if api_key: + effective["api_key"] = api_key + elif embed_url: + effective.pop("api_key", None) + return effective + + +def _apply_caption_endpoint_override( + base_caption: dict[str, Any] | None, + endpoint_overrides: dict[str, Any], +) -> dict[str, Any] | None: + """Fold a validated per-request caption (VLM) endpoint override onto the base. + + When the cluster has no caption endpoint configured, a client-supplied URL + still produces a usable base dict so the caption stage can run against the + client's own VLM (the policy layer gates whether this is permitted). + """ + caption_url = endpoint_overrides.get("caption_invoke_url") + caption_model = endpoint_overrides.get("caption_model_name") + if not (caption_url or caption_model): + return base_caption + + effective = dict(base_caption or {}) + if caption_url: + effective["endpoint_url"] = caption_url + if caption_model: + effective["model_name"] = caption_model + api_key = endpoint_overrides.get("caption_api_key") or endpoint_overrides.get("api_key") + if api_key: + effective["api_key"] = api_key + elif caption_url: + effective.pop("api_key", None) + return effective + + def _resolve_sidecar_in_spec(spec: dict[str, Any] | None) -> dict[str, Any] | None: """Resolve ``vdb_upload_params.meta_dataframe_id`` to in-band bytes. @@ -527,25 +592,33 @@ def _build_graph_ingestor_from_spec( extraction_mode = _resolve_service_extraction_mode(spec.get("extraction_mode", "auto"), filename) split_config = spec.get("split_config") + # Per-request model-endpoint overrides (validated by the policy layer) + # are folded onto the server-owned base dicts so they survive the + # _merge_server_owned belt-and-suspenders pass below. + endpoint_overrides = spec.get("endpoint_overrides") or {} + effective_base_embed = _apply_embed_endpoint_override(base_embed, endpoint_overrides) + effective_base_caption = _apply_caption_endpoint_override(base_caption, endpoint_overrides) + extract_kwargs = _merge_server_owned(base_extract, spec.get("extract_params"), _TRUST_OWNED_EXTRACT_KEYS) extract_params = ExtractParams(**extract_kwargs) embed_override = spec.get("embed_params") - embed_params = _resolve_embed_params(base_embed, embed_override) + embed_params = _resolve_embed_params(effective_base_embed, embed_override) # Caption baseline + per-request overrides. The base dict carries - # the server-owned endpoint/API key/model name; the override carries - # behavioural knobs (prompt, system_prompt, batch_size, …). + # the server-owned (or client-overridden) endpoint/API key/model name; + # the override carries behavioural knobs (prompt, system_prompt, + # batch_size, …). caption_override = spec.get("caption_params") - if base_caption is None and caption_override is None: + if effective_base_caption is None and caption_override is None: caption_params = None - elif base_caption is None and caption_override is not None: + elif effective_base_caption is None and caption_override is not None: raise RuntimeError( "caption_params provided but no caption endpoint is configured on " "this worker. The policy layer should have rejected this earlier." ) else: - caption_kwargs = _merge_server_owned(base_caption or {}, caption_override, _TRUST_OWNED_CAPTION_KEYS) + caption_kwargs = _merge_server_owned(effective_base_caption or {}, caption_override, _TRUST_OWNED_CAPTION_KEYS) caption_params = CaptionParams(**caption_kwargs) if caption_kwargs.get("endpoint_url") else None asr_params = ASRParams(**base_asr) if base_asr else None diff --git a/nemo_retriever/tests/test_helm_rerank_endpoint.py b/nemo_retriever/tests/test_helm_rerank_endpoint.py new file mode 100644 index 0000000000..1d0bd30cde --- /dev/null +++ b/nemo_retriever/tests/test_helm_rerank_endpoint.py @@ -0,0 +1,191 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. +# All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for the rerank-endpoint auto-wiring in the Helm chart. + +The chart can deploy the VL reranker (``llama-nemotron-rerank-vl-1b-v2``) +as a NIMService, but until now the retriever-service ConfigMap rendered no +``rerank_invoke_url``, so ``POST /v1/answer`` never reranked retrieval hits +even when the reranker was Ready in the cluster. + +These tests pin the chart-side wiring: + +* ``serviceConfig.nimEndpoints`` exposes ``rerankInvokeUrl`` and + ``rerankModelName`` overrides, defaulting empty; ``serviceConfig.rerank`` + exposes the behaviour knobs (``enabled`` / ``refineFactor`` / ``maxLength``). +* ``templates/configmap.yaml`` resolves the rerank URL via the standard + ``nim.endpointURL`` helper (operator-managed + ``llama-nemotron-rerank-vl-1b-v2`` at ``/v1/ranking``) and renders both + the ``nim_endpoints`` fields and the ``rerank`` section. +* Explicit ``rerankInvokeUrl`` overrides win; the model name defaults to the + canonical VL rerank model id whenever any rerank URL is resolved; and + ``rerank.enabled`` flips true automatically when a URL is resolved. + +The integration tests shell out to ``helm template`` when ``helm`` is on +``$PATH``; otherwise they skip cleanly. +""" + +from __future__ import annotations + +import shutil +import subprocess +from pathlib import Path +from typing import Sequence +from unittest import SkipTest, TestCase, main + + +# Must match nemo_retriever.models.__init__.VL_RERANK_MODEL. +_VL_RERANK_MODEL_ID = "nvidia/llama-nemotron-rerank-vl-1b-v2" +_RERANK_OPERATOR_SERVICE = "llama-nemotron-rerank-vl-1b-v2" +_RERANK_INVOKE_PATH = "/v1/ranking" + + +def _repo_root() -> Path: + return Path(__file__).resolve().parents[2] + + +def _read_required_file(path: Path) -> str: + if not path.is_file(): + raise SkipTest(f"Required file not present in this test environment: {path}") + return path.read_text(encoding="utf-8") + + +def _helm_template( + extra_args: Sequence[str] = (), + api_versions: Sequence[str] = (), +) -> subprocess.CompletedProcess[str]: + helm = shutil.which("helm") + if helm is None: + raise SkipTest("`helm` binary not available in this environment.") + chart_path = _repo_root() / "nemo_retriever/helm" + if not chart_path.is_dir(): + raise SkipTest(f"Chart directory missing: {chart_path}") + + cmd: list[str] = [ + helm, + "template", + "retriever", + str(chart_path), + "--set", + "ngcImagePullSecret.create=false", + "--set", + "ngcApiSecret.create=false", + ] + for v in api_versions: + cmd += ["--api-versions", v] + cmd += list(extra_args) + return subprocess.run(cmd, check=False, capture_output=True, text=True) + + +def _assert_helm_ok(self: TestCase, proc: subprocess.CompletedProcess[str]) -> None: + self.assertEqual( + proc.returncode, + 0, + f"`helm template` failed unexpectedly:\nSTDOUT:\n{proc.stdout}\nSTDERR:\n{proc.stderr}", + ) + + +class HelmRerankEndpointTests(TestCase): + """Source-level + integration coverage of the rerank auto-wiring.""" + + # ------------------------------------------------------------------ + # Source / values + # ------------------------------------------------------------------ + + def test_values_expose_rerank_endpoint_overrides(self) -> None: + values = _read_required_file(_repo_root() / "nemo_retriever/helm/values.yaml") + self.assertIn("rerankInvokeUrl:", values) + self.assertIn("rerankModelName:", values) + self.assertIn('rerankInvokeUrl: ""', values) + self.assertIn('rerankModelName: ""', values) + # Behaviour knobs. + self.assertIn("refineFactor:", values) + self.assertIn("maxLength:", values) + + def test_configmap_resolves_rerank_url_via_standard_helper(self) -> None: + body = _read_required_file(_repo_root() / "nemo_retriever/helm/templates/configmap.yaml") + self.assertIn('"key" "rerankqa"', body) + self.assertIn(f'"serviceName" "{_RERANK_OPERATOR_SERVICE}"', body) + self.assertIn(f'"invokePath" "{_RERANK_INVOKE_PATH}"', body) + self.assertIn('"configKey" "rerankInvokeUrl"', body) + # nim_endpoints fields + the rerank behaviour section must render. + self.assertIn("rerank_invoke_url:", body) + self.assertIn("rerank_model_name:", body) + self.assertIn("rerank:", body) + self.assertIn("refine_factor:", body) + + # ------------------------------------------------------------------ + # Integration: actual `helm template` against the chart + # ------------------------------------------------------------------ + + def test_helm_template_autowires_rerank_when_operator_enabled(self) -> None: + proc = _helm_template( + extra_args=("--set", "nimOperator.rerankqa.enabled=true"), + api_versions=("apps.nvidia.com/v1alpha1",), + ) + _assert_helm_ok(self, proc) + expected_url = f'rerank_invoke_url: "http://{_RERANK_OPERATOR_SERVICE}:8000{_RERANK_INVOKE_PATH}"' + expected_model = f'rerank_model_name: "{_VL_RERANK_MODEL_ID}"' + self.assertIn(expected_url, proc.stdout) + self.assertIn(expected_model, proc.stdout) + # Resolving a URL auto-enables reranking on /v1/answer. + self.assertIn(" rerank:\n enabled: true", proc.stdout) + + def test_helm_template_rerank_null_when_operator_disabled(self) -> None: + proc = _helm_template( + extra_args=("--set", "nimOperator.rerankqa.enabled=false"), + api_versions=("apps.nvidia.com/v1alpha1",), + ) + _assert_helm_ok(self, proc) + self.assertIn("rerank_invoke_url: null", proc.stdout) + self.assertIn("rerank_model_name: null", proc.stdout) + self.assertIn(" rerank:\n enabled: false", proc.stdout) + + def test_helm_template_explicit_rerank_url_wins(self) -> None: + proc = _helm_template( + extra_args=( + "--set", + "nimOperator.rerankqa.enabled=true", + "--set", + "serviceConfig.nimEndpoints.rerankInvokeUrl=https://integrate.api.nvidia.com/v1/ranking", + "--set", + "serviceConfig.nimEndpoints.rerankModelName=nvidia/some-other-rerank", + ), + api_versions=("apps.nvidia.com/v1alpha1",), + ) + _assert_helm_ok(self, proc) + self.assertIn('rerank_invoke_url: "https://integrate.api.nvidia.com/v1/ranking"', proc.stdout) + self.assertIn('rerank_model_name: "nvidia/some-other-rerank"', proc.stdout) + + def test_helm_template_explicit_url_defaults_model_to_vl_rerank(self) -> None: + proc = _helm_template( + extra_args=( + "--set", + "nimOperator.rerankqa.enabled=false", + "--set", + "serviceConfig.nimEndpoints.rerankInvokeUrl=https://integrate.api.nvidia.com/v1/ranking", + ), + api_versions=("apps.nvidia.com/v1alpha1",), + ) + _assert_helm_ok(self, proc) + self.assertIn('rerank_invoke_url: "https://integrate.api.nvidia.com/v1/ranking"', proc.stdout) + self.assertIn(f'rerank_model_name: "{_VL_RERANK_MODEL_ID}"', proc.stdout) + + def test_helm_template_rerank_url_renders_in_split_mode(self) -> None: + proc = _helm_template( + extra_args=( + "--set", + "nimOperator.rerankqa.enabled=true", + "--set", + "topology.mode=split", + ), + api_versions=("apps.nvidia.com/v1alpha1",), + ) + _assert_helm_ok(self, proc) + url_count = proc.stdout.count(f"http://{_RERANK_OPERATOR_SERVICE}:8000{_RERANK_INVOKE_PATH}") + self.assertGreaterEqual(url_count, 3) + + +if __name__ == "__main__": + main() diff --git a/nemo_retriever/tests/test_service_caption.py b/nemo_retriever/tests/test_service_caption.py index e463474971..ba9ec792a8 100644 --- a/nemo_retriever/tests/test_service_caption.py +++ b/nemo_retriever/tests/test_service_caption.py @@ -8,8 +8,12 @@ client": the remote VLM endpoint URL + API key + model name are configured by the service operator (``nim_endpoints.caption_invoke_url`` in ``retriever-service.yaml``). Clients may submit behavioural knobs — -``prompt``, ``system_prompt``, ``batch_size``, sampling params — but -*never* redirect the destination. +``prompt``, ``system_prompt``, ``batch_size``, sampling params — freely. +A client can redirect the destination only through the dedicated, +operator-gated ``endpoint_overrides`` channel (see +``test_service_endpoint_overrides.py``); passing ``endpoint_url`` / +``model_name`` to ``.caption(...)`` routes into that channel rather than +the ordinary caption params. The CPU-only worker pod cannot honor local-execution params (``device``, ``hf_cache_dir``, ``tensor_parallel_size``, ``gpu_memory_utilization``) @@ -79,23 +83,34 @@ def test_caption_populates_spec_with_behaviour_knobs_only() -> None: assert "caption" in payload["stage_order"] -def test_caption_rejects_endpoint_url_via_kwargs() -> None: - """kwargs are unambiguous caller intent; the client rejects immediately.""" +def test_caption_routes_endpoint_url_via_kwargs_to_overrides() -> None: + """endpoint_url is now a per-request override, routed to endpoint_overrides. + + The server still gates it behind pipeline_overrides.endpoint_overrides.caption + (see test_service_endpoint_overrides.py); the client no longer rejects it. + """ ing = ServiceIngestor(base_url="http://example:7670") - with pytest.raises(ValueError, match="server-owned"): - ing.caption(endpoint_url="http://attacker.evil/v1") + ing.caption(endpoint_url="http://client/v1") + payload = ing._pipeline_payload() + assert payload["endpoint_overrides"]["caption_invoke_url"] == "http://client/v1" + assert "endpoint_url" not in payload.get("caption_params", {}) -def test_caption_rejects_api_key_via_kwargs() -> None: +def test_caption_bare_api_key_is_dropped_without_endpoint_override() -> None: + """A bare api_key (no endpoint/model override) stays server-owned: dropped.""" ing = ServiceIngestor(base_url="http://example:7670") - with pytest.raises(ValueError, match="server-owned"): - ing.caption(api_key="leaked-secret") + ing.caption(api_key="leaked-secret", prompt="describe") + payload = ing._pipeline_payload() + assert "endpoint_overrides" not in payload + assert "api_key" not in payload["caption_params"] -def test_caption_rejects_model_name_via_kwargs() -> None: +def test_caption_routes_model_name_via_kwargs_to_overrides() -> None: ing = ServiceIngestor(base_url="http://example:7670") - with pytest.raises(ValueError, match="server-owned"): - ing.caption(model_name="evil/model") + ing.caption(model_name="my/vlm") + payload = ing._pipeline_payload() + assert payload["endpoint_overrides"]["caption_model_name"] == "my/vlm" + assert "model_name" not in payload.get("caption_params", {}) def test_caption_rejects_local_execution_keys_via_kwargs() -> None: diff --git a/nemo_retriever/tests/test_service_endpoint_overrides.py b/nemo_retriever/tests/test_service_endpoint_overrides.py new file mode 100644 index 0000000000..b04119a4b9 --- /dev/null +++ b/nemo_retriever/tests/test_service_endpoint_overrides.py @@ -0,0 +1,767 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. +# All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Per-request model-endpoint overrides for service run_mode. + +Covers the opt-in ``pipeline_overrides.endpoint_overrides`` channel across +every layer: + +* schema — :class:`EndpointOverrides` emptiness + :class:`PipelineSpec` wiring; +* config — ``EndpointOverridesConfig`` → ``EndpointOverridePolicy`` plumbing; +* policy — accept when the operator opted in, reject otherwise, prefix + allowlist enforcement, client-supplied caption endpoints unlocking the + caption stage; +* worker — the base embed / caption dicts are retargeted and win the + server-owned merge; +* client — ``ServiceIngestor.embed(...)`` / ``.caption(...)`` route model + endpoints into the spec's ``endpoint_overrides`` channel; and +* answer — ``POST /v1/answer`` honors a per-request LLM override only when + enabled, and reranks retrieval hits (server default or per-request + rerank endpoint override) before answer generation. +""" + +from __future__ import annotations + +import json +from types import SimpleNamespace +from typing import Any +from unittest.mock import patch + +import pytest +from fastapi.testclient import TestClient + +from nemo_retriever.common.params import EmbedParams +from nemo_retriever.common.policy import PolicyError, validate_pipeline_spec +from nemo_retriever.common.schemas.pipeline_spec import EndpointOverrides, PipelineSpec +from nemo_retriever.models.llm.types import GenerationResult +from nemo_retriever.service.app import create_app +from nemo_retriever.service.config import ( + EndpointOverridesConfig, + LLMConfig, + LoggingConfig, + NimEndpointsConfig, + PipelineOverridesConfig, + PipelinePoolConfig, + RerankConfig, + ServiceConfig, + VectorDbConfig, +) +from nemo_retriever.service.service_ingestor import ServiceIngestor +from nemo_retriever.service.services.pipeline_executor import ( + _apply_caption_endpoint_override, + _apply_embed_endpoint_override, + _build_graph_ingestor_from_spec, +) + + +@pytest.fixture(autouse=True) +def _no_remote_api_keys(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("NVIDIA_API_KEY", raising=False) + monkeypatch.delenv("NGC_API_KEY", raising=False) + + +# ---------------------------------------------------------------------- +# Schema +# ---------------------------------------------------------------------- + + +def test_endpoint_overrides_is_empty() -> None: + assert EndpointOverrides().is_empty() + assert not EndpointOverrides(embed_invoke_url="http://x/embed").is_empty() + assert not EndpointOverrides(caption_model_name="my-vlm").is_empty() + + +def test_pipeline_spec_with_endpoint_overrides_is_not_empty() -> None: + spec = PipelineSpec(endpoint_overrides=EndpointOverrides(embed_invoke_url="http://x/embed")) + assert not spec.is_empty() + # An all-None overrides block does not make the spec non-empty. + assert PipelineSpec(endpoint_overrides=EndpointOverrides()).is_empty() + + +# ---------------------------------------------------------------------- +# Config → policy plumbing +# ---------------------------------------------------------------------- + + +def test_config_defaults_disable_all_endpoint_overrides() -> None: + cfg = EndpointOverridesConfig() + assert (cfg.embed, cfg.caption, cfg.llm, cfg.rerank) == (False, False, False, False) + assert cfg.allowed_url_prefixes == [] + + +def test_to_policy_carries_endpoint_override_flags() -> None: + cfg = PipelineOverridesConfig( + endpoint_overrides=EndpointOverridesConfig( + embed=True, caption=True, rerank=True, allowed_url_prefixes=["https://"] + ) + ) + policy = cfg.to_policy() + assert policy.endpoint_overrides.embed is True + assert policy.endpoint_overrides.caption is True + assert policy.endpoint_overrides.llm is False + assert policy.endpoint_overrides.rerank is True + assert policy.endpoint_overrides.allowed_url_prefixes == ["https://"] + described = policy.describe()["endpoint_overrides"] + assert described["embed"] is True + assert described["rerank"] is True + assert described["allowed_url_prefixes"] == ["https://"] + + +def test_check_rerank_policy_accept_and_reject() -> None: + from nemo_retriever.common.policy import EndpointOverridePolicy + + disabled = EndpointOverridePolicy() + with pytest.raises(PolicyError) as exc: + disabled.check_rerank(url="http://x/rerank") + assert exc.value.status_code == 403 + + enabled = EndpointOverridePolicy(rerank=True, allowed_url_prefixes=["https://"]) + enabled.check_rerank(url="https://ok/rerank") # no raise + with pytest.raises(PolicyError): + enabled.check_rerank(url="http://blocked/rerank") + assert enabled.any_enabled() is True + + +# ---------------------------------------------------------------------- +# Policy: accept / reject +# ---------------------------------------------------------------------- + + +def test_policy_rejects_endpoint_override_when_disabled() -> None: + policy = PipelineOverridesConfig().to_policy() + spec = PipelineSpec(endpoint_overrides=EndpointOverrides(embed_invoke_url="http://x/embed")) + with pytest.raises(PolicyError) as exc: + validate_pipeline_spec(spec, policy) + assert exc.value.status_code == 403 + assert "endpoint_overrides" in exc.value.detail + + +def test_policy_accepts_embed_override_when_enabled() -> None: + cfg = PipelineOverridesConfig(endpoint_overrides=EndpointOverridesConfig(embed=True)) + spec = PipelineSpec(endpoint_overrides=EndpointOverrides(embed_invoke_url="http://x/embed", embed_model_name="m")) + assert validate_pipeline_spec(spec, cfg.to_policy()) is spec + + +def test_policy_rejects_caption_override_when_only_embed_enabled() -> None: + cfg = PipelineOverridesConfig(endpoint_overrides=EndpointOverridesConfig(embed=True)) + spec = PipelineSpec(endpoint_overrides=EndpointOverrides(caption_invoke_url="http://x/vlm")) + with pytest.raises(PolicyError) as exc: + validate_pipeline_spec(spec, cfg.to_policy()) + assert exc.value.status_code == 403 + + +def test_policy_enforces_url_prefix_allowlist() -> None: + cfg = PipelineOverridesConfig( + endpoint_overrides=EndpointOverridesConfig(embed=True, allowed_url_prefixes=["https://trusted/"]) + ) + ok = PipelineSpec(endpoint_overrides=EndpointOverrides(embed_invoke_url="https://trusted/embed")) + assert validate_pipeline_spec(ok, cfg.to_policy()) is ok + + bad = PipelineSpec(endpoint_overrides=EndpointOverrides(embed_invoke_url="http://evil/embed")) + with pytest.raises(PolicyError) as exc: + validate_pipeline_spec(bad, cfg.to_policy()) + assert exc.value.status_code == 403 + + +def test_policy_rejects_bare_api_key_without_endpoint() -> None: + cfg = PipelineOverridesConfig(endpoint_overrides=EndpointOverridesConfig(embed=True)) + spec = PipelineSpec(endpoint_overrides=EndpointOverrides(api_key="leaked")) + with pytest.raises(PolicyError) as exc: + validate_pipeline_spec(spec, cfg.to_policy()) + assert exc.value.status_code == 400 + + +def test_endpoint_only_override_allowed_under_reject_mode() -> None: + """endpoint_overrides is an independent gate, so mode='reject' does not block it.""" + cfg = PipelineOverridesConfig(mode="reject", endpoint_overrides=EndpointOverridesConfig(embed=True)) + spec = PipelineSpec(endpoint_overrides=EndpointOverrides(embed_invoke_url="http://x/embed")) + assert validate_pipeline_spec(spec, cfg.to_policy()) is spec + + +@pytest.mark.parametrize("extraction_mode", ["image", "text", "html", "audio"]) +def test_endpoint_only_override_allowed_for_non_pdf_extraction_modes(extraction_mode: str) -> None: + cfg = PipelineOverridesConfig(mode="reject", endpoint_overrides=EndpointOverridesConfig(embed=True)) + spec = PipelineSpec( + extraction_mode=extraction_mode, + endpoint_overrides=EndpointOverrides(embed_invoke_url="http://x/embed"), + ) + assert validate_pipeline_spec(spec, cfg.to_policy()) is spec + + +def test_client_caption_endpoint_unlocks_caption_stage() -> None: + """A client-supplied VLM endpoint enables caption even without a server NIM.""" + cfg = PipelineOverridesConfig(endpoint_overrides=EndpointOverridesConfig(caption=True)) + spec = PipelineSpec( + endpoint_overrides=EndpointOverrides(caption_invoke_url="http://x/vlm"), + caption_params={"prompt": "Describe"}, + stage_order=["extract", "caption"], + ) + # caption_enabled=False on the server, but the client brought its own endpoint. + out = validate_pipeline_spec(spec, cfg.to_policy(caption_enabled=False)) + assert out is spec + + +def test_client_caption_params_without_endpoint_still_rejected_when_server_has_none() -> None: + cfg = PipelineOverridesConfig(endpoint_overrides=EndpointOverridesConfig(caption=True)) + spec = PipelineSpec(caption_params={"prompt": "Describe"}, stage_order=["extract", "caption"]) + with pytest.raises(PolicyError) as exc: + validate_pipeline_spec(spec, cfg.to_policy(caption_enabled=False)) + assert exc.value.status_code == 403 + + +# ---------------------------------------------------------------------- +# Worker merge +# ---------------------------------------------------------------------- + + +def test_apply_embed_endpoint_override_retargets_base() -> None: + base = {"embed_invoke_url": "http://server/embed", "model_name": "server-m", "api_key": "server-key"} + ov = {"embed_invoke_url": "http://client/embed", "embed_model_name": "client-m", "api_key": "client-key"} + out = _apply_embed_endpoint_override(base, ov) + assert out["embed_invoke_url"] == "http://client/embed" + assert out["model_name"] == "client-m" + assert out["embed_model_name"] == "client-m" + assert out["api_key"] == "client-key" + # The base dict is not mutated in place. + assert base["embed_invoke_url"] == "http://server/embed" + + +def test_apply_embed_endpoint_override_noop_without_embed_fields() -> None: + base = {"embed_invoke_url": "http://server/embed"} + assert _apply_embed_endpoint_override(base, {"caption_invoke_url": "http://x/vlm"}) is base + + +def test_apply_caption_endpoint_override_creates_base_when_none() -> None: + ov = {"caption_invoke_url": "http://client/vlm", "caption_model_name": "vlm-x", "api_key": "k"} + out = _apply_caption_endpoint_override(None, ov) + assert out == {"endpoint_url": "http://client/vlm", "model_name": "vlm-x", "api_key": "k"} + + +def test_apply_embed_endpoint_override_clears_server_key_when_url_overridden_without_client_key() -> None: + base = {"embed_invoke_url": "http://server/embed", "model_name": "server-m", "api_key": "server-key"} + out = _apply_embed_endpoint_override(base, {"embed_invoke_url": "http://client/embed"}) + assert out["embed_invoke_url"] == "http://client/embed" + assert "api_key" not in out + + +def test_apply_embed_endpoint_override_keeps_server_key_for_model_only_override() -> None: + base = {"embed_invoke_url": "http://server/embed", "model_name": "server-m", "api_key": "server-key"} + out = _apply_embed_endpoint_override(base, {"embed_model_name": "client-m"}) + assert out["model_name"] == "client-m" + assert out["api_key"] == "server-key" + + +def test_apply_caption_endpoint_override_clears_server_key_when_url_overridden_without_client_key() -> None: + base = {"endpoint_url": "http://server/vlm", "model_name": "server-m", "api_key": "server-key"} + out = _apply_caption_endpoint_override(base, {"caption_invoke_url": "http://client/vlm"}) + assert out["endpoint_url"] == "http://client/vlm" + assert "api_key" not in out + + +def test_build_graph_ingestor_applies_embed_endpoint_override() -> None: + base_embed = {"embed_invoke_url": "http://server/embed", "model_name": "server-m", "api_key": "server-key"} + spec = { + "extraction_mode": "auto", + "stage_order": ["extract", "embed"], + "endpoint_overrides": {"embed_invoke_url": "http://client/embed", "embed_model_name": "client-m"}, + } + ingestor, _mode, _ = _build_graph_ingestor_from_spec( + "doc.pdf", + b"%PDF-1.4 stub", + {}, + base_embed, + spec, + ) + assert ingestor._embed_params is not None + assert ingestor._embed_params.embed_invoke_url == "http://client/embed" + assert ingestor._embed_params.model_name == "client-m" + assert ingestor._embed_params.api_key is None + + +def test_build_graph_ingestor_applies_caption_endpoint_override_without_server_endpoint() -> None: + spec = { + "extraction_mode": "auto", + "stage_order": ["extract", "caption"], + "caption_params": {"prompt": "Describe the figure"}, + "endpoint_overrides": {"caption_invoke_url": "http://client/vlm", "caption_model_name": "vlm-x"}, + } + ingestor, _mode, _ = _build_graph_ingestor_from_spec( + "doc.pdf", + b"%PDF-1.4 stub", + {}, + None, + spec, + base_caption=None, + ) + assert ingestor._caption_params is not None + assert ingestor._caption_params.endpoint_url == "http://client/vlm" + assert ingestor._caption_params.model_name == "vlm-x" + assert ingestor._caption_params.prompt == "Describe the figure" + + +def test_build_graph_ingestor_client_cannot_override_via_embed_params() -> None: + """The denylist still protects the ordinary embed_params path.""" + base_embed = {"embed_invoke_url": "http://server/embed", "api_key": "server-key"} + spec = { + "extraction_mode": "auto", + "stage_order": ["extract", "embed"], + # Even if a malicious embed_params slipped past validation, the + # server-owned merge restores the endpoint. + "embed_params": {"embed_invoke_url": "http://attacker/", "inference_batch_size": 8}, + } + ingestor, _mode, _ = _build_graph_ingestor_from_spec("doc.pdf", b"%PDF-1.4 stub", {}, base_embed, spec) + assert ingestor._embed_params.embed_invoke_url == "http://server/embed" + + +# ---------------------------------------------------------------------- +# Client SDK +# ---------------------------------------------------------------------- + + +def test_embed_routes_endpoint_fields_to_overrides() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.embed( + embed_invoke_url="http://client/embed", + embed_model_name="client-embed", + embed_model_provider_prefix="openai", + inference_batch_size=64, + ) + payload = ing._pipeline_payload() + assert payload is not None + ov = payload["endpoint_overrides"] + assert ov["embed_invoke_url"] == "http://client/embed" + assert ov["embed_model_name"] == "client-embed" + assert ov["embed_model_provider_prefix"] == "openai" + # Shape knobs stay in embed_params; the endpoint fields never leak there. + assert payload["embed_params"]["inference_batch_size"] == 64 + assert "embed_invoke_url" not in payload["embed_params"] + assert "embed" in payload["stage_order"] + # Round-trips through the wire schema. + assert PipelineSpec.model_validate(payload).endpoint_overrides.embed_invoke_url == "http://client/embed" + + +def test_embed_via_embed_params_model_routes_endpoint() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.embed(EmbedParams(embed_invoke_url="http://client/embed", inference_batch_size=8)) + payload = ing._pipeline_payload() + assert payload["endpoint_overrides"]["embed_invoke_url"] == "http://client/embed" + assert payload["embed_params"]["inference_batch_size"] == 8 + + +def test_embed_without_endpoint_sets_no_overrides() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.embed(inference_batch_size=32) + payload = ing._pipeline_payload() + assert "endpoint_overrides" not in payload + assert payload["embed_params"]["inference_batch_size"] == 32 + + +def test_caption_routes_endpoint_fields_to_overrides() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.caption(endpoint_url="http://client/vlm", model_name="vlm-x", prompt="Describe") + payload = ing._pipeline_payload() + ov = payload["endpoint_overrides"] + assert ov["caption_invoke_url"] == "http://client/vlm" + assert ov["caption_model_name"] == "vlm-x" + # Behavioural knobs stay in caption_params; endpoint/model do not leak. + assert payload["caption_params"]["prompt"] == "Describe" + assert "endpoint_url" not in payload["caption_params"] + assert "model_name" not in payload["caption_params"] + + +def test_caption_without_endpoint_sets_no_overrides() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.caption(prompt="Describe") + payload = ing._pipeline_payload() + assert "endpoint_overrides" not in payload + assert payload["caption_params"]["prompt"] == "Describe" + + +def test_caption_rejects_local_execution_keys() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + with pytest.raises(ValueError, match="local"): + ing.caption(device="cuda:0") + + +def test_embed_and_caption_endpoint_overrides_merge() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.embed(embed_invoke_url="http://client/embed").caption(endpoint_url="http://client/vlm") + ov = ing._pipeline_payload()["endpoint_overrides"] + assert ov["embed_invoke_url"] == "http://client/embed" + assert ov["caption_invoke_url"] == "http://client/vlm" + + +def test_embed_and_caption_endpoint_overrides_keep_separate_api_keys() -> None: + ing = ServiceIngestor(base_url="http://example:7670") + ing.embed(embed_invoke_url="http://client/embed", api_key="embed-key").caption( + endpoint_url="http://client/vlm", + api_key="caption-key", + ) + ov = ing._pipeline_payload()["endpoint_overrides"] + assert ov["embed_api_key"] == "embed-key" + assert ov["caption_api_key"] == "caption-key" + assert "api_key" not in ov + + embed_out = _apply_embed_endpoint_override({"embed_invoke_url": "http://server/embed"}, ov) + caption_out = _apply_caption_endpoint_override(None, ov) + assert embed_out["api_key"] == "embed-key" + assert caption_out["api_key"] == "caption-key" + + +# ---------------------------------------------------------------------- +# Answer path: LLM endpoint override +# ---------------------------------------------------------------------- + + +class _AnswerFakeResponse: + status_code = 200 + content = json.dumps({"results": [{"hits": [{"text": "context"}]}]}).encode() + + def json(self) -> dict[str, Any]: + return json.loads(self.content.decode()) + + +class _AnswerFakeAsyncClient: + 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) -> _AnswerFakeResponse: + return _AnswerFakeResponse() + + +def _make_answer_app(monkeypatch: pytest.MonkeyPatch, tmp_path, *, llm_override: bool, prefixes=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, + ) + 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="server/model", + api_base="http://server-llm:8000/v1", + api_key="server-key", + max_tokens=128, + ), + pipeline_overrides=PipelineOverridesConfig( + endpoint_overrides=EndpointOverridesConfig(llm=llm_override, allowed_url_prefixes=prefixes or []) + ), + ) + return create_app(cfg) + + +def test_answer_llm_override_rejected_when_disabled(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_answer_app(monkeypatch, tmp_path, llm_override=False) + monkeypatch.setattr("httpx.AsyncClient", _AnswerFakeAsyncClient) + with TestClient(app) as client: + resp = client.post( + "/v1/answer", + json={"query": "q", "llm_model": "client/model", "llm_api_base": "http://client-llm:8000/v1"}, + ) + assert resp.status_code == 403 + assert "LLM endpoint overrides are disabled" in resp.json()["detail"] + + +def test_answer_llm_override_applied_when_enabled(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_answer_app(monkeypatch, tmp_path, llm_override=True) + monkeypatch.setattr("httpx.AsyncClient", _AnswerFakeAsyncClient) + + fake_llm = SimpleNamespace( + generate=lambda query, chunks, *, reasoning_enabled=None: GenerationResult( + answer="ok", latency_s=0.1, model="client/model" + ) + ) + with patch("nemo_retriever.models.llm.clients.LiteLLMClient.from_kwargs", return_value=fake_llm) as from_kwargs: + with TestClient(app) as client: + resp = client.post( + "/v1/answer", + json={ + "query": "q", + "llm_model": "client/model", + "llm_api_base": "http://client-llm:8000/v1", + "llm_api_key": "client-key", + }, + ) + assert resp.status_code == 200, resp.text + kwargs = from_kwargs.call_args.kwargs + assert kwargs["model"] == "client/model" + assert kwargs["api_base"] == "http://client-llm:8000/v1" + assert kwargs["api_key"] == "client-key" + + +def test_answer_llm_override_prefix_allowlist_enforced(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_answer_app(monkeypatch, tmp_path, llm_override=True, prefixes=["https://"]) + monkeypatch.setattr("httpx.AsyncClient", _AnswerFakeAsyncClient) + with TestClient(app) as client: + resp = client.post( + "/v1/answer", + json={"query": "q", "llm_api_base": "http://client-llm:8000/v1"}, + ) + assert resp.status_code == 403 + assert "does not match any allowed prefix" in resp.json()["detail"] + + +# ---------------------------------------------------------------------- +# Answer path: reranking (server default + per-request endpoint override) +# ---------------------------------------------------------------------- + + +class _MultiHitResponse: + status_code = 200 + content = json.dumps({"results": [{"hits": [{"text": f"c{i}"} for i in range(8)]}]}).encode() + + def json(self) -> dict[str, Any]: + return json.loads(self.content.decode()) + + +class _CapturingAsyncClient: + """Records the JSON body posted to the vectordb so tests can assert over-fetch.""" + + last_json: dict[str, Any] | None = None + + 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) -> _MultiHitResponse: + _CapturingAsyncClient.last_json = kwargs.get("json") + return _MultiHitResponse() + + +def _make_rerank_app( + monkeypatch: pytest.MonkeyPatch, + tmp_path, + *, + rerank_override: bool = False, + rerank_enabled: bool = False, + server_rerank_url: str | None = None, + server_rerank_model: str | None = None, + refine_factor: int = 4, + prefixes=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, + ) + 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="server/model", api_base="http://server-llm:8000/v1", api_key="server-key"), + nim_endpoints=NimEndpointsConfig( + rerank_invoke_url=server_rerank_url, + rerank_model_name=server_rerank_model, + api_key="server-key", + ), + rerank=RerankConfig(enabled=rerank_enabled, refine_factor=refine_factor), + pipeline_overrides=PipelineOverridesConfig( + endpoint_overrides=EndpointOverridesConfig(rerank=rerank_override, allowed_url_prefixes=prefixes or []) + ), + ) + return create_app(cfg) + + +def _stub_llm(monkeypatch: pytest.MonkeyPatch) -> list[list[str]]: + """Patch the answer LLM; return a list that captures the chunks it receives.""" + captured: list[list[str]] = [] + + def _generate(query, chunks, *, reasoning_enabled=None): + captured.append(list(chunks)) + return GenerationResult(answer="ok", latency_s=0.1, model="server/model") + + fake_llm = SimpleNamespace(generate=_generate) + monkeypatch.setattr( + "nemo_retriever.models.llm.clients.LiteLLMClient.from_kwargs", + lambda **kwargs: fake_llm, + ) + return captured + + +def test_answer_reranks_with_server_endpoint_by_default(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app( + monkeypatch, + tmp_path, + rerank_enabled=True, + server_rerank_url="http://rerank.svc/v1", + server_rerank_model="server/rerank", + refine_factor=3, + ) + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + _stub_llm(monkeypatch) + + calls: dict[str, Any] = {} + + def _fake_rerank_hits(query, hits, **kwargs): + calls["query"] = query + calls["kwargs"] = kwargs + calls["n_hits"] = len(hits) + return list(reversed(hits))[: kwargs.get("top_n")] + + monkeypatch.setattr("nemo_retriever.operators.rerank.rerank_hits", _fake_rerank_hits) + + with TestClient(app) as client: + resp = client.post("/v1/answer", json={"query": "q", "top_k": 2}) + assert resp.status_code == 200, resp.text + # Over-fetch: top_k(2) * refine_factor(3) candidates requested from vectordb. + assert _CapturingAsyncClient.last_json == {"query": "q", "top_k": 6} + assert calls["kwargs"]["rerank_invoke_url"] == "http://rerank.svc/v1" + assert calls["kwargs"]["model_name"] == "server/rerank" + assert calls["kwargs"]["api_key"] == "server-key" + assert calls["kwargs"]["top_n"] == 2 + + +def test_answer_no_rerank_when_disabled(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app(monkeypatch, tmp_path, rerank_enabled=False, server_rerank_url="http://rerank.svc/v1") + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + _stub_llm(monkeypatch) + + def _boom(*a, **k): + raise AssertionError("rerank_hits should not be called when disabled") + + monkeypatch.setattr("nemo_retriever.operators.rerank.rerank_hits", _boom) + + with TestClient(app) as client: + resp = client.post("/v1/answer", json={"query": "q", "top_k": 5}) + assert resp.status_code == 200, resp.text + # No over-fetch when not reranking. + assert _CapturingAsyncClient.last_json == {"query": "q", "top_k": 5} + + +def test_answer_rerank_override_rejected_when_disabled(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app(monkeypatch, tmp_path, rerank_override=False) + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + with TestClient(app) as client: + resp = client.post("/v1/answer", json={"query": "q", "rerank_invoke_url": "http://client/rerank"}) + assert resp.status_code == 403 + assert "rerank endpoint overrides are disabled" in resp.json()["detail"] + + +def test_answer_rerank_override_applied_when_enabled(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app( + monkeypatch, + tmp_path, + rerank_override=True, + server_rerank_url="http://rerank.svc/v1", + ) + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + _stub_llm(monkeypatch) + + calls: dict[str, Any] = {} + + def _fake_rerank_hits(query, hits, **kwargs): + calls["kwargs"] = kwargs + return hits[: kwargs.get("top_n")] + + monkeypatch.setattr("nemo_retriever.operators.rerank.rerank_hits", _fake_rerank_hits) + + with TestClient(app) as client: + resp = client.post( + "/v1/answer", + json={ + "query": "q", + "top_k": 3, + "rerank_invoke_url": "http://client/rerank", + "rerank_model_name": "client/rerank", + "rerank_api_key": "client-key", + }, + ) + assert resp.status_code == 200, resp.text + assert calls["kwargs"]["rerank_invoke_url"] == "http://client/rerank" + assert calls["kwargs"]["model_name"] == "client/rerank" + assert calls["kwargs"]["api_key"] == "client-key" + + +def test_answer_rerank_url_override_does_not_forward_server_api_key(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app( + monkeypatch, + tmp_path, + rerank_override=True, + server_rerank_url="http://rerank.svc/v1", + ) + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + _stub_llm(monkeypatch) + + calls: dict[str, Any] = {} + + def _fake_rerank_hits(query, hits, **kwargs): + calls["kwargs"] = kwargs + return hits[: kwargs.get("top_n")] + + monkeypatch.setattr("nemo_retriever.operators.rerank.rerank_hits", _fake_rerank_hits) + + with TestClient(app) as client: + resp = client.post( + "/v1/answer", + json={"query": "q", "top_k": 3, "rerank_invoke_url": "http://client/rerank"}, + ) + assert resp.status_code == 200, resp.text + assert calls["kwargs"]["rerank_invoke_url"] == "http://client/rerank" + assert calls["kwargs"]["api_key"] == "" + + +def test_answer_llm_url_override_does_not_forward_server_api_key(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_answer_app(monkeypatch, tmp_path, llm_override=True) + monkeypatch.setattr("httpx.AsyncClient", _AnswerFakeAsyncClient) + + fake_llm = SimpleNamespace( + generate=lambda query, chunks, *, reasoning_enabled=None: GenerationResult( + answer="ok", latency_s=0.1, model="client/model" + ) + ) + with patch("nemo_retriever.models.llm.clients.LiteLLMClient.from_kwargs", return_value=fake_llm) as from_kwargs: + with TestClient(app) as client: + resp = client.post( + "/v1/answer", + json={"query": "q", "llm_api_base": "http://client-llm:8000/v1"}, + ) + assert resp.status_code == 200, resp.text + assert from_kwargs.call_args.kwargs["api_key"] is None + + +def test_answer_rerank_override_prefix_allowlist_enforced(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app(monkeypatch, tmp_path, rerank_override=True, prefixes=["https://"]) + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + with TestClient(app) as client: + resp = client.post("/v1/answer", json={"query": "q", "rerank_invoke_url": "http://client/rerank"}) + assert resp.status_code == 403 + assert "does not match any allowed prefix" in resp.json()["detail"] + + +def test_answer_rerank_requested_without_endpoint_returns_400(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None: + app = _make_rerank_app(monkeypatch, tmp_path, rerank_enabled=False, server_rerank_url=None) + monkeypatch.setattr("httpx.AsyncClient", _CapturingAsyncClient) + _stub_llm(monkeypatch) + with TestClient(app) as client: + resp = client.post("/v1/answer", json={"query": "q", "rerank": True}) + assert resp.status_code == 400 + assert "no rerank endpoint" in resp.json()["detail"]