diff --git a/openrag/components/llm.py b/openrag/components/llm.py index bfaf4ed76..6c8bb2451 100644 --- a/openrag/components/llm.py +++ b/openrag/components/llm.py @@ -32,12 +32,13 @@ def _extract_llm_overrides(self, request: dict): if llm_override.get("model"): payload["model"] = llm_override["model"] - base_url = (llm_override.get("base_url") or self._base_url).rstrip("/") - api_key = llm_override.get("api_key") or self._api_key - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}", - } + # Only `model` may be overridden by the client. `base_url` / `api_key` + # are deliberately NOT read from the request: honoring a client-supplied + # endpoint enables SSRF (the server would issue requests to an arbitrary + # host) and would leak the server's API key to that host. The endpoint + # and credentials always come from server configuration. + base_url = self._base_url.rstrip("/") + headers = self.headers return payload, base_url, headers diff --git a/openrag/components/test_llm.py b/openrag/components/test_llm.py index 822313be3..3ebb93332 100644 --- a/openrag/components/test_llm.py +++ b/openrag/components/test_llm.py @@ -29,34 +29,39 @@ def test_no_override_uses_defaults(self, llm): assert base_url == "http://default-llm:8000/v1" assert headers["Authorization"] == "Bearer default-key" - def test_override_all_fields(self, llm): + def test_model_override_applied(self, llm): request = { "model": "openrag-my-partition", "messages": [{"role": "user", "content": "hello"}], "stream": False, - "metadata": { - "llm_override": { - "base_url": "http://custom-llm:9000/v1", - "api_key": "custom-key", - "model": "custom-model", - } - }, + "metadata": {"llm_override": {"model": "custom-model"}}, } payload, base_url, headers = llm._extract_llm_overrides(request) assert payload["model"] == "custom-model" - assert base_url == "http://custom-llm:9000/v1" - assert headers["Authorization"] == "Bearer custom-key" + # Endpoint and credentials always come from server config. + assert base_url == "http://default-llm:8000/v1" + assert headers["Authorization"] == "Bearer default-key" - def test_trailing_slash_stripped_from_base_url(self, llm): + def test_client_base_url_and_api_key_override_ignored(self, llm): + # SSRF / key-exfiltration guard: a client-supplied base_url/api_key + # must never be honored. request = { "model": "openrag-my-partition", "stream": False, - "metadata": {"llm_override": {"base_url": "http://custom:8000/v1///"}}, + "metadata": { + "llm_override": { + "base_url": "http://169.254.169.254/latest/meta-data", + "api_key": "attacker-key", + "model": "custom-model", + } + }, } - _, base_url, _ = llm._extract_llm_overrides(request) + payload, base_url, headers = llm._extract_llm_overrides(request) - assert base_url == "http://custom:8000/v1" + assert payload["model"] == "custom-model" + assert base_url == "http://default-llm:8000/v1" + assert headers["Authorization"] == "Bearer default-key" def test_request_params_forwarded_to_payload(self, llm): request = { diff --git a/openrag/models/openai.py b/openrag/models/openai.py index b77305f70..88915c252 100644 --- a/openrag/models/openai.py +++ b/openrag/models/openai.py @@ -32,7 +32,7 @@ class OpenAIChatCompletionRequest(BaseModel): "websearch": False, "llm_override": None, }, - description="Extra custom parameters. Supports 'llm_override' object with optional 'base_url', 'api_key', and 'model' to override the downstream LLM endpoint.", + description="Extra custom parameters. Supports 'llm_override' object with an optional 'model' to override the downstream model name. The LLM endpoint and credentials are fixed by server configuration and cannot be overridden by the client.", )