openai/openai-python

Public

mirrored from https://github.com/openai/openai-pythonAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
15c6aebc6cda59e8d884e1fbc63fd1f6eddbc196

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

openai/api_requestor.py

695lines · modecode

1import asyncio
2import json
3import platform
4import sys
5import threading
6import warnings
7from contextlib import asynccontextmanager
8from json import JSONDecodeError
9from typing import (
10 AsyncGenerator,
11 AsyncIterator,
12 Dict,
13 Iterator,
14 Optional,
15 Tuple,
16 Union,
17 overload,
18)
19from urllib.parse import urlencode, urlsplit, urlunsplit
20
21import aiohttp
22import requests
23
24if sys.version_info >= (3, 8):
25 from typing import Literal
26else:
27 from typing_extensions import Literal
28
29import openai
30from openai import error, util, version
31from openai.openai_response import OpenAIResponse
32from openai.util import ApiType
33
34TIMEOUT_SECS = 600
35MAX_CONNECTION_RETRIES = 2
36
37# Has one attribute per thread, 'session'.
38_thread_context = threading.local()
39
40
41def _build_api_url(url, query):
42 scheme, netloc, path, base_query, fragment = urlsplit(url)
43
44 if base_query:
45 query = "%s&%s" % (base_query, query)
46
47 return urlunsplit((scheme, netloc, path, query, fragment))
48
49
50def _requests_proxies_arg(proxy) -> Optional[Dict[str, str]]:
51 """Returns a value suitable for the 'proxies' argument to 'requests.request."""
52 if proxy is None:
53 return None
54 elif isinstance(proxy, str):
55 return {"http": proxy, "https": proxy}
56 elif isinstance(proxy, dict):
57 return proxy.copy()
58 else:
59 raise ValueError(
60 "'openai.proxy' must be specified as either a string URL or a dict with string URL under the https and/or http keys."
61 )
62
63
64def _aiohttp_proxies_arg(proxy) -> Optional[str]:
65 """Returns a value suitable for the 'proxies' argument to 'aiohttp.ClientSession.request."""
66 if proxy is None:
67 return None
68 elif isinstance(proxy, str):
69 return proxy
70 elif isinstance(proxy, dict):
71 return proxy["https"] if "https" in proxy else proxy["http"]
72 else:
73 raise ValueError(
74 "'openai.proxy' must be specified as either a string URL or a dict with string URL under the https and/or http keys."
75 )
76
77
78def _make_session() -> requests.Session:
79 if not openai.verify_ssl_certs:
80 warnings.warn("verify_ssl_certs is ignored; openai always verifies.")
81 s = requests.Session()
82 proxies = _requests_proxies_arg(openai.proxy)
83 if proxies:
84 s.proxies = proxies
85 s.mount(
86 "https://",
87 requests.adapters.HTTPAdapter(max_retries=MAX_CONNECTION_RETRIES),
88 )
89 return s
90
91
92def parse_stream_helper(line: bytes) -> Optional[str]:
93 if line:
94 if line.strip() == b"data: [DONE]":
95 # return here will cause GeneratorExit exception in urllib3
96 # and it will close http connection with TCP Reset
97 return None
98 if line.startswith(b"data: "):
99 line = line[len(b"data: "):]
100 return line.decode("utf-8")
101 else:
102 return None
103 return None
104
105
106def parse_stream(rbody: Iterator[bytes]) -> Iterator[str]:
107 for line in rbody:
108 _line = parse_stream_helper(line)
109 if _line is not None:
110 yield _line
111
112
113async def parse_stream_async(rbody: aiohttp.StreamReader):
114 async for line in rbody:
115 _line = parse_stream_helper(line)
116 if _line is not None:
117 yield _line
118
119
120class APIRequestor:
121 def __init__(
122 self,
123 key=None,
124 api_base=None,
125 api_type=None,
126 api_version=None,
127 organization=None,
128 ):
129 self.api_base = api_base or openai.api_base
130 self.api_key = key or util.default_api_key()
131 self.api_type = (
132 ApiType.from_str(api_type)
133 if api_type
134 else ApiType.from_str(openai.api_type)
135 )
136 self.api_version = api_version or openai.api_version
137 self.organization = organization or openai.organization
138
139 @classmethod
140 def format_app_info(cls, info):
141 str = info["name"]
142 if info["version"]:
143 str += "/%s" % (info["version"],)
144 if info["url"]:
145 str += " (%s)" % (info["url"],)
146 return str
147
148 @overload
149 def request(
150 self,
151 method,
152 url,
153 params,
154 headers,
155 files,
156 stream: Literal[True],
157 request_id: Optional[str] = ...,
158 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
159 ) -> Tuple[Iterator[OpenAIResponse], bool, str]:
160 pass
161
162 @overload
163 def request(
164 self,
165 method,
166 url,
167 params=...,
168 headers=...,
169 files=...,
170 *,
171 stream: Literal[True],
172 request_id: Optional[str] = ...,
173 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
174 ) -> Tuple[Iterator[OpenAIResponse], bool, str]:
175 pass
176
177 @overload
178 def request(
179 self,
180 method,
181 url,
182 params=...,
183 headers=...,
184 files=...,
185 stream: Literal[False] = ...,
186 request_id: Optional[str] = ...,
187 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
188 ) -> Tuple[OpenAIResponse, bool, str]:
189 pass
190
191 @overload
192 def request(
193 self,
194 method,
195 url,
196 params=...,
197 headers=...,
198 files=...,
199 stream: bool = ...,
200 request_id: Optional[str] = ...,
201 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
202 ) -> Tuple[Union[OpenAIResponse, Iterator[OpenAIResponse]], bool, str]:
203 pass
204
205 def request(
206 self,
207 method,
208 url,
209 params=None,
210 headers=None,
211 files=None,
212 stream: bool = False,
213 request_id: Optional[str] = None,
214 request_timeout: Optional[Union[float, Tuple[float, float]]] = None,
215 ) -> Tuple[Union[OpenAIResponse, Iterator[OpenAIResponse]], bool, str]:
216 result = self.request_raw(
217 method.lower(),
218 url,
219 params=params,
220 supplied_headers=headers,
221 files=files,
222 stream=stream,
223 request_id=request_id,
224 request_timeout=request_timeout,
225 )
226 resp, got_stream = self._interpret_response(result, stream)
227 return resp, got_stream, self.api_key
228
229 @overload
230 async def arequest(
231 self,
232 method,
233 url,
234 params,
235 headers,
236 files,
237 stream: Literal[True],
238 request_id: Optional[str] = ...,
239 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
240 ) -> Tuple[AsyncGenerator[OpenAIResponse, None], bool, str]:
241 pass
242
243 @overload
244 async def arequest(
245 self,
246 method,
247 url,
248 params=...,
249 headers=...,
250 files=...,
251 *,
252 stream: Literal[True],
253 request_id: Optional[str] = ...,
254 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
255 ) -> Tuple[AsyncGenerator[OpenAIResponse, None], bool, str]:
256 pass
257
258 @overload
259 async def arequest(
260 self,
261 method,
262 url,
263 params=...,
264 headers=...,
265 files=...,
266 stream: Literal[False] = ...,
267 request_id: Optional[str] = ...,
268 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
269 ) -> Tuple[OpenAIResponse, bool, str]:
270 pass
271
272 @overload
273 async def arequest(
274 self,
275 method,
276 url,
277 params=...,
278 headers=...,
279 files=...,
280 stream: bool = ...,
281 request_id: Optional[str] = ...,
282 request_timeout: Optional[Union[float, Tuple[float, float]]] = ...,
283 ) -> Tuple[Union[OpenAIResponse, AsyncGenerator[OpenAIResponse, None]], bool, str]:
284 pass
285
286 async def arequest(
287 self,
288 method,
289 url,
290 params=None,
291 headers=None,
292 files=None,
293 stream: bool = False,
294 request_id: Optional[str] = None,
295 request_timeout: Optional[Union[float, Tuple[float, float]]] = None,
296 ) -> Tuple[Union[OpenAIResponse, AsyncGenerator[OpenAIResponse, None]], bool, str]:
297 ctx = aiohttp_session()
298 session = await ctx.__aenter__()
299 try:
300 result = await self.arequest_raw(
301 method.lower(),
302 url,
303 session,
304 params=params,
305 supplied_headers=headers,
306 files=files,
307 request_id=request_id,
308 request_timeout=request_timeout,
309 )
310 resp, got_stream = await self._interpret_async_response(result, stream)
311 except Exception:
312 await ctx.__aexit__(None, None, None)
313 raise
314 if got_stream:
315
316 async def wrap_resp():
317 assert isinstance(resp, AsyncGenerator)
318 try:
319 async for r in resp:
320 yield r
321 finally:
322 await ctx.__aexit__(None, None, None)
323
324 return wrap_resp(), got_stream, self.api_key
325 else:
326 await ctx.__aexit__(None, None, None)
327 return resp, got_stream, self.api_key
328
329 def handle_error_response(self, rbody, rcode, resp, rheaders, stream_error=False):
330 try:
331 error_data = resp["error"]
332 except (KeyError, TypeError):
333 raise error.APIError(
334 "Invalid response object from API: %r (HTTP response code "
335 "was %d)" % (rbody, rcode),
336 rbody,
337 rcode,
338 resp,
339 )
340
341 if "internal_message" in error_data:
342 error_data["message"] += "\n\n" + error_data["internal_message"]
343
344 util.log_info(
345 "OpenAI API error received",
346 error_code=error_data.get("code"),
347 error_type=error_data.get("type"),
348 error_message=error_data.get("message"),
349 error_param=error_data.get("param"),
350 stream_error=stream_error,
351 )
352
353 # Rate limits were previously coded as 400's with code 'rate_limit'
354 if rcode == 429:
355 return error.RateLimitError(
356 error_data.get("message"), rbody, rcode, resp, rheaders
357 )
358 elif rcode in [400, 404, 415]:
359 return error.InvalidRequestError(
360 error_data.get("message"),
361 error_data.get("param"),
362 error_data.get("code"),
363 rbody,
364 rcode,
365 resp,
366 rheaders,
367 )
368 elif rcode == 401:
369 return error.AuthenticationError(
370 error_data.get("message"), rbody, rcode, resp, rheaders
371 )
372 elif rcode == 403:
373 return error.PermissionError(
374 error_data.get("message"), rbody, rcode, resp, rheaders
375 )
376 elif rcode == 409:
377 return error.TryAgain(
378 error_data.get("message"), rbody, rcode, resp, rheaders
379 )
380 elif stream_error:
381 # TODO: we will soon attach status codes to stream errors
382 parts = [error_data.get("message"), "(Error occurred while streaming.)"]
383 message = " ".join([p for p in parts if p is not None])
384 return error.APIError(message, rbody, rcode, resp, rheaders)
385 else:
386 return error.APIError(
387 f"{error_data.get('message')} {rbody} {rcode} {resp} {rheaders}",
388 rbody,
389 rcode,
390 resp,
391 rheaders,
392 )
393
394 def request_headers(
395 self, method: str, extra, request_id: Optional[str]
396 ) -> Dict[str, str]:
397 user_agent = "OpenAI/v1 PythonBindings/%s" % (version.VERSION,)
398 if openai.app_info:
399 user_agent += " " + self.format_app_info(openai.app_info)
400
401 uname_without_node = " ".join(
402 v for k, v in platform.uname()._asdict().items() if k != "node"
403 )
404 ua = {
405 "bindings_version": version.VERSION,
406 "httplib": "requests",
407 "lang": "python",
408 "lang_version": platform.python_version(),
409 "platform": platform.platform(),
410 "publisher": "openai",
411 "uname": uname_without_node,
412 }
413 if openai.app_info:
414 ua["application"] = openai.app_info
415
416 headers = {
417 "X-OpenAI-Client-User-Agent": json.dumps(ua),
418 "User-Agent": user_agent,
419 }
420
421 headers.update(util.api_key_to_header(self.api_type, self.api_key))
422
423 if self.organization:
424 headers["OpenAI-Organization"] = self.organization
425
426 if self.api_version is not None and self.api_type == ApiType.OPEN_AI:
427 headers["OpenAI-Version"] = self.api_version
428 if request_id is not None:
429 headers["X-Request-Id"] = request_id
430 if openai.debug:
431 headers["OpenAI-Debug"] = "true"
432 headers.update(extra)
433
434 return headers
435
436 def _validate_headers(
437 self, supplied_headers: Optional[Dict[str, str]]
438 ) -> Dict[str, str]:
439 headers: Dict[str, str] = {}
440 if supplied_headers is None:
441 return headers
442
443 if not isinstance(supplied_headers, dict):
444 raise TypeError("Headers must be a dictionary")
445
446 for k, v in supplied_headers.items():
447 if not isinstance(k, str):
448 raise TypeError("Header keys must be strings")
449 if not isinstance(v, str):
450 raise TypeError("Header values must be strings")
451 headers[k] = v
452
453 # NOTE: It is possible to do more validation of the headers, but a request could always
454 # be made to the API manually with invalid headers, so we need to handle them server side.
455
456 return headers
457
458 def _prepare_request_raw(
459 self,
460 url,
461 supplied_headers,
462 method,
463 params,
464 files,
465 request_id: Optional[str],
466 ) -> Tuple[str, Dict[str, str], Optional[bytes]]:
467 abs_url = "%s%s" % (self.api_base, url)
468 headers = self._validate_headers(supplied_headers)
469
470 data = None
471 if method == "get" or method == "delete":
472 if params:
473 encoded_params = urlencode(
474 [(k, v) for k, v in params.items() if v is not None]
475 )
476 abs_url = _build_api_url(abs_url, encoded_params)
477 elif method in {"post", "put"}:
478 if params and files:
479 data = params
480 if params and not files:
481 data = json.dumps(params).encode()
482 headers["Content-Type"] = "application/json"
483 else:
484 raise error.APIConnectionError(
485 "Unrecognized HTTP method %r. This may indicate a bug in the "
486 "OpenAI bindings. Please contact support@openai.com for "
487 "assistance." % (method,)
488 )
489
490 headers = self.request_headers(method, headers, request_id)
491
492 util.log_debug("Request to OpenAI API", method=method, path=abs_url)
493 util.log_debug("Post details", data=data, api_version=self.api_version)
494
495 return abs_url, headers, data
496
497 def request_raw(
498 self,
499 method,
500 url,
501 *,
502 params=None,
503 supplied_headers: Optional[Dict[str, str]] = None,
504 files=None,
505 stream: bool = False,
506 request_id: Optional[str] = None,
507 request_timeout: Optional[Union[float, Tuple[float, float]]] = None,
508 ) -> requests.Response:
509 abs_url, headers, data = self._prepare_request_raw(
510 url, supplied_headers, method, params, files, request_id
511 )
512
513 if not hasattr(_thread_context, "session"):
514 _thread_context.session = _make_session()
515 try:
516 result = _thread_context.session.request(
517 method,
518 abs_url,
519 headers=headers,
520 data=data,
521 files=files,
522 stream=stream,
523 timeout=request_timeout if request_timeout else TIMEOUT_SECS,
524 )
525 except requests.exceptions.Timeout as e:
526 raise error.Timeout("Request timed out: {}".format(e)) from e
527 except requests.exceptions.RequestException as e:
528 raise error.APIConnectionError(
529 "Error communicating with OpenAI: {}".format(e)
530 ) from e
531 util.log_debug(
532 "OpenAI API response",
533 path=abs_url,
534 response_code=result.status_code,
535 processing_ms=result.headers.get("OpenAI-Processing-Ms"),
536 request_id=result.headers.get("X-Request-Id"),
537 )
538 # Don't read the whole stream for debug logging unless necessary.
539 if openai.log == "debug":
540 util.log_debug(
541 "API response body", body=result.content, headers=result.headers
542 )
543 return result
544
545 async def arequest_raw(
546 self,
547 method,
548 url,
549 session,
550 *,
551 params=None,
552 supplied_headers: Optional[Dict[str, str]] = None,
553 files=None,
554 request_id: Optional[str] = None,
555 request_timeout: Optional[Union[float, Tuple[float, float]]] = None,
556 ) -> aiohttp.ClientResponse:
557 abs_url, headers, data = self._prepare_request_raw(
558 url, supplied_headers, method, params, files, request_id
559 )
560
561 if isinstance(request_timeout, tuple):
562 timeout = aiohttp.ClientTimeout(
563 connect=request_timeout[0],
564 total=request_timeout[1],
565 )
566 else:
567 timeout = aiohttp.ClientTimeout(
568 total=request_timeout if request_timeout else TIMEOUT_SECS
569 )
570
571 if files:
572 # TODO: Use `aiohttp.MultipartWriter` to create the multipart form data here.
573 # For now we use the private `requests` method that is known to have worked so far.
574 data, content_type = requests.models.RequestEncodingMixin._encode_files( # type: ignore
575 files, data
576 )
577 headers["Content-Type"] = content_type
578 request_kwargs = {
579 "method": method,
580 "url": abs_url,
581 "headers": headers,
582 "data": data,
583 "proxy": _aiohttp_proxies_arg(openai.proxy),
584 "timeout": timeout,
585 }
586 try:
587 result = await session.request(**request_kwargs)
588 util.log_info(
589 "OpenAI API response",
590 path=abs_url,
591 response_code=result.status,
592 processing_ms=result.headers.get("OpenAI-Processing-Ms"),
593 request_id=result.headers.get("X-Request-Id"),
594 )
595 # Don't read the whole stream for debug logging unless necessary.
596 if openai.log == "debug":
597 util.log_debug(
598 "API response body", body=result.content, headers=result.headers
599 )
600 return result
601 except (aiohttp.ServerTimeoutError, asyncio.TimeoutError) as e:
602 raise error.Timeout("Request timed out") from e
603 except aiohttp.ClientError as e:
604 raise error.APIConnectionError("Error communicating with OpenAI") from e
605
606 def _interpret_response(
607 self, result: requests.Response, stream: bool
608 ) -> Tuple[Union[OpenAIResponse, Iterator[OpenAIResponse]], bool]:
609 """Returns the response(s) and a bool indicating whether it is a stream."""
610 if stream and "text/event-stream" in result.headers.get("Content-Type", ""):
611 return (
612 self._interpret_response_line(
613 line, result.status_code, result.headers, stream=True
614 )
615 for line in parse_stream(result.iter_lines())
616 ), True
617 else:
618 return (
619 self._interpret_response_line(
620 result.content.decode("utf-8"),
621 result.status_code,
622 result.headers,
623 stream=False,
624 ),
625 False,
626 )
627
628 async def _interpret_async_response(
629 self, result: aiohttp.ClientResponse, stream: bool
630 ) -> Tuple[Union[OpenAIResponse, AsyncGenerator[OpenAIResponse, None]], bool]:
631 """Returns the response(s) and a bool indicating whether it is a stream."""
632 if stream and "text/event-stream" in result.headers.get("Content-Type", ""):
633 return (
634 self._interpret_response_line(
635 line, result.status, result.headers, stream=True
636 )
637 async for line in parse_stream_async(result.content)
638 ), True
639 else:
640 try:
641 await result.read()
642 except aiohttp.ClientError as e:
643 util.log_warn(e, body=result.content)
644 return (
645 self._interpret_response_line(
646 (await result.read()).decode("utf-8"),
647 result.status,
648 result.headers,
649 stream=False,
650 ),
651 False,
652 )
653
654 def _interpret_response_line(
655 self, rbody: str, rcode: int, rheaders, stream: bool
656 ) -> OpenAIResponse:
657 # HTTP 204 response code does not have any content in the body.
658 if rcode == 204:
659 return OpenAIResponse(None, rheaders)
660
661 if rcode == 503:
662 raise error.ServiceUnavailableError(
663 "The server is overloaded or not ready yet.",
664 rbody,
665 rcode,
666 headers=rheaders,
667 )
668 try:
669 if 'text/plain' in rheaders.get('Content-Type'):
670 data = rbody
671 else:
672 data = json.loads(rbody)
673 except (JSONDecodeError, UnicodeDecodeError) as e:
674 raise error.APIError(
675 f"HTTP code {rcode} from API ({rbody})", rbody, rcode, headers=rheaders
676 ) from e
677 resp = OpenAIResponse(data, rheaders)
678 # In the future, we might add a "status" parameter to errors
679 # to better handle the "error while streaming" case.
680 stream_error = stream and "error" in resp.data
681 if stream_error or not 200 <= rcode < 300:
682 raise self.handle_error_response(
683 rbody, rcode, resp.data, rheaders, stream_error=stream_error
684 )
685 return resp
686
687
688@asynccontextmanager
689async def aiohttp_session() -> AsyncIterator[aiohttp.ClientSession]:
690 user_set_session = openai.aiosession.get()
691 if user_set_session:
692 yield user_set_session
693 else:
694 async with aiohttp.ClientSession() as session:
695 yield session
696