@@ -51,6 +51,8 @@ class Message:
5151 content : Optional [Union [str , List [Union [dict , Parts ]], dict ]] = None
5252 parts : Optional [List [Parts ]] = None
5353 tool_calls : List [ToolCall ] = None
54+ metadata : Optional [dict ] = None
55+ reasoning_content : Optional [str ] = None
5456
5557 def __post_init__ (self ):
5658 if self .role is not None and (self .role == 'AIMessageChunk' or self .role == 'ai' ):
@@ -193,12 +195,7 @@ def to_json(self):
193195 for i , generation in enumerate (self .generations ):
194196 choice : Choice = None
195197 if isinstance (generation , ChatGeneration ):
196- tool_calls = convert_tool_calls_by_additional_kwargs (generation .message .additional_kwargs .get ('tool_calls' , []))
197- if len (tool_calls ) == 0 and 'function_call' in generation .message .additional_kwargs :
198- function_call = generation .message .additional_kwargs .get ('function_call' , {})
199- function = ToolFunction (name = function_call .get ('name' , '' ), arguments = json .loads (function_call .get ('arguments' , {})))
200- tool_calls .append (ToolCall (function = function , type = 'function_call(deprecated)' ))
201- message = Message (role = generation .message .type , content = generation .message .content , tool_calls = tool_calls )
198+ message = convert_output_message (generation .message )
202199 choice = Choice (index = i , message = message , finish_reason = generation .generation_info .get ('finish_reason' , '' ))
203200 elif isinstance (generation , Generation ):
204201 choice = Choice (index = i , message = Message (content = generation .text ))
@@ -234,4 +231,35 @@ def convert_tool_calls_by_additional_kwargs(tool_calls: list) -> List[ToolCall]:
234231 logger .error (f"convert_tool_calls_by_additional_kwargs failed, error: { e } , tool_call.function.arguments: { raw_args } " )
235232 function = ToolFunction (name = tool_call .get ('function' , {}).get ('name' , '' ), arguments = final_args )
236233 format_tool_calls .append (ToolCall (id = tool_call .get ('id' , '' ), type = tool_call .get ('type' , '' ), function = function ))
237- return format_tool_calls
234+ return format_tool_calls
235+
236+
237+ def convert_output_message (message : BaseMessage ) -> Message :
238+ if message is None :
239+ return None
240+ tool_calls = convert_tool_calls_by_additional_kwargs (message .additional_kwargs .get ('tool_calls' , []))
241+ if len (tool_calls ) == 0 and isinstance (message , (AIMessage , AIMessageChunk )):
242+ tool_calls = convert_tool_calls_by_raw (message .tool_calls )
243+ if len (tool_calls ) == 0 and 'function_call' in message .additional_kwargs :
244+ function_call = message .additional_kwargs .get ('function_call' , {})
245+ try :
246+ arg = json .loads (function_call .get ('arguments' , {}))
247+ except Exception as e :
248+ logging .error (f"ModelTraceOutput.to_json arguments loads failed, exception: { e } " )
249+ arg = {}
250+ function = ToolFunction (name = function_call .get ('name' , '' ), arguments = arg )
251+ tool_calls .append (ToolCall (function = function , type = 'function_call(deprecated)' ))
252+ metadata = {}
253+ if message .response_metadata is not None :
254+ if message .response_metadata .get ('id' , '' ):
255+ response_id = message .response_metadata .get ('id' , '' )
256+ metadata ['id' ] = response_id
257+ message = Message (
258+ role = message .type ,
259+ content = message .content ,
260+ tool_calls = tool_calls ,
261+ metadata = metadata ,
262+ reasoning_content = message .additional_kwargs .get ('reasoning_content' , '' ),
263+ )
264+
265+ return message
0 commit comments