AITECHPRODUCTS/githubtest
0
1import linecache2import re3from typing import Dict, List, Optional4 5import openai6 7 8class ChatCompletion:9 def __init__(self, model: str = 'gpt-3.5-turbo',10 api_key: Optional[str] = None, api_key_path: str = './openai_api_key'):11 if api_key is None:12 openai.api_key = api_key13 api_key = linecache.getline(api_key_path, 2).strip('\n')14 if len(api_key) == 0:15 raise EnvironmentError16 openai.api_key = api_key17 18 self.model = model19 self.system_messages = []20 self.user_messages = []21 22 def chat(self, msg: str, setting: Optional[str] = None, model: Optional[str] = None) -> str:23 if self._context_length() > 2048:24 self.reset()25 if setting is not None:26 if setting not in self.system_messages:27 self.system_messages.append(setting)28 if not self.user_messages or msg != self.user_messages[-1]:29 self.user_messages.append(msg)30 31 return self._run(model)32 33 def retry(self, model: Optional[str] = None) -> str:34 return self._run(model)35 36 def reset(self):37 self.system_messages.clear()38 self.user_messages.clear()39 40 def _make_message(self) -> List[Dict]:41 sys_messages = [{'role': 'system', 'content': msg} for msg in self.system_messages]42 user_messages = [{'role': 'user', 'content': msg} for msg in self.user_messages]43 return sys_messages + user_messages44 45 def _context_length(self) -> int:46 return len(''.join(self.system_messages)) + len(''.join(self.user_messages))47 48 def _run(self, model: Optional[str] = None) -> str:49 if model is None:50 model = self.model51 try:52 response = openai.ChatCompletion.create(model=model, messages=self._make_message())53 ans = response['choices'][0]['message']['content']54 ans = re.sub(r'^\n+', '', ans)55 except openai.error.OpenAIError as e:56 ans = e57 except Exception as e:58 print(e)59 return ans60 61 def __call__(self, msg: str, setting: Optional[str] = None, model: Optional[str] = None) -> str:62 return self.chat(msg, setting, model)63 