openai/openai-python

Public

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

CodeCommitsIssuesPull requestsActionsInsightsSecurity
v1.3.5

Branches

Tags

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

Clone

HTTPS

Download ZIP

tests/test_client.py

1164lines · modeblame

08b8179aDavid Schnurr2 years ago1# File generated from our OpenAPI spec by Stainless.
2
3from __future__ import annotations
4
5import os
6import json
7import asyncio
8import inspect
9from typing import Any, Dict, Union, cast
10from unittest import mock
11
12import httpx
13import pytest
14from respx import MockRouter
15from pydantic import ValidationError
16
17from openai import OpenAI, AsyncOpenAI, APIResponseValidationError
18from openai._client import OpenAI, AsyncOpenAI
19from openai._models import BaseModel, FinalRequestOptions
20from openai._streaming import Stream, AsyncStream
21from openai._exceptions import APIResponseValidationError
22from openai._base_client import (
23DEFAULT_TIMEOUT,
24HTTPX_DEFAULT_TIMEOUT,
25BaseClient,
26make_request_options,
27)
28
0733934fStainless Bot2 years ago29from .utils import update_env
30
08b8179aDavid Schnurr2 years ago31base_url = os.environ.get("TEST_API_BASE_URL", "http://127.0.0.1:4010")
32api_key = "My API Key"
33
34
35def _get_params(client: BaseClient[Any, Any]) -> dict[str, str]:
36request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
37url = httpx.URL(request.url)
38return dict(url.params)
39
40
41class TestOpenAI:
42client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
43
44@pytest.mark.respx(base_url=base_url)
45def test_raw_response(self, respx_mock: MockRouter) -> None:
c5975bd0Stainless Bot2 years ago46respx_mock.post("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
08b8179aDavid Schnurr2 years ago47
48response = self.client.post("/foo", cast_to=httpx.Response)
49assert response.status_code == 200
50assert isinstance(response, httpx.Response)
c5975bd0Stainless Bot2 years ago51assert response.json() == {"foo": "bar"}
08b8179aDavid Schnurr2 years ago52
53@pytest.mark.respx(base_url=base_url)
54def test_raw_response_for_binary(self, respx_mock: MockRouter) -> None:
55respx_mock.post("/foo").mock(
56return_value=httpx.Response(200, headers={"Content-Type": "application/binary"}, content='{"foo": "bar"}')
57)
58
59response = self.client.post("/foo", cast_to=httpx.Response)
60assert response.status_code == 200
61assert isinstance(response, httpx.Response)
c5975bd0Stainless Bot2 years ago62assert response.json() == {"foo": "bar"}
08b8179aDavid Schnurr2 years ago63
64def test_copy(self) -> None:
65copied = self.client.copy()
66assert id(copied) != id(self.client)
67
68copied = self.client.copy(api_key="another My API Key")
69assert copied.api_key == "another My API Key"
70assert self.client.api_key == "My API Key"
71
72def test_copy_default_options(self) -> None:
73# options that have a default are overridden correctly
74copied = self.client.copy(max_retries=7)
75assert copied.max_retries == 7
76assert self.client.max_retries == 2
77
78copied2 = copied.copy(max_retries=6)
79assert copied2.max_retries == 6
80assert copied.max_retries == 7
81
82# timeout
83assert isinstance(self.client.timeout, httpx.Timeout)
84copied = self.client.copy(timeout=None)
85assert copied.timeout is None
86assert isinstance(self.client.timeout, httpx.Timeout)
87
88def test_copy_default_headers(self) -> None:
89client = OpenAI(
90base_url=base_url, api_key=api_key, _strict_response_validation=True, default_headers={"X-Foo": "bar"}
91)
92assert client.default_headers["X-Foo"] == "bar"
93
94# does not override the already given value when not specified
95copied = client.copy()
96assert copied.default_headers["X-Foo"] == "bar"
97
98# merges already given headers
99copied = client.copy(default_headers={"X-Bar": "stainless"})
100assert copied.default_headers["X-Foo"] == "bar"
101assert copied.default_headers["X-Bar"] == "stainless"
102
103# uses new values for any already given headers
104copied = client.copy(default_headers={"X-Foo": "stainless"})
105assert copied.default_headers["X-Foo"] == "stainless"
106
107# set_default_headers
108
109# completely overrides already set values
110copied = client.copy(set_default_headers={})
111assert copied.default_headers.get("X-Foo") is None
112
113copied = client.copy(set_default_headers={"X-Bar": "Robert"})
114assert copied.default_headers["X-Bar"] == "Robert"
115
116with pytest.raises(
117ValueError,
118match="`default_headers` and `set_default_headers` arguments are mutually exclusive",
119):
120client.copy(set_default_headers={}, default_headers={"X-Foo": "Bar"})
121
122def test_copy_default_query(self) -> None:
123client = OpenAI(
124base_url=base_url, api_key=api_key, _strict_response_validation=True, default_query={"foo": "bar"}
125)
126assert _get_params(client)["foo"] == "bar"
127
128# does not override the already given value when not specified
129copied = client.copy()
130assert _get_params(copied)["foo"] == "bar"
131
132# merges already given params
133copied = client.copy(default_query={"bar": "stainless"})
134params = _get_params(copied)
135assert params["foo"] == "bar"
136assert params["bar"] == "stainless"
137
138# uses new values for any already given headers
139copied = client.copy(default_query={"foo": "stainless"})
140assert _get_params(copied)["foo"] == "stainless"
141
142# set_default_query
143
144# completely overrides already set values
145copied = client.copy(set_default_query={})
146assert _get_params(copied) == {}
147
148copied = client.copy(set_default_query={"bar": "Robert"})
149assert _get_params(copied)["bar"] == "Robert"
150
151with pytest.raises(
152ValueError,
153# TODO: update
154match="`default_query` and `set_default_query` arguments are mutually exclusive",
155):
156client.copy(set_default_query={}, default_query={"foo": "Bar"})
157
158def test_copy_signature(self) -> None:
159# ensure the same parameters that can be passed to the client are defined in the `.copy()` method
160init_signature = inspect.signature(
161# mypy doesn't like that we access the `__init__` property.
162self.client.__init__, # type: ignore[misc]
163)
164copy_signature = inspect.signature(self.client.copy)
165exclude_params = {"transport", "proxies", "_strict_response_validation"}
166
167for name in init_signature.parameters.keys():
168if name in exclude_params:
169continue
170
171copy_param = copy_signature.parameters.get(name)
172assert copy_param is not None, f"copy() signature is missing the {name} param"
173
174def test_request_timeout(self) -> None:
175request = self.client._build_request(FinalRequestOptions(method="get", url="/foo"))
176timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
177assert timeout == DEFAULT_TIMEOUT
178
179request = self.client._build_request(
180FinalRequestOptions(method="get", url="/foo", timeout=httpx.Timeout(100.0))
181)
182timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
183assert timeout == httpx.Timeout(100.0)
184
185def test_client_timeout_option(self) -> None:
186client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True, timeout=httpx.Timeout(0))
187
188request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
189timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
190assert timeout == httpx.Timeout(0)
191
192def test_http_client_timeout_option(self) -> None:
193# custom timeout given to the httpx client should be used
194with httpx.Client(timeout=None) as http_client:
195client = OpenAI(
196base_url=base_url, api_key=api_key, _strict_response_validation=True, http_client=http_client
197)
198
199request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
200timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
201assert timeout == httpx.Timeout(None)
202
203# no timeout given to the httpx client should not use the httpx default
204with httpx.Client() as http_client:
205client = OpenAI(
206base_url=base_url, api_key=api_key, _strict_response_validation=True, http_client=http_client
207)
208
209request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
210timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
211assert timeout == DEFAULT_TIMEOUT
212
213# explicitly passing the default timeout currently results in it being ignored
214with httpx.Client(timeout=HTTPX_DEFAULT_TIMEOUT) as http_client:
215client = OpenAI(
216base_url=base_url, api_key=api_key, _strict_response_validation=True, http_client=http_client
217)
218
219request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
220timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
221assert timeout == DEFAULT_TIMEOUT # our default
222
223def test_default_headers_option(self) -> None:
224client = OpenAI(
225base_url=base_url, api_key=api_key, _strict_response_validation=True, default_headers={"X-Foo": "bar"}
226)
227request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
228assert request.headers.get("x-foo") == "bar"
229assert request.headers.get("x-stainless-lang") == "python"
230
231client2 = OpenAI(
232base_url=base_url,
233api_key=api_key,
234_strict_response_validation=True,
235default_headers={
236"X-Foo": "stainless",
237"X-Stainless-Lang": "my-overriding-header",
238},
239)
240request = client2._build_request(FinalRequestOptions(method="get", url="/foo"))
241assert request.headers.get("x-foo") == "stainless"
242assert request.headers.get("x-stainless-lang") == "my-overriding-header"
243
244def test_validate_headers(self) -> None:
245client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
246request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
247assert request.headers.get("Authorization") == f"Bearer {api_key}"
248
249with pytest.raises(Exception):
250client2 = OpenAI(base_url=base_url, api_key=None, _strict_response_validation=True)
251_ = client2
252
253def test_default_query_option(self) -> None:
254client = OpenAI(
255base_url=base_url, api_key=api_key, _strict_response_validation=True, default_query={"query_param": "bar"}
256)
257request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
258url = httpx.URL(request.url)
259assert dict(url.params) == {"query_param": "bar"}
260
261request = client._build_request(
262FinalRequestOptions(
263method="get",
264url="/foo",
265params={"foo": "baz", "query_param": "overriden"},
266)
267)
268url = httpx.URL(request.url)
269assert dict(url.params) == {"foo": "baz", "query_param": "overriden"}
270
271def test_request_extra_json(self) -> None:
272request = self.client._build_request(
273FinalRequestOptions(
274method="post",
275url="/foo",
276json_data={"foo": "bar"},
277extra_json={"baz": False},
278),
279)
280data = json.loads(request.content.decode("utf-8"))
281assert data == {"foo": "bar", "baz": False}
282
283request = self.client._build_request(
284FinalRequestOptions(
285method="post",
286url="/foo",
287extra_json={"baz": False},
288),
289)
290data = json.loads(request.content.decode("utf-8"))
291assert data == {"baz": False}
292
293# `extra_json` takes priority over `json_data` when keys clash
294request = self.client._build_request(
295FinalRequestOptions(
296method="post",
297url="/foo",
298json_data={"foo": "bar", "baz": True},
299extra_json={"baz": None},
300),
301)
302data = json.loads(request.content.decode("utf-8"))
303assert data == {"foo": "bar", "baz": None}
304
305def test_request_extra_headers(self) -> None:
306request = self.client._build_request(
307FinalRequestOptions(
308method="post",
309url="/foo",
310**make_request_options(extra_headers={"X-Foo": "Foo"}),
311),
312)
313assert request.headers.get("X-Foo") == "Foo"
314
315# `extra_headers` takes priority over `default_headers` when keys clash
316request = self.client.with_options(default_headers={"X-Bar": "true"})._build_request(
317FinalRequestOptions(
318method="post",
319url="/foo",
320**make_request_options(
321extra_headers={"X-Bar": "false"},
322),
323),
324)
325assert request.headers.get("X-Bar") == "false"
326
327def test_request_extra_query(self) -> None:
328request = self.client._build_request(
329FinalRequestOptions(
330method="post",
331url="/foo",
332**make_request_options(
333extra_query={"my_query_param": "Foo"},
334),
335),
336)
337params = cast(Dict[str, str], dict(request.url.params))
338assert params == {"my_query_param": "Foo"}
339
340# if both `query` and `extra_query` are given, they are merged
341request = self.client._build_request(
342FinalRequestOptions(
343method="post",
344url="/foo",
345**make_request_options(
346query={"bar": "1"},
347extra_query={"foo": "2"},
348),
349),
350)
351params = cast(Dict[str, str], dict(request.url.params))
352assert params == {"bar": "1", "foo": "2"}
353
354# `extra_query` takes priority over `query` when keys clash
355request = self.client._build_request(
356FinalRequestOptions(
357method="post",
358url="/foo",
359**make_request_options(
360query={"foo": "1"},
361extra_query={"foo": "2"},
362),
363),
364)
365params = cast(Dict[str, str], dict(request.url.params))
366assert params == {"foo": "2"}
367
368@pytest.mark.respx(base_url=base_url)
369def test_basic_union_response(self, respx_mock: MockRouter) -> None:
370class Model1(BaseModel):
371name: str
372
373class Model2(BaseModel):
374foo: str
375
376respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
377
378response = self.client.get("/foo", cast_to=cast(Any, Union[Model1, Model2]))
379assert isinstance(response, Model2)
380assert response.foo == "bar"
381
382@pytest.mark.respx(base_url=base_url)
383def test_union_response_different_types(self, respx_mock: MockRouter) -> None:
384"""Union of objects with the same field name using a different type"""
385
386class Model1(BaseModel):
387foo: int
388
389class Model2(BaseModel):
390foo: str
391
392respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
393
394response = self.client.get("/foo", cast_to=cast(Any, Union[Model1, Model2]))
395assert isinstance(response, Model2)
396assert response.foo == "bar"
397
398respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": 1}))
399
400response = self.client.get("/foo", cast_to=cast(Any, Union[Model1, Model2]))
401assert isinstance(response, Model1)
402assert response.foo == 1
403
c26014e2Stainless Bot2 years ago404@pytest.mark.respx(base_url=base_url)
405def test_non_application_json_content_type_for_json_data(self, respx_mock: MockRouter) -> None:
406"""
407Response that sets Content-Type to something other than application/json but returns json data
408"""
409
410class Model(BaseModel):
411foo: int
412
413respx_mock.get("/foo").mock(
414return_value=httpx.Response(
415200,
416content=json.dumps({"foo": 2}),
417headers={"Content-Type": "application/text"},
418)
419)
420
421response = self.client.get("/foo", cast_to=Model)
422assert isinstance(response, Model)
423assert response.foo == 2
424
0733934fStainless Bot2 years ago425def test_base_url_env(self) -> None:
426with update_env(OPENAI_BASE_URL="http://localhost:5000/from/env"):
427client = OpenAI(api_key=api_key, _strict_response_validation=True)
428assert client.base_url == "http://localhost:5000/from/env/"
429
08b8179aDavid Schnurr2 years ago430@pytest.mark.parametrize(
431"client",
432[
433OpenAI(base_url="http://localhost:5000/custom/path/", api_key=api_key, _strict_response_validation=True),
434OpenAI(
435base_url="http://localhost:5000/custom/path/",
436api_key=api_key,
437_strict_response_validation=True,
438http_client=httpx.Client(),
439),
440],
441ids=["standard", "custom http client"],
442)
443def test_base_url_trailing_slash(self, client: OpenAI) -> None:
444request = client._build_request(
445FinalRequestOptions(
446method="post",
447url="/foo",
448json_data={"foo": "bar"},
449),
450)
451assert request.url == "http://localhost:5000/custom/path/foo"
452
453@pytest.mark.parametrize(
454"client",
455[
456OpenAI(base_url="http://localhost:5000/custom/path/", api_key=api_key, _strict_response_validation=True),
457OpenAI(
458base_url="http://localhost:5000/custom/path/",
459api_key=api_key,
460_strict_response_validation=True,
461http_client=httpx.Client(),
462),
463],
464ids=["standard", "custom http client"],
465)
466def test_base_url_no_trailing_slash(self, client: OpenAI) -> None:
467request = client._build_request(
468FinalRequestOptions(
469method="post",
470url="/foo",
471json_data={"foo": "bar"},
472),
473)
474assert request.url == "http://localhost:5000/custom/path/foo"
475
476@pytest.mark.parametrize(
477"client",
478[
479OpenAI(base_url="http://localhost:5000/custom/path/", api_key=api_key, _strict_response_validation=True),
480OpenAI(
481base_url="http://localhost:5000/custom/path/",
482api_key=api_key,
483_strict_response_validation=True,
484http_client=httpx.Client(),
485),
486],
487ids=["standard", "custom http client"],
488)
489def test_absolute_request_url(self, client: OpenAI) -> None:
490request = client._build_request(
491FinalRequestOptions(
492method="post",
493url="https://myapi.com/foo",
494json_data={"foo": "bar"},
495),
496)
497assert request.url == "https://myapi.com/foo"
498
499def test_client_del(self) -> None:
500client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
501assert not client.is_closed()
502
503client.__del__()
504
505assert client.is_closed()
506
507def test_copied_client_does_not_close_http(self) -> None:
508client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
509assert not client.is_closed()
510
511copied = client.copy()
512assert copied is not client
513
514copied.__del__()
515
516assert not copied.is_closed()
517assert not client.is_closed()
518
519def test_client_context_manager(self) -> None:
520client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
521with client as c2:
522assert c2 is client
523assert not c2.is_closed()
524assert not client.is_closed()
525assert client.is_closed()
526
527@pytest.mark.respx(base_url=base_url)
528def test_client_response_validation_error(self, respx_mock: MockRouter) -> None:
529class Model(BaseModel):
530foo: str
531
532respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": {"invalid": True}}))
533
534with pytest.raises(APIResponseValidationError) as exc:
535self.client.get("/foo", cast_to=Model)
536
537assert isinstance(exc.value.__cause__, ValidationError)
538
539@pytest.mark.respx(base_url=base_url)
540def test_default_stream_cls(self, respx_mock: MockRouter) -> None:
541class Model(BaseModel):
542name: str
543
544respx_mock.post("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
545
546response = self.client.post("/foo", cast_to=Model, stream=True)
547assert isinstance(response, Stream)
548
549@pytest.mark.respx(base_url=base_url)
550def test_received_text_for_expected_json(self, respx_mock: MockRouter) -> None:
551class Model(BaseModel):
552name: str
553
554respx_mock.get("/foo").mock(return_value=httpx.Response(200, text="my-custom-format"))
555
556strict_client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
557
558with pytest.raises(APIResponseValidationError):
559strict_client.get("/foo", cast_to=Model)
560
561client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=False)
562
563response = client.get("/foo", cast_to=Model)
564assert isinstance(response, str) # type: ignore[unreachable]
565
566@pytest.mark.parametrize(
567"remaining_retries,retry_after,timeout",
568[
569[3, "20", 20],
570[3, "0", 0.5],
571[3, "-10", 0.5],
572[3, "60", 60],
573[3, "61", 0.5],
574[3, "Fri, 29 Sep 2023 16:26:57 GMT", 20],
575[3, "Fri, 29 Sep 2023 16:26:37 GMT", 0.5],
576[3, "Fri, 29 Sep 2023 16:26:27 GMT", 0.5],
577[3, "Fri, 29 Sep 2023 16:27:37 GMT", 60],
578[3, "Fri, 29 Sep 2023 16:27:38 GMT", 0.5],
579[3, "99999999999999999999999999999999999", 0.5],
580[3, "Zun, 29 Sep 2023 16:26:27 GMT", 0.5],
581[3, "", 0.5],
582[2, "", 0.5 * 2.0],
583[1, "", 0.5 * 4.0],
584],
585)
586@mock.patch("time.time", mock.MagicMock(return_value=1696004797))
587def test_parse_retry_after_header(self, remaining_retries: int, retry_after: str, timeout: float) -> None:
588client = OpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
589
590headers = httpx.Headers({"retry-after": retry_after})
591options = FinalRequestOptions(method="get", url="/foo", max_retries=3)
592calculated = client._calculate_retry_timeout(remaining_retries, options, headers)
593assert calculated == pytest.approx(timeout, 0.5 * 0.875) # pyright: ignore[reportUnknownMemberType]
594
595
596class TestAsyncOpenAI:
597client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
598
599@pytest.mark.respx(base_url=base_url)
600@pytest.mark.asyncio
601async def test_raw_response(self, respx_mock: MockRouter) -> None:
c5975bd0Stainless Bot2 years ago602respx_mock.post("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
08b8179aDavid Schnurr2 years ago603
604response = await self.client.post("/foo", cast_to=httpx.Response)
605assert response.status_code == 200
606assert isinstance(response, httpx.Response)
c5975bd0Stainless Bot2 years ago607assert response.json() == {"foo": "bar"}
08b8179aDavid Schnurr2 years ago608
609@pytest.mark.respx(base_url=base_url)
610@pytest.mark.asyncio
611async def test_raw_response_for_binary(self, respx_mock: MockRouter) -> None:
612respx_mock.post("/foo").mock(
613return_value=httpx.Response(200, headers={"Content-Type": "application/binary"}, content='{"foo": "bar"}')
614)
615
616response = await self.client.post("/foo", cast_to=httpx.Response)
617assert response.status_code == 200
618assert isinstance(response, httpx.Response)
c5975bd0Stainless Bot2 years ago619assert response.json() == {"foo": "bar"}
08b8179aDavid Schnurr2 years ago620
621def test_copy(self) -> None:
622copied = self.client.copy()
623assert id(copied) != id(self.client)
624
625copied = self.client.copy(api_key="another My API Key")
626assert copied.api_key == "another My API Key"
627assert self.client.api_key == "My API Key"
628
629def test_copy_default_options(self) -> None:
630# options that have a default are overridden correctly
631copied = self.client.copy(max_retries=7)
632assert copied.max_retries == 7
633assert self.client.max_retries == 2
634
635copied2 = copied.copy(max_retries=6)
636assert copied2.max_retries == 6
637assert copied.max_retries == 7
638
639# timeout
640assert isinstance(self.client.timeout, httpx.Timeout)
641copied = self.client.copy(timeout=None)
642assert copied.timeout is None
643assert isinstance(self.client.timeout, httpx.Timeout)
644
645def test_copy_default_headers(self) -> None:
646client = AsyncOpenAI(
647base_url=base_url, api_key=api_key, _strict_response_validation=True, default_headers={"X-Foo": "bar"}
648)
649assert client.default_headers["X-Foo"] == "bar"
650
651# does not override the already given value when not specified
652copied = client.copy()
653assert copied.default_headers["X-Foo"] == "bar"
654
655# merges already given headers
656copied = client.copy(default_headers={"X-Bar": "stainless"})
657assert copied.default_headers["X-Foo"] == "bar"
658assert copied.default_headers["X-Bar"] == "stainless"
659
660# uses new values for any already given headers
661copied = client.copy(default_headers={"X-Foo": "stainless"})
662assert copied.default_headers["X-Foo"] == "stainless"
663
664# set_default_headers
665
666# completely overrides already set values
667copied = client.copy(set_default_headers={})
668assert copied.default_headers.get("X-Foo") is None
669
670copied = client.copy(set_default_headers={"X-Bar": "Robert"})
671assert copied.default_headers["X-Bar"] == "Robert"
672
673with pytest.raises(
674ValueError,
675match="`default_headers` and `set_default_headers` arguments are mutually exclusive",
676):
677client.copy(set_default_headers={}, default_headers={"X-Foo": "Bar"})
678
679def test_copy_default_query(self) -> None:
680client = AsyncOpenAI(
681base_url=base_url, api_key=api_key, _strict_response_validation=True, default_query={"foo": "bar"}
682)
683assert _get_params(client)["foo"] == "bar"
684
685# does not override the already given value when not specified
686copied = client.copy()
687assert _get_params(copied)["foo"] == "bar"
688
689# merges already given params
690copied = client.copy(default_query={"bar": "stainless"})
691params = _get_params(copied)
692assert params["foo"] == "bar"
693assert params["bar"] == "stainless"
694
695# uses new values for any already given headers
696copied = client.copy(default_query={"foo": "stainless"})
697assert _get_params(copied)["foo"] == "stainless"
698
699# set_default_query
700
701# completely overrides already set values
702copied = client.copy(set_default_query={})
703assert _get_params(copied) == {}
704
705copied = client.copy(set_default_query={"bar": "Robert"})
706assert _get_params(copied)["bar"] == "Robert"
707
708with pytest.raises(
709ValueError,
710# TODO: update
711match="`default_query` and `set_default_query` arguments are mutually exclusive",
712):
713client.copy(set_default_query={}, default_query={"foo": "Bar"})
714
715def test_copy_signature(self) -> None:
716# ensure the same parameters that can be passed to the client are defined in the `.copy()` method
717init_signature = inspect.signature(
718# mypy doesn't like that we access the `__init__` property.
719self.client.__init__, # type: ignore[misc]
720)
721copy_signature = inspect.signature(self.client.copy)
722exclude_params = {"transport", "proxies", "_strict_response_validation"}
723
724for name in init_signature.parameters.keys():
725if name in exclude_params:
726continue
727
728copy_param = copy_signature.parameters.get(name)
729assert copy_param is not None, f"copy() signature is missing the {name} param"
730
731async def test_request_timeout(self) -> None:
732request = self.client._build_request(FinalRequestOptions(method="get", url="/foo"))
733timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
734assert timeout == DEFAULT_TIMEOUT
735
736request = self.client._build_request(
737FinalRequestOptions(method="get", url="/foo", timeout=httpx.Timeout(100.0))
738)
739timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
740assert timeout == httpx.Timeout(100.0)
741
742async def test_client_timeout_option(self) -> None:
743client = AsyncOpenAI(
744base_url=base_url, api_key=api_key, _strict_response_validation=True, timeout=httpx.Timeout(0)
745)
746
747request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
748timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
749assert timeout == httpx.Timeout(0)
750
751async def test_http_client_timeout_option(self) -> None:
752# custom timeout given to the httpx client should be used
753async with httpx.AsyncClient(timeout=None) as http_client:
754client = AsyncOpenAI(
755base_url=base_url, api_key=api_key, _strict_response_validation=True, http_client=http_client
756)
757
758request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
759timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
760assert timeout == httpx.Timeout(None)
761
762# no timeout given to the httpx client should not use the httpx default
763async with httpx.AsyncClient() as http_client:
764client = AsyncOpenAI(
765base_url=base_url, api_key=api_key, _strict_response_validation=True, http_client=http_client
766)
767
768request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
769timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
770assert timeout == DEFAULT_TIMEOUT
771
772# explicitly passing the default timeout currently results in it being ignored
773async with httpx.AsyncClient(timeout=HTTPX_DEFAULT_TIMEOUT) as http_client:
774client = AsyncOpenAI(
775base_url=base_url, api_key=api_key, _strict_response_validation=True, http_client=http_client
776)
777
778request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
779timeout = httpx.Timeout(**request.extensions["timeout"]) # type: ignore
780assert timeout == DEFAULT_TIMEOUT # our default
781
782def test_default_headers_option(self) -> None:
783client = AsyncOpenAI(
784base_url=base_url, api_key=api_key, _strict_response_validation=True, default_headers={"X-Foo": "bar"}
785)
786request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
787assert request.headers.get("x-foo") == "bar"
788assert request.headers.get("x-stainless-lang") == "python"
789
790client2 = AsyncOpenAI(
791base_url=base_url,
792api_key=api_key,
793_strict_response_validation=True,
794default_headers={
795"X-Foo": "stainless",
796"X-Stainless-Lang": "my-overriding-header",
797},
798)
799request = client2._build_request(FinalRequestOptions(method="get", url="/foo"))
800assert request.headers.get("x-foo") == "stainless"
801assert request.headers.get("x-stainless-lang") == "my-overriding-header"
802
803def test_validate_headers(self) -> None:
804client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
805request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
806assert request.headers.get("Authorization") == f"Bearer {api_key}"
807
808with pytest.raises(Exception):
809client2 = AsyncOpenAI(base_url=base_url, api_key=None, _strict_response_validation=True)
810_ = client2
811
812def test_default_query_option(self) -> None:
813client = AsyncOpenAI(
814base_url=base_url, api_key=api_key, _strict_response_validation=True, default_query={"query_param": "bar"}
815)
816request = client._build_request(FinalRequestOptions(method="get", url="/foo"))
817url = httpx.URL(request.url)
818assert dict(url.params) == {"query_param": "bar"}
819
820request = client._build_request(
821FinalRequestOptions(
822method="get",
823url="/foo",
824params={"foo": "baz", "query_param": "overriden"},
825)
826)
827url = httpx.URL(request.url)
828assert dict(url.params) == {"foo": "baz", "query_param": "overriden"}
829
830def test_request_extra_json(self) -> None:
831request = self.client._build_request(
832FinalRequestOptions(
833method="post",
834url="/foo",
835json_data={"foo": "bar"},
836extra_json={"baz": False},
837),
838)
839data = json.loads(request.content.decode("utf-8"))
840assert data == {"foo": "bar", "baz": False}
841
842request = self.client._build_request(
843FinalRequestOptions(
844method="post",
845url="/foo",
846extra_json={"baz": False},
847),
848)
849data = json.loads(request.content.decode("utf-8"))
850assert data == {"baz": False}
851
852# `extra_json` takes priority over `json_data` when keys clash
853request = self.client._build_request(
854FinalRequestOptions(
855method="post",
856url="/foo",
857json_data={"foo": "bar", "baz": True},
858extra_json={"baz": None},
859),
860)
861data = json.loads(request.content.decode("utf-8"))
862assert data == {"foo": "bar", "baz": None}
863
864def test_request_extra_headers(self) -> None:
865request = self.client._build_request(
866FinalRequestOptions(
867method="post",
868url="/foo",
869**make_request_options(extra_headers={"X-Foo": "Foo"}),
870),
871)
872assert request.headers.get("X-Foo") == "Foo"
873
874# `extra_headers` takes priority over `default_headers` when keys clash
875request = self.client.with_options(default_headers={"X-Bar": "true"})._build_request(
876FinalRequestOptions(
877method="post",
878url="/foo",
879**make_request_options(
880extra_headers={"X-Bar": "false"},
881),
882),
883)
884assert request.headers.get("X-Bar") == "false"
885
886def test_request_extra_query(self) -> None:
887request = self.client._build_request(
888FinalRequestOptions(
889method="post",
890url="/foo",
891**make_request_options(
892extra_query={"my_query_param": "Foo"},
893),
894),
895)
896params = cast(Dict[str, str], dict(request.url.params))
897assert params == {"my_query_param": "Foo"}
898
899# if both `query` and `extra_query` are given, they are merged
900request = self.client._build_request(
901FinalRequestOptions(
902method="post",
903url="/foo",
904**make_request_options(
905query={"bar": "1"},
906extra_query={"foo": "2"},
907),
908),
909)
910params = cast(Dict[str, str], dict(request.url.params))
911assert params == {"bar": "1", "foo": "2"}
912
913# `extra_query` takes priority over `query` when keys clash
914request = self.client._build_request(
915FinalRequestOptions(
916method="post",
917url="/foo",
918**make_request_options(
919query={"foo": "1"},
920extra_query={"foo": "2"},
921),
922),
923)
924params = cast(Dict[str, str], dict(request.url.params))
925assert params == {"foo": "2"}
926
927@pytest.mark.respx(base_url=base_url)
928async def test_basic_union_response(self, respx_mock: MockRouter) -> None:
929class Model1(BaseModel):
930name: str
931
932class Model2(BaseModel):
933foo: str
934
935respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
936
937response = await self.client.get("/foo", cast_to=cast(Any, Union[Model1, Model2]))
938assert isinstance(response, Model2)
939assert response.foo == "bar"
940
941@pytest.mark.respx(base_url=base_url)
942async def test_union_response_different_types(self, respx_mock: MockRouter) -> None:
943"""Union of objects with the same field name using a different type"""
944
945class Model1(BaseModel):
946foo: int
947
948class Model2(BaseModel):
949foo: str
950
951respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
952
953response = await self.client.get("/foo", cast_to=cast(Any, Union[Model1, Model2]))
954assert isinstance(response, Model2)
955assert response.foo == "bar"
956
957respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": 1}))
958
959response = await self.client.get("/foo", cast_to=cast(Any, Union[Model1, Model2]))
960assert isinstance(response, Model1)
961assert response.foo == 1
962
c26014e2Stainless Bot2 years ago963@pytest.mark.respx(base_url=base_url)
964async def test_non_application_json_content_type_for_json_data(self, respx_mock: MockRouter) -> None:
965"""
966Response that sets Content-Type to something other than application/json but returns json data
967"""
968
969class Model(BaseModel):
970foo: int
971
972respx_mock.get("/foo").mock(
973return_value=httpx.Response(
974200,
975content=json.dumps({"foo": 2}),
976headers={"Content-Type": "application/text"},
977)
978)
979
980response = await self.client.get("/foo", cast_to=Model)
981assert isinstance(response, Model)
982assert response.foo == 2
983
0733934fStainless Bot2 years ago984def test_base_url_env(self) -> None:
985with update_env(OPENAI_BASE_URL="http://localhost:5000/from/env"):
986client = AsyncOpenAI(api_key=api_key, _strict_response_validation=True)
987assert client.base_url == "http://localhost:5000/from/env/"
988
08b8179aDavid Schnurr2 years ago989@pytest.mark.parametrize(
990"client",
991[
992AsyncOpenAI(
993base_url="http://localhost:5000/custom/path/", api_key=api_key, _strict_response_validation=True
994),
995AsyncOpenAI(
996base_url="http://localhost:5000/custom/path/",
997api_key=api_key,
998_strict_response_validation=True,
999http_client=httpx.AsyncClient(),
1000),
1001],
1002ids=["standard", "custom http client"],
1003)
1004def test_base_url_trailing_slash(self, client: AsyncOpenAI) -> None:
1005request = client._build_request(
1006FinalRequestOptions(
1007method="post",
1008url="/foo",
1009json_data={"foo": "bar"},
1010),
1011)
1012assert request.url == "http://localhost:5000/custom/path/foo"
1013
1014@pytest.mark.parametrize(
1015"client",
1016[
1017AsyncOpenAI(
1018base_url="http://localhost:5000/custom/path/", api_key=api_key, _strict_response_validation=True
1019),
1020AsyncOpenAI(
1021base_url="http://localhost:5000/custom/path/",
1022api_key=api_key,
1023_strict_response_validation=True,
1024http_client=httpx.AsyncClient(),
1025),
1026],
1027ids=["standard", "custom http client"],
1028)
1029def test_base_url_no_trailing_slash(self, client: AsyncOpenAI) -> None:
1030request = client._build_request(
1031FinalRequestOptions(
1032method="post",
1033url="/foo",
1034json_data={"foo": "bar"},
1035),
1036)
1037assert request.url == "http://localhost:5000/custom/path/foo"
1038
1039@pytest.mark.parametrize(
1040"client",
1041[
1042AsyncOpenAI(
1043base_url="http://localhost:5000/custom/path/", api_key=api_key, _strict_response_validation=True
1044),
1045AsyncOpenAI(
1046base_url="http://localhost:5000/custom/path/",
1047api_key=api_key,
1048_strict_response_validation=True,
1049http_client=httpx.AsyncClient(),
1050),
1051],
1052ids=["standard", "custom http client"],
1053)
1054def test_absolute_request_url(self, client: AsyncOpenAI) -> None:
1055request = client._build_request(
1056FinalRequestOptions(
1057method="post",
1058url="https://myapi.com/foo",
1059json_data={"foo": "bar"},
1060),
1061)
1062assert request.url == "https://myapi.com/foo"
1063
1064async def test_client_del(self) -> None:
1065client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
1066assert not client.is_closed()
1067
1068client.__del__()
1069
1070await asyncio.sleep(0.2)
1071assert client.is_closed()
1072
1073async def test_copied_client_does_not_close_http(self) -> None:
1074client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
1075assert not client.is_closed()
1076
1077copied = client.copy()
1078assert copied is not client
1079
1080copied.__del__()
1081
1082await asyncio.sleep(0.2)
1083assert not copied.is_closed()
1084assert not client.is_closed()
1085
1086async def test_client_context_manager(self) -> None:
1087client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
1088async with client as c2:
1089assert c2 is client
1090assert not c2.is_closed()
1091assert not client.is_closed()
1092assert client.is_closed()
1093
1094@pytest.mark.respx(base_url=base_url)
1095@pytest.mark.asyncio
1096async def test_client_response_validation_error(self, respx_mock: MockRouter) -> None:
1097class Model(BaseModel):
1098foo: str
1099
1100respx_mock.get("/foo").mock(return_value=httpx.Response(200, json={"foo": {"invalid": True}}))
1101
1102with pytest.raises(APIResponseValidationError) as exc:
1103await self.client.get("/foo", cast_to=Model)
1104
1105assert isinstance(exc.value.__cause__, ValidationError)
1106
1107@pytest.mark.respx(base_url=base_url)
1108@pytest.mark.asyncio
1109async def test_default_stream_cls(self, respx_mock: MockRouter) -> None:
1110class Model(BaseModel):
1111name: str
1112
1113respx_mock.post("/foo").mock(return_value=httpx.Response(200, json={"foo": "bar"}))
1114
1115response = await self.client.post("/foo", cast_to=Model, stream=True)
1116assert isinstance(response, AsyncStream)
1117
1118@pytest.mark.respx(base_url=base_url)
1119@pytest.mark.asyncio
1120async def test_received_text_for_expected_json(self, respx_mock: MockRouter) -> None:
1121class Model(BaseModel):
1122name: str
1123
1124respx_mock.get("/foo").mock(return_value=httpx.Response(200, text="my-custom-format"))
1125
1126strict_client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
1127
1128with pytest.raises(APIResponseValidationError):
1129await strict_client.get("/foo", cast_to=Model)
1130
1131client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=False)
1132
1133response = await client.get("/foo", cast_to=Model)
1134assert isinstance(response, str) # type: ignore[unreachable]
1135
1136@pytest.mark.parametrize(
1137"remaining_retries,retry_after,timeout",
1138[
1139[3, "20", 20],
1140[3, "0", 0.5],
1141[3, "-10", 0.5],
1142[3, "60", 60],
1143[3, "61", 0.5],
1144[3, "Fri, 29 Sep 2023 16:26:57 GMT", 20],
1145[3, "Fri, 29 Sep 2023 16:26:37 GMT", 0.5],
1146[3, "Fri, 29 Sep 2023 16:26:27 GMT", 0.5],
1147[3, "Fri, 29 Sep 2023 16:27:37 GMT", 60],
1148[3, "Fri, 29 Sep 2023 16:27:38 GMT", 0.5],
1149[3, "99999999999999999999999999999999999", 0.5],
1150[3, "Zun, 29 Sep 2023 16:26:27 GMT", 0.5],
1151[3, "", 0.5],
1152[2, "", 0.5 * 2.0],
1153[1, "", 0.5 * 4.0],
1154],
1155)
1156@mock.patch("time.time", mock.MagicMock(return_value=1696004797))
1157@pytest.mark.asyncio
1158async def test_parse_retry_after_header(self, remaining_retries: int, retry_after: str, timeout: float) -> None:
1159client = AsyncOpenAI(base_url=base_url, api_key=api_key, _strict_response_validation=True)
1160
1161headers = httpx.Headers({"retry-after": retry_after})
1162options = FinalRequestOptions(method="get", url="/foo", max_retries=3)
1163calculated = client._calculate_retry_timeout(remaining_retries, options, headers)
1164assert calculated == pytest.approx(timeout, 0.5 * 0.875) # pyright: ignore[reportUnknownMemberType]