Skip to content

Commit aee16a8

Browse files
authored
fix(server): correct thread and model injection errors (#2240)
Fixes two server issues: Return HTTP 400 when thread_id is used without a configured datastore. Preserve the main model’s api_key_env_var during request model injection. Signed-off-by: Pouyanpi <13303554+Pouyanpi@users.noreply.github.com>
1 parent 3e81d0e commit aee16a8

2 files changed

Lines changed: 90 additions & 8 deletions

File tree

nemoguardrails/server/api.py

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -381,9 +381,14 @@ def _update_models_in_config(config: RailsConfig, main_model: Model) -> RailsCon
381381
break
382382

383383
if main_model_index is not None:
384-
parameters = {**models[main_model_index].parameters, **main_model.parameters}
385-
models[main_model_index] = main_model
386-
models[main_model_index].parameters = parameters
384+
configured_model = models[main_model_index]
385+
parameters = {**configured_model.parameters, **main_model.parameters}
386+
models[main_model_index] = main_model.model_copy(
387+
update={
388+
"api_key_env_var": main_model.api_key_env_var or configured_model.api_key_env_var,
389+
"parameters": parameters,
390+
}
391+
)
387392
else:
388393
models.append(main_model)
389394

@@ -636,12 +641,9 @@ async def chat_completion(body: GuardrailsChatCompletionRequest, request: Reques
636641

637642
if body.guardrails.thread_id:
638643
if datastore is None:
639-
raise RuntimeError("No DataStore has been configured.")
640-
# We make sure the `thread_id` meets the minimum complexity requirement.
641-
if len(body.guardrails.thread_id) < 16:
642644
raise HTTPException(
643-
status_code=422,
644-
detail="The `thread_id` must have a minimum length of 16 characters.",
645+
status_code=400,
646+
detail="Conversation threads are not enabled on this server.",
645647
)
646648

647649
# Fetch the existing thread messages. For easier management, we prepend

tests/server/test_api.py

Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,10 @@
2424
pytest.importorskip("openai", reason="openai is required for server tests")
2525
from fastapi.testclient import TestClient
2626

27+
from nemoguardrails import RailsConfig
28+
from nemoguardrails.guardrails.model_engine import ModelEngine
29+
from nemoguardrails.llm.models.openai_chat import OpenAIChatModel
30+
from nemoguardrails.rails import LLMRails
2731
from nemoguardrails.server import api
2832
from nemoguardrails.server.api import _format_streaming_response
2933
from nemoguardrails.server.schemas.openai import GuardrailsChatCompletionRequest
@@ -166,6 +170,82 @@ def test_model_field_independent_of_config_id():
166170
assert request_body.guardrails.config_ids == ["test_config"]
167171

168172

173+
@pytest.fixture
174+
def injected_model_config(monkeypatch):
175+
monkeypatch.setenv("CUSTOM_MAIN_API_KEY", "main-key")
176+
monkeypatch.setenv("MAIN_MODEL_BASE_URL", "https://request.example/v1")
177+
config = RailsConfig.from_content(
178+
config={
179+
"models": [
180+
{
181+
"type": "main",
182+
"engine": "nim",
183+
"model": "configured-model",
184+
"api_key_env_var": "CUSTOM_MAIN_API_KEY",
185+
"parameters": {
186+
"base_url": "https://configured.example/v1",
187+
"default_headers": {"X-Tenant": "acme"},
188+
},
189+
}
190+
]
191+
}
192+
)
193+
return api._inject_model(config, "requested-model")
194+
195+
196+
def test_inject_model_preserves_main_model_api_key_env_var(injected_model_config):
197+
main_model = injected_model_config.models[0]
198+
headers = ModelEngine(main_model)._prepare_request([{"role": "user", "content": "hi"}]).headers
199+
200+
assert main_model.model == "requested-model"
201+
assert main_model.engine == "custom_llm"
202+
assert main_model.api_key_env_var == "CUSTOM_MAIN_API_KEY"
203+
assert main_model.parameters == {
204+
"base_url": "https://request.example/v1",
205+
"default_headers": {"X-Tenant": "acme"},
206+
}
207+
assert headers["Authorization"] == "Bearer main-key"
208+
assert headers["X-Tenant"] == "acme"
209+
210+
211+
def test_inject_model_preserves_main_model_api_key_for_llmrails(injected_model_config):
212+
rails = LLMRails(config=injected_model_config.model_copy(deep=True))
213+
214+
assert isinstance(rails.llm, OpenAIChatModel)
215+
headers = rails.llm._client._build_headers()
216+
assert rails.llm.model_name == "requested-model"
217+
assert rails.llm.provider_name == "custom_llm"
218+
assert headers["Authorization"] == "Bearer main-key"
219+
assert headers["X-Tenant"] == "acme"
220+
221+
222+
def test_thread_id_without_datastore_returns_400(monkeypatch):
223+
mock_rails = AsyncMock()
224+
mock_rails.config = RailsConfig.from_content(config={"models": []})
225+
monkeypatch.setattr(api, "datastore", None)
226+
227+
with patch("nemoguardrails.server.api._get_rails", new=AsyncMock(return_value=mock_rails)):
228+
response = client.post(
229+
"/v1/chat/completions",
230+
json={
231+
"model": "gpt-4o",
232+
"messages": [{"role": "user", "content": "Hello"}],
233+
"guardrails": {
234+
"config_id": "test_config",
235+
"thread_id": "0123456789abcdef",
236+
},
237+
},
238+
)
239+
240+
assert response.status_code == 400
241+
assert response.json()["error"] == {
242+
"message": "Conversation threads are not enabled on this server.",
243+
"type": "invalid_request_error",
244+
"param": None,
245+
"code": None,
246+
}
247+
248+
169249
def test_request_body_state():
170250
"""Test GuardrailsChatCompletionRequest state handling."""
171251
data = {

0 commit comments

Comments
 (0)