Skip to content

Commit 1c4dd37

Browse files
Matt CoulterMatt Coulter
authored andcommitted
feat: support **kwargs with request payload to function parameters
1 parent 636b806 commit 1c4dd37

5 files changed

Lines changed: 54 additions & 10 deletions

File tree

src/pydantic_openapi_generator/language_converters/python/client_generator.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -807,6 +807,7 @@ def _generate_service_operation(
807807
query_params = generate_operation_parameters(op, generator_config, "query")
808808
header_params = generate_operation_parameters(op, generator_config, "header")
809809
all_params = path_params + query_params + header_params
810+
call_kwargs = [param.code_name for param in all_params]
810811
params = _method_signature(all_params, _body_signature_param(op) if body_param is not None else None)
811812
path_name = _resolved_path_name(path_name, path_params)
812813

@@ -816,12 +817,16 @@ def _generate_service_operation(
816817
norm_ph = common.normalize_symbol(ph)
817818
if norm_ph not in existing_param_names and norm_ph:
818819
params = f"{norm_ph}: Any, " + params
820+
call_kwargs.insert(0, norm_ph)
821+
if body_param is not None:
822+
call_kwargs.append("data")
819823

820824
operation_id = generate_operation_id(op, http_operation, path_name)
821825
return_type = generate_return_type(op)
822826

823827
so = ServiceOperation(
824828
params=params,
829+
call_kwargs=call_kwargs,
825830
operation_id=operation_id,
826831
path_params=path_params,
827832
query_params=query_params,

src/pydantic_openapi_generator/language_converters/python/templates/async_client_httpx_pydantic_2.jinja2

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@ class AsyncClient(BaseModel):
2727
{{ field.code_name }}: {{ field.field_type_hint }}{% if field.field_default is not none %} = {{ field.field_default }}{% endif %}{{ "\n" }}
2828
{% endfor %}
2929

30-
async def get_access_token(self) -> Optional[str]:
30+
async def get_access_token(self, **kwargs: Any) -> Optional[str]:
3131
{% if env_token_name is not none %}
3232
try:
3333
return os.environ["{{ env_token_name }}"]
@@ -48,7 +48,7 @@ class AsyncClient(BaseModel):
4848
{% endif %}
4949

5050
{% for function in client_functions %}
51-
async def {{ function.getter_name }}(self) -> {{ function.base_type_hint }}:
51+
async def {{ function.getter_name }}(self, **kwargs: Any) -> {{ function.base_type_hint }}:
5252
raise NotImplementedError
5353

5454
{% endfor %}
@@ -59,6 +59,11 @@ AsyncGenerator[str | {% if op.sse_data_handler and op.sse_data_handler.type %}{%
5959
{{ op.return_type.return_type_hint or "None" }}
6060
{%- endif -%}:
6161
base_url = self.base_url
62+
_request_kwargs = {
63+
{% for kwarg in op.call_kwargs %}
64+
"{{ kwarg }}": {{ kwarg }},
65+
{% endfor %}
66+
}
6267
{% for param in op.path_params + op.query_params + op.header_params %}
6368
{% if param.source == "ClassVar" %}
6469
{% if param.is_secret %}
@@ -67,7 +72,7 @@ AsyncGenerator[str | {% if op.sse_data_handler and op.sse_data_handler.type %}{%
6772
{{ param.local_name }} = {{ param.code_name }} if {{ param.code_name }} is not None else self.{{ param.code_name }}
6873
{% endif %}
6974
{% elif param.source == "Function" %}
70-
{{ param.local_name }} = {{ param.code_name }} if {{ param.code_name }} is not None else await self.{{ param.getter_name }}()
75+
{{ param.local_name }} = {{ param.code_name }} if {{ param.code_name }} is not None else await self.{{ param.getter_name }}(**_request_kwargs)
7176
{% endif %}
7277
{% endfor %}
7378
path = f"{{ op.path_name }}"
@@ -89,7 +94,7 @@ AsyncGenerator[str | {% if op.sse_data_handler and op.sse_data_handler.type %}{%
8994
}
9095
headers = {k: v for (k, v) in headers.items() if v is not None}
9196

92-
_token = await self.get_access_token()
97+
_token = await self.get_access_token(**_request_kwargs)
9398
if _token:
9499
headers["Authorization"] = f"Bearer {_token}"
95100

src/pydantic_openapi_generator/language_converters/python/templates/sync_client_httpx_pydantic_2.jinja2

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@ class SyncClient(BaseModel):
2727
{{ field.code_name }}: {{ field.field_type_hint }}{% if field.field_default is not none %} = {{ field.field_default }}{% endif %}{{ "\n" }}
2828
{% endfor %}
2929

30-
def get_access_token(self) -> Optional[str]:
30+
def get_access_token(self, **kwargs: Any) -> Optional[str]:
3131
{% if env_token_name is not none %}
3232
try:
3333
return os.environ["{{ env_token_name }}"]
@@ -48,7 +48,7 @@ class SyncClient(BaseModel):
4848
{% endif %}
4949

5050
{% for function in client_functions %}
51-
def {{ function.getter_name }}(self) -> {{ function.base_type_hint }}:
51+
def {{ function.getter_name }}(self, **kwargs: Any) -> {{ function.base_type_hint }}:
5252
raise NotImplementedError
5353

5454
{% endfor %}
@@ -59,6 +59,11 @@ Generator[str | {% if op.sse_data_handler and op.sse_data_handler.type %}{% if o
5959
{{ op.return_type.return_type_hint or "None" }}
6060
{%- endif -%}:
6161
base_url = self.base_url
62+
_request_kwargs = {
63+
{% for kwarg in op.call_kwargs %}
64+
"{{ kwarg }}": {{ kwarg }},
65+
{% endfor %}
66+
}
6267
{% for param in op.path_params + op.query_params + op.header_params %}
6368
{% if param.source == "ClassVar" %}
6469
{% if param.is_secret %}
@@ -67,7 +72,7 @@ Generator[str | {% if op.sse_data_handler and op.sse_data_handler.type %}{% if o
6772
{{ param.local_name }} = {{ param.code_name }} if {{ param.code_name }} is not None else self.{{ param.code_name }}
6873
{% endif %}
6974
{% elif param.source == "Function" %}
70-
{{ param.local_name }} = {{ param.code_name }} if {{ param.code_name }} is not None else self.{{ param.getter_name }}()
75+
{{ param.local_name }} = {{ param.code_name }} if {{ param.code_name }} is not None else self.{{ param.getter_name }}(**_request_kwargs)
7176
{% endif %}
7277
{% endfor %}
7378
path = f"{{ op.path_name }}"
@@ -89,7 +94,7 @@ Generator[str | {% if op.sse_data_handler and op.sse_data_handler.type %}{% if o
8994
}
9095
headers = {k: v for (k, v) in headers.items() if v is not None}
9196

92-
_token = self.get_access_token()
97+
_token = self.get_access_token(**_request_kwargs)
9398
if _token:
9499
headers["Authorization"] = f"Bearer {_token}"
95100

src/pydantic_openapi_generator/models.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -95,6 +95,7 @@ class GeneratedParameter(BaseModel):
9595

9696
class ServiceOperation(BaseModel):
9797
params: str
98+
call_kwargs: List[str] = Field(default_factory=list)
9899
operation_id: str
99100
path_params: List[GeneratedParameter] = Field(default_factory=list)
100101
query_params: List[GeneratedParameter]

tests/test_client_generator_contracts.py

Lines changed: 30 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
import sys
66
import tempfile
77
from pathlib import Path
8+
from typing import Any
89
from uuid import uuid4
910

1011
import httpx
@@ -384,11 +385,18 @@ def thing_handler(request: httpx.Request) -> httpx.Response:
384385

385386
class Client(client_module.AsyncClient):
386387
getter_calls: int = 0
388+
getter_kwargs: dict[str, Any] = {}
389+
token_kwargs: dict[str, Any] = {}
387390

388-
async def get_dynamic_header(self) -> str:
391+
async def get_dynamic_header(self, **kwargs: Any) -> str:
389392
self.getter_calls += 1
393+
self.getter_kwargs = kwargs
390394
return "from-getter"
391395

396+
async def get_access_token(self, **kwargs: Any) -> str | None:
397+
self.token_kwargs = kwargs
398+
return None
399+
392400
client = Client(test_header="from-client", optional_header="from-class")
393401
asyncio.run(_run_async_configured_contract(client))
394402
assert client.getter_calls == 1
@@ -399,11 +407,18 @@ async def get_dynamic_header(self) -> str:
399407

400408
class Client(client_module.SyncClient):
401409
getter_calls: int = 0
410+
getter_kwargs: dict[str, Any] = {}
411+
token_kwargs: dict[str, Any] = {}
402412

403-
def get_dynamic_header(self) -> str:
413+
def get_dynamic_header(self, **kwargs: Any) -> str:
404414
self.getter_calls += 1
415+
self.getter_kwargs = kwargs
405416
return "from-getter"
406417

418+
def get_access_token(self, **kwargs: Any) -> str | None:
419+
self.token_kwargs = kwargs
420+
return None
421+
407422
client = Client(test_header="from-client", optional_header="from-class")
408423
client.getThing(product_brand="CGU")
409424
client.getThing(
@@ -414,6 +429,19 @@ def get_dynamic_header(self) -> str:
414429
)
415430
assert client.getter_calls == 1
416431

432+
assert client.getter_kwargs == {
433+
"product_brand": "CGU",
434+
"test_header": None,
435+
"optional_header": None,
436+
"dynamic_header": None,
437+
}
438+
assert client.token_kwargs == {
439+
"product_brand": "AMI",
440+
"test_header": "from-override",
441+
"optional_header": "optional-override",
442+
"dynamic_header": "function-override",
443+
}
444+
417445
first_request = requests[0]
418446
assert first_request.headers["X-Test-Header"] == "from-client"
419447
assert first_request.headers["X-Optional-Header"] == "from-class"

0 commit comments

Comments
 (0)