Skip to content

Commit 0f51e08

Browse files
committed
fix(server): validate chat completion messages
Signed-off-by: Pouyanpi <13303554+Pouyanpi@users.noreply.github.com>
1 parent 663efef commit 0f51e08

4 files changed

Lines changed: 404 additions & 12 deletions

File tree

nemoguardrails/rails/llm/llmrails.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -807,7 +807,7 @@ def _get_events_for_messages(self, messages: List[dict], state: Any):
807807
events.append({"type": "ContextUpdate", "data": msg["content"]})
808808
elif msg["role"] == "event":
809809
events.append(msg["event"])
810-
elif msg["role"] == "system":
810+
elif msg["role"] in ("developer", "system"):
811811
# Handle system messages - convert them to SystemMessage events
812812
events.append({"type": "SystemMessage", "content": msg["content"]})
813813
elif msg["role"] == "tool":
@@ -864,7 +864,7 @@ def _get_events_for_messages(self, messages: List[dict], state: Any):
864864
events.append({"type": "ContextUpdate", "data": msg["content"]})
865865
elif msg["role"] == "event":
866866
events.append(msg["event"])
867-
elif msg["role"] == "system":
867+
elif msg["role"] in ("developer", "system"):
868868
# Handle system messages - convert them to SystemMessage events
869869
events.append({"type": "SystemMessage", "content": msg["content"]})
870870
elif msg["role"] == "tool":

nemoguardrails/server/schemas/openai.py

Lines changed: 179 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,16 @@
1919
from typing import Annotated, Any, List, Literal, Optional, Union
2020

2121
from openai.types.chat.chat_completion import ChatCompletion
22-
from pydantic import BaseModel, BeforeValidator, ConfigDict, Field, ValidationInfo, field_validator, model_validator
22+
from pydantic import (
23+
BaseModel,
24+
BeforeValidator,
25+
ConfigDict,
26+
Field,
27+
ValidationError,
28+
ValidationInfo,
29+
field_validator,
30+
model_validator,
31+
)
2332

2433
from nemoguardrails.rails.llm.options import GenerationOptions, RailType
2534

@@ -43,29 +52,193 @@ class GuardrailsChatCompletion(ChatCompletion):
4352
guardrails: Optional[GuardrailsDataOutput] = Field(default=None, description="Guardrails specific output data.")
4453

4554

46-
class _OpenAIChatMessageSchema(BaseModel):
47-
model_config = ConfigDict(extra="allow")
55+
class _OpenAIChatMessageBase(BaseModel):
56+
model_config = ConfigDict(extra="forbid")
4857

