1919from typing import Annotated , Any , List , Literal , Optional , Union
2020
2121from 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
2433from 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
52215def _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
57221OpenAIChatMessage = 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
66230class 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