Skip to content

Commit cb9fb82

Browse files
committed
little update
update update update update
1 parent 573f001 commit cb9fb82

2 files changed

Lines changed: 41 additions & 10 deletions

File tree

‎cozeloop/integration/langchain/trace_callback.py‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -127,10 +127,13 @@ def on_chain_start(self, serialized: Dict[str, Any], inputs: Dict[str, Any], **k
127127
flow_span = None
128128
try:
129129
if kwargs.get('run_type', '') == 'prompt' or kwargs.get('name', '') == 'ChatPromptTemplate':
130-
flow_span = self._new_flow_span(kwargs['name'], kwargs['name'], **kwargs)
130+
flow_span = self._new_flow_span(kwargs['name'], 'prompt', **kwargs)
131131
self._on_prompt_start(flow_span, serialized, inputs, **kwargs)
132132
else:
133-
flow_span = self._new_flow_span(kwargs['name'], kwargs['name'], **kwargs)
133+
span_type = 'chain'
134+
if kwargs['name'] == 'LangGraph': # LangGraph is agent span_type,for trajectory evaluation aggregate to an agent
135+
span_type = 'agent'
136+
flow_span = self._new_flow_span(kwargs['name'], span_type, **kwargs)
134137
flow_span.set_tags({'input': _convert_2_json(inputs)})
135138
except Exception as e:
136139
if flow_span is not None:
@@ -505,7 +508,7 @@ def _convert_inputs(inputs: Any) -> Any:
505508
for each in inputs:
506509
format_inputs.append(_convert_inputs(each))
507510
return format_inputs
508-
if isinstance(inputs, AIMessageChunk):
511+
if isinstance(inputs, (AIMessageChunk, AIMessage)):
509512
"""
510513
Must be before BaseMessage.
511514
"""

‎cozeloop/integration/langchain/trace_model/llm_model.py‎

Lines changed: 35 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)