1212
1313logger = get_logger ()
1414
15+ MAX_CONTINUE_RUNS = 3
16+
1517
1618class 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
0 commit comments