| 156 | |
| 157 | |
| 158 | class ChatOpenAI(BaseOpenAIEngine): |
| 159 | def __init__( |
| 160 | self, |
| 161 | model_string: str = "gpt-3.5-turbo-0613", |
| 162 | system_prompt: str = BaseOpenAIEngine.DEFAULT_SYSTEM_PROMPT, |
| 163 | is_multimodal: bool = False, |
| 164 | base_url: str = None, |
| 165 | **kwargs, |
| 166 | ): |
| 167 | """ |
| 168 | :param model_string: |
| 169 | :param system_prompt: |
| 170 | :param base_url: Used to support Ollama |
| 171 | """ |
| 172 | root = platformdirs.user_cache_dir("textgrad") |
| 173 | cache_path = os.path.join(root, f"cache_openai_{model_string}.db") |
| 174 | |
| 175 | super().__init__(cache_path, system_prompt, model_string, is_multimodal) |
| 176 | |
| 177 | self.base_url = base_url |
| 178 | |
| 179 | if not base_url: |
| 180 | if os.getenv("OPENAI_API_KEY") is None: |
| 181 | raise ValueError( |
| 182 | "Please set the OPENAI_API_KEY environment variable if you'd like to use OpenAI models." |
| 183 | ) |
| 184 | |
| 185 | # Check if OPENAI_API_BASE is set in environment variables |
| 186 | api_base = os.getenv("OPENAI_API_BASE") |
| 187 | if api_base: |
| 188 | self.client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"), base_url=api_base) |
| 189 | else: |
| 190 | self.client = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) |
| 191 | elif base_url and base_url == OLLAMA_BASE_URL: |
| 192 | self.client = OpenAI(base_url=base_url, api_key="ollama") |
| 193 | else: |
| 194 | raise ValueError( |
| 195 | "Invalid base URL provided. Please use the default OLLAMA base URL or None." |
| 196 | ) |
| 197 | |
| 198 | |
| 199 | class AzureChatOpenAI(BaseOpenAIEngine): |
no outgoing calls
no test coverage detected