Skip to content

Commit 804602c

Browse files
suluyanaluyan
andauthored
feat: add max_continue_runs (#685)
Co-authored-by: luyan <suluyan.sly@aliabab-inc.com>
1 parent 15c2bf3 commit 804602c

3 files changed

Lines changed: 57 additions & 21 deletions

File tree

.github/workflows/citest.yaml

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,8 @@ jobs:
6666
- name: Run tests
6767
env:
6868
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
69+
DASHSCOPE_API_KEY: ${{ secrets.DASHSCOPE_API_KEY }}
70+
MODELSCOPE_API_KEY: ${{ secrets.MODELSCOPE_API_KEY }}
6971

7072
shell: bash
7173
run: bash .dev_scripts/dockerci.sh

ms_agent/llm/openai_llm.py

Lines changed: 34 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,8 @@
1212

1313
logger = get_logger()
1414

15+
MAX_CONTINUE_RUNS = 3
16+
1517

1618
class OpenAI(LLM):
1719
"""Base Class for OpenAI SDK LLMs.
@@ -28,14 +30,18 @@ class OpenAI(LLM):
2830
'role', 'content', 'tool_calls', 'partial', 'prefix', 'tool_call_id'
2931
}
3032

31-
def __init__(self,
32-
config: DictConfig,
33-
base_url: Optional[str] = None,
34-
api_key: Optional[str] = None):
33+
def __init__(
34+
self,
35+
config: DictConfig,
36+
base_url: Optional[str] = None,
37+
api_key: Optional[str] = None,
38+
):
3539
super().__init__(config)
3640
assert_package_exist('openai')
3741
import openai
3842
self.model: str = config.llm.model
43+
self.max_continue_runs = getattr(config.llm, 'max_continue_runs',
44+
None) or MAX_CONTINUE_RUNS
3945
base_url = base_url or config.llm.openai_base_url
4046
api_key = api_key or config.llm.openai_api_key
4147

@@ -76,6 +82,7 @@ def format_tools(self,
7682
def generate(self,
7783
messages: List[Message],
7884
tools: Optional[List[Tool]] = None,
85+
max_continue_runs: Optional[int] = None,
7986
**kwargs) -> Message | Generator[Message, None, None]:
8087
"""Generates a response based on the given conversation history and optional tools.
8188
@@ -99,11 +106,14 @@ def generate(self,
99106

100107
# Complex task may produce long response
101108
# Call continue_generate to keep generating if the finish_reason is `length`
109+
max_continue_runs = max_continue_runs or self.max_continue_runs
102110
if stream:
103111
return self._stream_continue_generate(messages, completion, tools,
112+
max_continue_runs - 1,
104113
**args)
105114
else:
106-
return self._continue_generate(messages, completion, tools, **args)
115+
return self._continue_generate(messages, completion, tools,
116+
max_continue_runs - 1, **args)
107117

108118
def _call_llm(self,
109119
messages: List[Message],
@@ -180,6 +190,7 @@ def _stream_continue_generate(self,
180190
messages: List[Message],
181191
completion: Iterable,
182192
tools: Optional[List[Tool]] = None,
193+
max_runs: Optional[int] = None,
183194
**kwargs) -> Generator[Message, None, None]:
184195
"""Recursively continues generating until the model finishes naturally in streaming mode.
185196
@@ -193,7 +204,6 @@ def _stream_continue_generate(self,
193204
Message: Incremental chunks of the generated message.
194205
"""
195206
message = None
196-
197207
for chunk in completion:
198208
message_chunk = self._stream_format_output_message(chunk)
199209
message = self._merge_stream_message(message, message_chunk)
@@ -208,15 +218,18 @@ def _stream_continue_generate(self,
208218
# The stream may end without a final usage chunk, which is acceptable.
209219
pass
210220
first_run = not messages[-1].to_dict().get('partial', False)
211-
if chunk.choices[0].finish_reason in ['length', 'null']:
212-
print(
213-
f'finish_reason: {chunk.choices[0].finish_reason}, continue generate.'
221+
if chunk.choices[0].finish_reason in [
222+
'length', 'null'
223+
] and (max_runs is None or max_runs != 0):
224+
logger.info(
225+
f'finish_reason: {chunk.choices[0].finish_reason}, continue generate.'
214226
)
215-
216227
completion = self._call_llm_for_continue_gen(
217228
messages, message, tools, **kwargs)
218229
for chunk in self._stream_continue_generate(
219-
messages, completion, tools, **kwargs):
230+
messages, completion, tools,
231+
max_runs - 1 if max_runs is not None else None,
232+
**kwargs):
220233
if first_run:
221234
yield self._merge_stream_message(
222235
messages[-1], chunk)
@@ -265,8 +278,9 @@ def _stream_format_output_message(completion_chunk) -> Message:
265278
reasoning_content=reasoning_content,
266279
tool_calls=tool_calls,
267280
id=completion_chunk.id,
268-
prompt_tokens=completion_chunk.usage.prompt_tokens,
269-
completion_tokens=completion_chunk.usage.completion_tokens)
281+
prompt_tokens=getattr(completion_chunk.usage, 'prompt_tokens', 0),
282+
completion_tokens=getattr(completion_chunk.usage,
283+
'completion_tokens', 0))
270284

271285
@staticmethod
272286
def _format_output_message(completion) -> Message:
@@ -359,6 +373,7 @@ def _continue_generate(self,
359373
messages: List[Message],
360374
completion,
361375
tools: List[Tool] = None,
376+
max_runs: Optional[int] = None,
362377
**kwargs) -> Message:
363378
"""Recursively continues generating until the model finishes naturally.
364379
@@ -375,14 +390,17 @@ def _continue_generate(self,
375390
Message: A fully formed Message object containing the complete response.
376391
"""
377392
new_message = self._format_output_message(completion)
378-
if completion.choices[0].finish_reason in ['length', 'null']:
393+
if completion.choices[0].finish_reason in [
394+
'length', 'null'
395+
] and (max_runs is None or max_runs != 0):
379396
logger.info(
380397
f'finish_reason: {completion.choices[0].finish_reason}, continue generate.'
381398
)
382399
completion = self._call_llm_for_continue_gen(
383400
messages, new_message, tools, **kwargs)
384-
return self._continue_generate(messages, completion, tools,
385-
**kwargs)
401+
return self._continue_generate(
402+
messages, completion, tools,
403+
max_runs - 1 if max_runs is not None else None, **kwargs)
386404
elif messages[-1].to_dict().get('partial', False):
387405
self._merge_partial_message(messages, new_message)
388406
messages[-1].partial = False

tests/llm/test_openai.py

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -3,19 +3,22 @@
33
import os
44
import unittest
55

6-
from ms_agent.llm.openai_llm import OpenAI
7-
from ms_agent.llm.utils import Message, Tool, ToolCall
6+
from ms_agent.llm.openai_llm import MAX_CONTINUE_RUNS, OpenAI
7+
from ms_agent.llm.utils import Message, Tool
88
from omegaconf import DictConfig, OmegaConf
99

10+
from modelscope.utils.test_utils import test_level
11+
1012
API_CALL_MAX_TOKEN = 50
1113

1214

1315
class OpenaiLLM(unittest.TestCase):
1416
conf: DictConfig = OmegaConf.create({
1517
'llm': {
16-
'model': 'Qwen/Qwen3-235B-A22B',
17-
'openai_base_url': 'https://api-inference.modelscope.cn/v1',
18-
'openai_api_key': os.getenv('MODELSCOPE_API_KEY'),
18+
'model': 'qwen3-235b-a22b',
19+
'openai_base_url':
20+
'https://dashscope.aliyuncs.com/compatible-mode/v1',
21+
'openai_api_key': os.getenv('DASHSCOPE_API_KEY'),
1922
},
2023
'generation_config': {
2124
'stream': False,
@@ -68,62 +71,73 @@ class OpenaiLLM(unittest.TestCase):
6871
})
6972
]
7073

74+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
7175
def test_call_no_stream(self):
7276
llm = OpenAI(self.conf)
7377
res = llm.generate(messages=self.messages, tools=None)
7478
print(res)
7579
assert (res.content)
7680

81+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
7782
def test_call_stream(self):
7883
llm = OpenAI(self.conf)
7984
res = llm.generate(messages=self.messages, tools=None, stream=True)
8085
for chunk in res:
8186
print(chunk)
8287
assert (len(chunk.content))
8388

89+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
8490
def test_call_thinking(self):
8591
llm = OpenAI(self.conf)
8692
res = llm.generate(
8793
messages=self.messages,
8894
tools=None,
95+
max_continue_runs=1,
8996
stream=True,
9097
extra_body={'enable_thinking': True})
9198
for chunk in res:
9299
print(chunk)
93100
assert (chunk.reasoning_content)
94101

102+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
95103
def test_continue_run(self):
96104
llm = OpenAI(self.conf)
97105
res = llm.generate(messages=self.continue_messages, tools=None)
98106
print(res)
99107
assert (res.completion_tokens > 100)
100108

109+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
101110
def test_call_tool(self):
102111
llm = OpenAI(self.conf)
103112
res = llm.generate(messages=self.tool_messages, tools=self.tools)
104113
print(res)
105114
assert (len(res.tool_calls))
106115

116+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
107117
def test_call_apis_count(self):
108118
llm = OpenAI(self.conf)
109119
res = llm.generate(messages=self.messages, tools=None)
110120
print(res)
111121
assert res.api_calls == 1
112122

123+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
113124
def test_call_apis_count_stream(self):
114125
llm = OpenAI(self.conf)
115126
res = llm.generate(messages=self.messages, stream=True, tools=None)
116127
for chunk in res:
117128
print(chunk)
118129
assert chunk.api_calls == 1
119130

131+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
120132
def test_call_apis_count_continue(self):
121133
llm = OpenAI(self.conf)
122134
res = llm.generate(messages=self.continue_messages, tools=None)
123135
print(res)
124136
assert math.ceil(res.completion_tokens
125137
/ API_CALL_MAX_TOKEN) == res.api_calls
138+
assert res.api_calls <= MAX_CONTINUE_RUNS
126139

140+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
127141
def test_call_apis_count_continue_stream(self):
128142
llm = OpenAI(self.conf)
129143
res = llm.generate(
@@ -132,7 +146,9 @@ def test_call_apis_count_continue_stream(self):
132146
print(chunk)
133147
assert math.ceil(chunk.completion_tokens
134148
/ API_CALL_MAX_TOKEN) == chunk.api_calls
149+
assert chunk.api_calls <= MAX_CONTINUE_RUNS
135150

151+
@unittest.skipUnless(test_level() >= 0, 'skip test in current test level')
136152
def test_call_tool_stream(self):
137153
llm = OpenAI(self.conf)
138154
res = llm.generate(

0 commit comments

Comments
 (0)