|
24 | 24 | pytest.importorskip("openai", reason="openai is required for server tests") |
25 | 25 | from fastapi.testclient import TestClient |
26 | 26 |
|
| 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 |
27 | 31 | from nemoguardrails.server import api |
28 | 32 | from nemoguardrails.server.api import _format_streaming_response |
29 | 33 | from nemoguardrails.server.schemas.openai import GuardrailsChatCompletionRequest |
@@ -166,6 +170,82 @@ def test_model_field_independent_of_config_id(): |
166 | 170 | assert request_body.guardrails.config_ids == ["test_config"] |
167 | 171 |
|
168 | 172 |
|
| 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 | + |
169 | 249 | def test_request_body_state(): |
170 | 250 | """Test GuardrailsChatCompletionRequest state handling.""" |
171 | 251 | data = { |
|
0 commit comments