|
18 | 18 | import os |
19 | 19 | import pytest |
20 | 20 | import pytest_asyncio |
| 21 | +from requests.adapters import HTTPAdapter |
| 22 | +from urllib3.poolmanager import PoolManager |
21 | 23 |
|
22 | 24 | from typing import Sequence, Tuple |
23 | 25 |
|
@@ -328,7 +330,7 @@ def _read_response_metadata_stream(self): |
328 | 330 | def intercept_unary_unary(self, continuation, client_call_details, request): |
329 | 331 | self._add_request_metadata(client_call_details) |
330 | 332 | response = continuation(client_call_details, request) |
331 | | - metadata = [(k, str(v)) for k, v in response.trailing_metadata()] |
| 333 | + metadata = [(k, str(v)) for k, v in response.initial_metadata()] + [(k, str(v)) for k, v in response.trailing_metadata()] |
332 | 334 | self.response_metadata = metadata |
333 | 335 | return response |
334 | 336 |
|
@@ -453,36 +455,64 @@ async def intercepted_echo_grpc_async(): |
453 | 455 | return EchoAsyncClient(transport=transport), interceptor |
454 | 456 |
|
455 | 457 |
|
| 458 | +class HostNameIgnoringAdapter(HTTPAdapter): |
| 459 | + """Custom HTTPAdapter that disables hostname verification for local self-signed certs.""" |
| 460 | + def init_poolmanager(self, connections, maxsize, block=False, **pool_kwargs): |
| 461 | + self.poolmanager = PoolManager( |
| 462 | + num_pools=connections, |
| 463 | + maxsize=maxsize, |
| 464 | + block=block, |
| 465 | + assert_hostname=False, |
| 466 | + **pool_kwargs |
| 467 | + ) |
| 468 | + |
| 469 | + |
456 | 470 | @pytest.fixture |
457 | | -def intercepted_echo_rest(): |
| 471 | +def intercepted_echo_rest(use_mtls): |
458 | 472 | transport_name = "rest" |
459 | 473 | transport_cls = EchoClient.get_transport_class(transport_name) |
460 | 474 | interceptor = EchoMetadataClientRestInterceptor() |
461 | 475 |
|
462 | | - # The custom host explicitly bypasses https. |
| 476 | + url_scheme = "https" if use_mtls else "http" |
463 | 477 | transport = transport_cls( |
464 | 478 | credentials=ga_credentials.AnonymousCredentials(), |
465 | 479 | host="localhost:7469", |
466 | | - url_scheme="http", |
| 480 | + url_scheme=url_scheme, |
467 | 481 | interceptor=interceptor, |
468 | 482 | ) |
| 483 | + if use_mtls: |
| 484 | + dir = os.path.dirname(__file__) |
| 485 | + cert_path = os.path.join(dir, "../cert/mtls.crt") |
| 486 | + key_path = os.path.join(dir, "../cert/mtls.key") |
| 487 | + transport._session.verify = cert_path |
| 488 | + transport._session.cert = (cert_path, key_path) |
| 489 | + transport._session.mount("https://", HostNameIgnoringAdapter()) |
| 490 | + |
469 | 491 | return EchoClient(transport=transport), interceptor |
470 | 492 |
|
471 | 493 |
|
472 | 494 | @pytest.fixture |
473 | | -def intercepted_echo_rest_async(): |
| 495 | +def intercepted_echo_rest_async(use_mtls): |
474 | 496 | if not HAS_ASYNC_REST_ECHO_TRANSPORT: |
475 | 497 | pytest.skip("Skipping test with async rest.") |
476 | 498 |
|
477 | 499 | transport_name = "rest_asyncio" |
478 | 500 | transport_cls = EchoAsyncClient.get_transport_class(transport_name) |
479 | 501 | interceptor = EchoMetadataClientRestAsyncInterceptor() |
480 | 502 |
|
481 | | - # The custom host explicitly bypasses https. |
| 503 | + url_scheme = "https" if use_mtls else "http" |
482 | 504 | transport = transport_cls( |
483 | 505 | credentials=async_anonymous_credentials(), |
484 | 506 | host="localhost:7469", |
485 | | - url_scheme="http", |
| 507 | + url_scheme=url_scheme, |
486 | 508 | interceptor=interceptor, |
487 | 509 | ) |
| 510 | + if use_mtls: |
| 511 | + dir = os.path.dirname(__file__) |
| 512 | + cert_path = os.path.join(dir, "../cert/mtls.crt") |
| 513 | + key_path = os.path.join(dir, "../cert/mtls.key") |
| 514 | + transport._session.verify = cert_path |
| 515 | + transport._session.cert = (cert_path, key_path) |
| 516 | + transport._session.mount("https://", HostNameIgnoringAdapter()) |
| 517 | + |
488 | 518 | return EchoAsyncClient(transport=transport), interceptor |
0 commit comments