| 15 | |
| 16 | |
| 17 | class OpenAIAPI(BaseLLM): |
| 18 | def __init__(self, config: LLMConfig) -> None: |
| 19 | super().__init__(config) |
| 20 | self.pool = SpinPool(config["client_args"]) |
| 21 | |
| 22 | def get_client_args(self): |
| 23 | return self.pool |
| 24 | |
| 25 | def generate( |
| 26 | self, |
| 27 | messages: List[Union[Message, Dict]], |
| 28 | tools: Optional[List[Union[Tool, Dict]]] = None, |
| 29 | ) -> Tuple[Message, DebugInfo]: |
| 30 | if len(messages) > 0 and isinstance(messages[0], Message): |
| 31 | kwargs = dict( |
| 32 | messages=[message2dict(msg) for msg in messages], |
| 33 | ) |
| 34 | elif isinstance(messages[0], dict): |
| 35 | kwargs = dict( |
| 36 | messages=messages, |
| 37 | ) |
| 38 | kwargs.update(self.config["model_args"]) |
| 39 | if tools is not None: |
| 40 | if len(tools) > 0 and isinstance(tools[0], Tool): |
| 41 | kwargs["tools"] = [tool2dict(tool) for tool in tools] |
| 42 | kwargs["tool_choice"] = "auto" |
| 43 | else: |
| 44 | kwargs["tools"] = tools |
| 45 | kwargs["tool_choice"] = "auto" |
| 46 | |
| 47 | new_message, debug_info = self._post_request(kwargs) |
| 48 | |
| 49 | new_message = Message(**new_message.model_dump()) |
| 50 | return new_message, debug_info |
| 51 | |
| 52 | @retry(stop=stop_after_attempt(10), wait=wait_random(min=5, max=10)) |
| 53 | def _post_request(self, kwargs: Dict[str, Any]) -> Tuple[OpenAIMessage, DebugInfo]: |
| 54 | try: |
| 55 | with self.get_client_args() as client_args: |
| 56 | logger.debug(f"getting client args: {client_args}") |
| 57 | client = openai.OpenAI( |
| 58 | api_key=client_args.get("api_key", None), |
| 59 | base_url=client_args.get("base_url", None), |
| 60 | ) |
| 61 | completion = client.chat.completions.create(**kwargs) |
| 62 | logger.debug(f"{completion}") |
| 63 | except openai.RateLimitError as e: |
| 64 | logger.warning(f"OpenAI rate limit error: {e}") |
| 65 | raise e |
| 66 | except openai.APIError as e: |
| 67 | logger.error(f"OpenAI API error: {e}") |
| 68 | raise e |
| 69 | except Exception as e: |
| 70 | logger.error(f"Unexpected error: {e}, type{e}, args{e.args}") |
| 71 | logger.error(f"Traceback: {traceback.format_exc()}") |
| 72 | raise e |
| 73 | |
| 74 | new_message = completion.choices[0].message |