49-
role: str
58+
59+
class _OpenAIChatMessageRoleSchema(BaseModel):
60+
role: Literal["developer", "system", "user", "assistant", "tool", "function", "context"]
61+
62+
63+
class _OpenAIPromptCacheBreakpointSchema(_OpenAIChatMessageBase):
64+
mode: Literal["explicit"]
65+
66+
67+
class _OpenAIImageURLSchema(_OpenAIChatMessageBase):
68+
url: str
69+
detail: Optional[Literal["auto", "low", "high"]] = None
70+
71+
72+
class _OpenAITextContentPartSchema(_OpenAIChatMessageBase):
73+
text: str
74+
type: Literal["text"]
75+
prompt_cache_breakpoint: Optional[_OpenAIPromptCacheBreakpointSchema] = None
76+
77+
78+
class _OpenAIImageContentPartSchema(_OpenAIChatMessageBase):
79+
image_url: _OpenAIImageURLSchema
80+
type: Literal["image_url"]
81+
prompt_cache_breakpoint: Optional[_OpenAIPromptCacheBreakpointSchema] = None
82+
83+
84+
class _OpenAIFileSchema(_OpenAIChatMessageBase):
85+
file_data: Optional[str] = None
86+
file_id: Optional[str] = None
87+
filename: Optional[str] = None
88+
89+
90+
class _OpenAIFileContentPartSchema(_OpenAIChatMessageBase):
91+
file: _OpenAIFileSchema
92+
type: Literal["file"]
93+
prompt_cache_breakpoint: Optional[_OpenAIPromptCacheBreakpointSchema] = None
94+
95+
96+
class _OpenAIRefusalContentPartSchema(_OpenAIChatMessageBase):
97+
refusal: str
98+
type: Literal["refusal"]
99+
100+
101+
_OpenAITextContentSchema = Union[str, List[_OpenAITextContentPartSchema]]
102+
_OpenAIUserContentPartSchema = Union[
103+
_OpenAITextContentPartSchema,
104+
_OpenAIImageContentPartSchema,
105+
_OpenAIFileContentPartSchema,
106+
]
107+
108+
109+
class _OpenAIDeveloperMessageSchema(_OpenAIChatMessageBase):
110+
content: _OpenAITextContentSchema
111+
role: Literal["developer"]
112+
name: Optional[str] = None
113+
114+
115+
class _OpenAISystemMessageSchema(_OpenAIChatMessageBase):
116+
content: _OpenAITextContentSchema
117+
role: Literal["system"]
118+
name: Optional[str] = None
119+
120+
121+
class _OpenAIUserMessageSchema(_OpenAIChatMessageBase):
122+
content: Union[str, List[_OpenAIUserContentPartSchema]]
123+
role: Literal["user"]
124+
name: Optional[str] = None
125+
126+
127+
class _OpenAIFunctionCallSchema(_OpenAIChatMessageBase):
128+
arguments: str
129+
name: str
130+
131+
132+
class _OpenAIFunctionToolCallSchema(_OpenAIChatMessageBase):
133+
id: str
134+
function: _OpenAIFunctionCallSchema
135+
type: Literal["function"]
136+
137+
138+
class _OpenAICustomToolSchema(_OpenAIChatMessageBase):
139+
input: str
140+
name: str
141+
142+
143+
class _OpenAICustomToolCallSchema(_OpenAIChatMessageBase):
144+
id: str
145+
custom: _OpenAICustomToolSchema
146+
type: Literal["custom"]
147+
148+
149+
_OpenAIToolCallSchema = Union[_OpenAIFunctionToolCallSchema, _OpenAICustomToolCallSchema]
150+
151+
152+
class _OpenAIAssistantMessageSchema(_OpenAIChatMessageBase):
153+
role: Literal["assistant"]
154+
content: Optional[Union[str, List[Union[_OpenAITextContentPartSchema, _OpenAIRefusalContentPartSchema]]]] = None
155+
function_call: Optional[_OpenAIFunctionCallSchema] = None
156+
name: Optional[str] = None
157+
tool_calls: Optional[List[_OpenAIToolCallSchema]] = None
158+
refusal: Optional[str] = None
159+
160+
@model_validator(mode="before")
161+
@classmethod
162+
def validate_content_or_call(cls, value: Any) -> Any:
163+
if (
164+
isinstance(value, dict)
165+
and value.get("content") is None
166+
and value.get("function_call") is None
167+
and not value.get("tool_calls")
168+
and value.get("refusal") is None
169+
):
170+
raise ValidationError.from_exception_data(
171+
cls.__name__,
172+
[{"type": "missing", "loc": ("content",), "input": value}],
173+
)
174+
return value
175+
176+
177+
class _OpenAIToolMessageSchema(_OpenAIChatMessageBase):
178+
content: _OpenAITextContentSchema
179+
role: Literal["tool"]
180+
tool_call_id: str
181+
182+
183+
class _OpenAIFunctionMessageSchema(_OpenAIChatMessageBase):
184+
content: Optional[str]
185+
name: str
186+
role: Literal["function"]
187+
188+
189+
class _GuardrailsContextMessageSchema(_OpenAIChatMessageBase):
190+
content: dict[str, Any]
191+
role: Literal["context"]
192+
193+
194+
_OpenAIChatMessageInput = Union[
195+
_OpenAIDeveloperMessageSchema,
196+
_OpenAISystemMessageSchema,
197+
_OpenAIUserMessageSchema,
198+
_OpenAIAssistantMessageSchema,
199+
_OpenAIToolMessageSchema,
200+
_OpenAIFunctionMessageSchema,
201+
_GuardrailsContextMessageSchema,
202+
]
203+
204+
_OPENAI_CHAT_MESSAGE_SCHEMAS: dict[str, type[BaseModel]] = {
205+
"developer": _OpenAIDeveloperMessageSchema,
206+
"system": _OpenAISystemMessageSchema,
207+
"user": _OpenAIUserMessageSchema,
208+
"assistant": _OpenAIAssistantMessageSchema,
209+
"tool": _OpenAIToolMessageSchema,
210+
"function": _OpenAIFunctionMessageSchema,
211+
"context": _GuardrailsContextMessageSchema,
212+
}
50213

51214

52215
def _validate_openai_chat_message(message: Any) -> Any:
53-
_OpenAIChatMessageSchema.model_validate(message)
216+
role = _OpenAIChatMessageRoleSchema.model_validate(message).role
217+
_OPENAI_CHAT_MESSAGE_SCHEMAS[role].model_validate(message)
54218
return message
55219

56220

57221
OpenAIChatMessage = Annotated[
58222
dict[str, Any],
59223
BeforeValidator(
60224
_validate_openai_chat_message,
61-
json_schema_input_type=_OpenAIChatMessageSchema,
225+
json_schema_input_type=_OpenAIChatMessageInput,
62226
),
63227
]
64228

65229

66230
class OpenAIChatCompletionRequest(BaseModel):
67231
"""Standard OpenAI chat completion request parameters."""
68232

233+
@model_validator(mode="before")
234+
@classmethod
235+
def reject_unsupported_audio(cls, data: Any) -> Any:
236+
if isinstance(data, dict):
237+
modalities = data.get("modalities")
238+
if "audio" in data or (isinstance(modalities, list) and "audio" in modalities):
239+
raise ValueError("Audio input and output are not supported.")
240+
return data
241+
69242
messages: Optional[List[OpenAIChatMessage]] = Field(
70243
default=None,
71244
description="The list of messages in the current conversation.",

0 commit comments

Comments
 (0)