From 9cac04610e8065cb24b4c20013dbc101277ad0c8 Mon Sep 17 00:00:00 2001 From: LingYi-Liang <270316896+LingYi-Liang@users.noreply.github.com> Date: Mon, 5 Oct 2026 03:51:50 +0800 Subject: [PATCH] =?UTF-8?q?fix(llm):=20=E4=BF=AE=E5=A4=8D=E6=B5=81?= =?UTF-8?q?=E5=BC=8F=E5=B7=A5=E5=85=B7=E8=B0=83=E7=94=A8=E7=9A=84=E5=90=88?= =?UTF-8?q?=E5=B9=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ms_agent/llm/openai_llm.py | 66 +++---- ms_agent/llm/transport/openai_compat.py | 71 ++++--- tests/llm/test_stream_tool_calls.py | 239 ++++++++++++++++++++++++ 3 files changed, 302 insertions(+), 74 deletions(-) create mode 100644 tests/llm/test_stream_tool_calls.py diff --git a/ms_agent/llm/openai_llm.py b/ms_agent/llm/openai_llm.py index 730205e2a..1208943ce 100644 --- a/ms_agent/llm/openai_llm.py +++ b/ms_agent/llm/openai_llm.py @@ -437,46 +437,40 @@ def _merge_stream_message(self, pre_message_chunk: Optional[Message], Note: - **Content Merging**: Textual content (`content`, `reasoning_content`) is appended cumulatively. - - **Tool Call Merging**: If the same tool call index appears in consecutive chunks, - its `arguments` and `tool_name` will be updated incrementally. + - **Tool Call Merging**: Calls are matched by index, including multiple + calls in one chunk and calls interleaved across chunks. - If a new tool call index is found, it will be added as a new entry in `tool_calls`. """ if not pre_message_chunk: - return message_chunk - message = deepcopy(pre_message_chunk) - message.reasoning_content += message_chunk.reasoning_content - message.content += message_chunk.content + message = deepcopy(message_chunk) + if message_chunk.tool_calls: + message.tool_calls = [] + else: + message = deepcopy(pre_message_chunk) + message.reasoning_content += message_chunk.reasoning_content + message.content += message_chunk.content if message_chunk.tool_calls: - if message.tool_calls: - if message.tool_calls[-1]['index'] == message_chunk.tool_calls[ - 0]['index']: - if message_chunk.tool_calls[0]['id']: - message.tool_calls[-1][ - 'id'] = message_chunk.tool_calls[0]['id'] - if message_chunk.tool_calls[0]['arguments']: - if message.tool_calls[-1]['arguments']: - message.tool_calls[-1][ - 'arguments'] += message_chunk.tool_calls[0][ - 'arguments'] - else: - # message.tool_calls[-1]['arguments'] may be None - message.tool_calls[-1][ - 'arguments'] = message_chunk.tool_calls[0][ - 'arguments'] - if message_chunk.tool_calls[0]['tool_name']: - message.tool_calls[-1][ - 'tool_name'] = message_chunk.tool_calls[0][ - 'tool_name'] - else: - message.tool_calls.append( - ToolCall( - id=message_chunk.tool_calls[0]['id'], - arguments=message_chunk.tool_calls[0]['arguments'], - type='function', - tool_name=message_chunk.tool_calls[0]['tool_name'], - index=message_chunk.tool_calls[0]['index'])) - else: - message.tool_calls = message_chunk.tool_calls + if not message.tool_calls: + message.tool_calls = [] + calls_by_index = { + call['index']: call + for call in message.tool_calls + } + for delta in message_chunk.tool_calls: + call = calls_by_index.get(delta['index']) + if call is None: + call = deepcopy(delta) + call['type'] = call.get('type') or 'function' + message.tool_calls.append(call) + calls_by_index[delta['index']] = call + continue + if delta['id']: + call['id'] = delta['id'] + if delta['arguments']: + call['arguments'] = (call['arguments'] + or '') + delta['arguments'] + if delta['tool_name']: + call['tool_name'] = delta['tool_name'] return message def _stream_continue_generate(self, diff --git a/ms_agent/llm/transport/openai_compat.py b/ms_agent/llm/transport/openai_compat.py index 4dbe214ee..6de730e53 100644 --- a/ms_agent/llm/transport/openai_compat.py +++ b/ms_agent/llm/transport/openai_compat.py @@ -556,45 +556,40 @@ def _extract_reasoning_tokens(usage_obj: Any) -> int: def _merge_stream_message(self, pre_message_chunk: Optional[Message], message_chunk: Message) -> Optional[Message]: if not pre_message_chunk: - return message_chunk - message = deepcopy(pre_message_chunk) - # Coalesce to '' first: either side can be None (first chunk of a tool - # call / reasoning, or a prior message without reasoning content), and - # None += str raises TypeError. - message.reasoning_content = (message.reasoning_content or '') + ( - message_chunk.reasoning_content or '') - message.content = (message.content or '') + ( - message_chunk.content or '') + message = deepcopy(message_chunk) + if message_chunk.tool_calls: + message.tool_calls = [] + else: + message = deepcopy(pre_message_chunk) + # Coalesce to '' first: either side can be None (first chunk of a tool + # call / reasoning, or a prior message without reasoning content), and + # None += str raises TypeError. + message.reasoning_content = (message.reasoning_content or '') + ( + message_chunk.reasoning_content or '') + message.content = (message.content or '') + ( + message_chunk.content or '') if message_chunk.tool_calls: - if message.tool_calls: - if message.tool_calls[-1]['index'] == message_chunk.tool_calls[ - 0]['index']: - if message_chunk.tool_calls[0]['id']: - message.tool_calls[-1][ - 'id'] = message_chunk.tool_calls[0]['id'] - if message_chunk.tool_calls[0]['arguments']: - if message.tool_calls[-1]['arguments']: - message.tool_calls[-1][ - 'arguments'] += message_chunk.tool_calls[0][ - 'arguments'] - else: - message.tool_calls[-1][ - 'arguments'] = message_chunk.tool_calls[0][ - 'arguments'] - if message_chunk.tool_calls[0]['tool_name']: - message.tool_calls[-1][ - 'tool_name'] = message_chunk.tool_calls[0][ - 'tool_name'] - else: - message.tool_calls.append( - ToolCall( - id=message_chunk.tool_calls[0]['id'], - arguments=message_chunk.tool_calls[0]['arguments'], - type='function', - tool_name=message_chunk.tool_calls[0]['tool_name'], - index=message_chunk.tool_calls[0]['index'])) - else: - message.tool_calls = message_chunk.tool_calls + if not message.tool_calls: + message.tool_calls = [] + calls_by_index = { + call['index']: call + for call in message.tool_calls + } + for delta in message_chunk.tool_calls: + call = calls_by_index.get(delta['index']) + if call is None: + call = deepcopy(delta) + call['type'] = call.get('type') or 'function' + message.tool_calls.append(call) + calls_by_index[delta['index']] = call + continue + if delta['id']: + call['id'] = delta['id'] + if delta['arguments']: + call['arguments'] = (call['arguments'] + or '') + delta['arguments'] + if delta['tool_name']: + call['tool_name'] = delta['tool_name'] return message def _stream_continue_generate(self, diff --git a/tests/llm/test_stream_tool_calls.py b/tests/llm/test_stream_tool_calls.py new file mode 100644 index 000000000..da522b10e --- /dev/null +++ b/tests/llm/test_stream_tool_calls.py @@ -0,0 +1,239 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Replay Chat Completions streams through the SDK and public LLM entry point.""" +import json +import threading +import unittest +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from unittest.mock import patch + +from omegaconf import OmegaConf + +from ms_agent.llm import LLM +from ms_agent.llm.utils import Message, Tool + + +def tool_delta(index, arguments=None, call_id=None, name=None): + delta = {'index': index, 'function': {'arguments': arguments}} + if call_id is not None: + delta.update(id=call_id, type='function') + if name is not None: + delta['function']['name'] = name + return delta + + +def event(delta, finish_reason=None): + return { + 'id': 'stream-test', 'object': 'chat.completion.chunk', + 'created': 0, 'model': 'test-model', + 'choices': [{'index': 0, 'delta': delta, + 'finish_reason': finish_reason}], + } + + +class StreamToolCallCases: + use_router = False + + def setUp(self): + self.requests = [] + self.events = [] + self.json_response = None + case = self + + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + body = self.rfile.read(int(self.headers['Content-Length'])) + case.requests.append((self.path, json.loads(body))) + self.send_response(200) + if case.json_response is not None: + payload = json.dumps(case.json_response).encode() + self.send_header('Content-Type', 'application/json') + else: + payload = ''.join('data: ' + json.dumps(item) + '\n\n' + for item in case.events).encode() + payload += b'data: [DONE]\n\n' + self.send_header('Content-Type', 'text/event-stream') + self.send_header('Content-Length', str(len(payload))) + self.end_headers() + self.wfile.write(payload) + + def log_message(self, *args): + pass + + self.server = ThreadingHTTPServer(('127.0.0.1', 0), Handler) + self.thread = threading.Thread( + target=self.server.serve_forever, + kwargs={'poll_interval': 0.01}, daemon=True) + self.thread.start() + self.addCleanup(self._close_server) + proxy = patch.dict('os.environ', {'NO_PROXY': '127.0.0.1'}) + proxy.start() + self.addCleanup(proxy.stop) + config = OmegaConf.create({ + 'llm': { + 'service': 'openai', 'model': 'test-model', + 'use_provider_router': self.use_router, + 'openai_api_key': 'local-test-key', + 'openai_base_url': + f'http://127.0.0.1:{self.server.server_port}/v1', + }, + 'generation_config': {}, + }) + self.llm = LLM.from_config(config) + self.engine = self.llm.transport if self.use_router else self.llm + self.addCleanup(self.engine.client.close) + self.tools = [Tool( + tool_name='weather', description='Get weather for a city', + parameters={'type': 'object', + 'properties': {'city': {'type': 'string'}}, + 'required': ['city']})] + + def _close_server(self): + self.server.shutdown() + self.server.server_close() + self.thread.join(timeout=5) + self.assertFalse(self.thread.is_alive()) + + def stream(self, deltas): + self.events = [event(delta) for delta in deltas] + finish_reason = 'tool_calls' if any( + delta.get('tool_calls') for delta in deltas) else 'stop' + self.events.append(event({}, finish_reason)) + self.events.append({ + 'id': 'stream-test', 'object': 'chat.completion.chunk', + 'created': 0, 'model': 'test-model', 'choices': [], + 'usage': {'prompt_tokens': 7, 'completion_tokens': 5, + 'total_tokens': 12}, + }) + messages = list(self.llm.generate( + [Message(role='user', content='Check the weather')], + tools=self.tools, stream=True)) + self.assertEqual(len(self.requests), 1) + path, request = self.requests[0] + self.assertEqual(path, '/v1/chat/completions') + self.assertTrue(request['stream']) + self.assertEqual(request['tools'][0]['function']['name'], 'weather') + self.assertEqual(messages[-1].prompt_tokens, 7) + self.assertEqual(messages[-1].completion_tokens, 5) + return messages + + def assert_calls(self, message, cities): + self.assertEqual(len(message.tool_calls), len(cities)) + for index, (call, city) in enumerate(zip(message.tool_calls, cities)): + self.assertEqual(call['index'], index) + self.assertEqual(call['id'], f'call-{index}') + self.assertEqual(call['type'], 'function') + self.assertEqual(call['tool_name'], 'weather') + self.assertEqual(json.loads(call['arguments']), {'city': city}) + + def test_multiple_calls_in_one_later_chunk(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, '{', 'call-0', 'weather')]}, + {'tool_calls': [tool_delta(0, '"city":"Paris"}'), + tool_delta(1, '{"city":"Tokyo"}', + 'call-1', 'weather')]}, + ]) + self.assert_calls(messages[-1], ['Paris', 'Tokyo']) + self.assertEqual(messages[0].tool_calls[0]['arguments'], '{') + + def test_interleaved_calls(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, '{', 'call-0', 'weather')]}, + {'tool_calls': [tool_delta(1, '{', 'call-1', 'weather')]}, + {'tool_calls': [tool_delta(0, '"city":"Paris"}')]}, + {'tool_calls': [tool_delta(1, '"city":"Tokyo"}')]}, + ]) + self.assert_calls(messages[-1], ['Paris', 'Tokyo']) + + def test_multiple_calls_in_first_and_following_chunks(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, '{', 'call-0', 'weather'), + tool_delta(1, '{', 'call-1', 'weather')]}, + {'tool_calls': [tool_delta(1, '"city":"Tokyo"}'), + tool_delta(0, '"city":"Paris"}')]}, + ]) + self.assert_calls(messages[-1], ['Paris', 'Tokyo']) + + def test_sequential_calls_still_work(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, '{', 'call-0', 'weather')]}, + {'tool_calls': [tool_delta(0, '"city":"Paris"}')]}, + {'tool_calls': [tool_delta(1, '{', 'call-1', 'weather')]}, + {'tool_calls': [tool_delta(1, '"city":"Tokyo"}')]}, + ]) + self.assert_calls(messages[-1], ['Paris', 'Tokyo']) + + def test_repeated_index_in_first_chunk(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, '{', 'call-0', 'weather'), + tool_delta(0, '"city":"Paris"}')]}, + ]) + self.assert_calls(messages[-1], ['Paris']) + + def test_none_arguments_then_single_call(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, None, 'call-0', 'weather')]}, + {'tool_calls': [tool_delta(0, '{"city":"Paris"}')]}, + ]) + self.assert_calls(messages[-1], ['Paris']) + + def test_metadata_arriving_after_arguments(self): + messages = self.stream([ + {'tool_calls': [tool_delta(0, '{"city":"Paris"}', + 'call-0', 'weather')]}, + {'tool_calls': [tool_delta(1, '{')]}, + {'tool_calls': [tool_delta(1, '"city":"Tokyo"}', + 'call-1', 'weather')]}, + ]) + self.assert_calls(messages[-1], ['Paris', 'Tokyo']) + self.assertEqual( + messages[-1].to_dict_clean()['tool_calls'][1]['type'], 'function') + + def test_text_before_multiple_calls(self): + messages = self.stream([ + {'content': 'Checking ', 'reasoning_content': 'Need '}, + {'content': 'weather.', 'reasoning_content': 'two cities.'}, + {'tool_calls': [tool_delta(0, '{"city":"Paris"}', + 'call-0', 'weather'), + tool_delta(1, '{"city":"Tokyo"}', + 'call-1', 'weather')]}, + ]) + self.assert_calls(messages[-1], ['Paris', 'Tokyo']) + self.assertEqual(messages[-1].content, 'Checking weather.') + self.assertEqual(messages[-1].reasoning_content, 'Need two cities.') + + def test_text_only_stream(self): + messages = self.stream([{'content': 'Hello '}, {'content': 'world'}]) + self.assertEqual(messages[-1].content, 'Hello world') + self.assertFalse(messages[-1].tool_calls) + + def test_non_streaming_calls_still_work(self): + calls = [tool_delta(index, json.dumps({'city': city}), + f'call-{index}', 'weather') + for index, city in enumerate(['Paris', 'Tokyo'])] + for call in calls: + call.pop('index') + self.json_response = { + 'id': 'non-stream-test', 'object': 'chat.completion', + 'created': 0, 'model': 'test-model', + 'choices': [{'index': 0, 'finish_reason': 'tool_calls', + 'message': {'role': 'assistant', 'content': None, + 'tool_calls': calls}}], + 'usage': {'prompt_tokens': 7, 'completion_tokens': 5, + 'total_tokens': 12}, + } + result = self.llm.generate( + [Message(role='user', content='Check the weather')], + tools=self.tools, stream=False) + self.assert_calls(result, ['Paris', 'Tokyo']) + + +class TestLegacyStreamToolCalls(StreamToolCallCases, unittest.TestCase): + pass + + +class TestRouterStreamToolCalls(StreamToolCallCases, unittest.TestCase): + use_router = True + + +if __name__ == '__main__': + unittest.main()