Skip to content
This repository was archived by the owner on May 22, 2026. It is now read-only.
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions test/azure/legacy/AcceptanceTests/asynctests/test_paging.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,3 +240,10 @@ async def test_item_name_with_xms_client_name(self, client):
async for item in pages:
items.append(item)
assert len(items) == 1

@pytest.mark.asyncio
async def test_duplicate_params(self, client):
pages = [p async for p in client.paging.duplicate_params(filter="foo")]
assert len(pages) == 1
assert pages[0].properties.id == 1
assert pages[0].properties.name == "Product"
7 changes: 7 additions & 0 deletions test/azure/legacy/AcceptanceTests/test_paging.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,6 +180,13 @@ def test_initial_response_no_items(self, client):
items = [i for i in pages]
assert len(items) == 1

@pytest.mark.skipif(sys.version_info < (3,5), reason="2.7 does different url encoding")
def test_duplicate_params(self, client):
pages = list(client.paging.duplicate_params(filter="foo"))
assert len(pages) == 1
assert pages[0].properties.id == 1
assert pages[0].properties.name == "Product"

def test_models(self):
from paging.models import OperationResult
if sys.version_info >= (3,5):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,47 @@
# IN THE SOFTWARE.
#
# --------------------------------------------------------------------------
from typing import List
import importlib
from ._auto_rest_paging_test_service import AutoRestPagingTestService as AutoRestPagingTestServiceGenerated
from azure.core.pipeline.policies import SansIOHTTPPolicy
try:
binary_type = str
import urlparse # type: ignore
from urllib import urlencode
except ImportError:
binary_type = bytes # type: ignore
from urllib import parse as urlparse
from urllib.parse import urlencode

class RemoveDuplicateParamsPolicy(SansIOHTTPPolicy):
def __init__(self, duplicate_param_names):
# type: (List[str]) -> None
self.duplicate_param_names = duplicate_param_names

def on_request(self, request):
parsed_url = urlparse.urlparse(request.http_request.url)
query_params = urlparse.parse_qs(parsed_url.query)
# service returned will be later in the url because of how we format
filtered_query_params = {
k: v[-1:] if k in self.duplicate_param_names else v
for k, v in query_params.items()
}
request.http_request.url = request.http_request.url.replace(parsed_url.query, "") + urlencode(filtered_query_params, doseq=True)
return super(RemoveDuplicateParamsPolicy, self).on_request(request)

class AutoRestPagingTestService(AutoRestPagingTestServiceGenerated):
def __init__(self, *args, **kwargs):
per_call_policies = kwargs.pop("per_call_policies", [])
params_policy = RemoveDuplicateParamsPolicy(duplicate_param_names=["$filter", "$skiptoken"])
try:
per_call_policies.append(params_policy)
except AttributeError:
per_call_policies = [per_call_policies, params_policy]
super(AutoRestPagingTestService, self).__init__(*args, per_call_policies=per_call_policies, **kwargs)

# This file is used for handwritten extensions to the generated code. Example:
# https://github.com/Azure/azure-sdk-for-python/blob/main/doc/dev/customize_code/how-to-patch-sdk-code.md
def patch_sdk():
pass
curr_package = importlib.import_module("paging")
curr_package.AutoRestPagingTestService = AutoRestPagingTestService
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,22 @@
# IN THE SOFTWARE.
#
# --------------------------------------------------------------------------
import importlib
from .._patch import RemoveDuplicateParamsPolicy
from ._auto_rest_paging_test_service import AutoRestPagingTestService as AutoRestPagingTestServiceGenerated

class AutoRestPagingTestService(AutoRestPagingTestServiceGenerated):
def __init__(self, *args, **kwargs):
per_call_policies = kwargs.pop("per_call_policies", [])
params_policy = RemoveDuplicateParamsPolicy(duplicate_param_names=["$filter", "$skiptoken"])
try:
per_call_policies.append(params_policy)
except AttributeError:
per_call_policies = [per_call_policies, params_policy]
super().__init__(*args, per_call_policies=per_call_policies, **kwargs)

# This file is used for handwritten extensions to the generated code. Example:
# https://github.com/Azure/azure-sdk-for-python/blob/main/doc/dev/customize_code/how-to-patch-sdk-code.md
def patch_sdk():
pass
curr_package = importlib.import_module("paging.aio")
curr_package.AutoRestPagingTestService = AutoRestPagingTestService
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
from ... import models as _models
from ..._vendor import _convert_request
from ...operations._paging_operations import (
build_duplicate_params_request,
build_first_response_empty_request,
build_get_multiple_pages_failure_request,
build_get_multiple_pages_failure_uri_request,
Expand Down Expand Up @@ -458,6 +459,69 @@ async def get_next(next_link=None):

get_with_query_params.metadata = {"url": "/paging/multiple/getWithQueryParams"} # type: ignore

@distributed_trace
def duplicate_params(self, filter: Optional[str] = None, **kwargs: Any) -> AsyncIterable["_models.ProductResult"]:
"""Define ``filter`` as a query param for all calls. However, the returned next link will also
include the ``filter`` as part of it. Make sure you don't end up duplicating the ``filter``
param in the url sent.

:param filter: OData filter options. Pass in 'foo'.
:type filter: str
:keyword callable cls: A custom type or function that will be passed the direct response
:return: An iterator like instance of either ProductResult or the result of cls(response)
:rtype: ~azure.core.async_paging.AsyncItemPaged[~paging.models.ProductResult]
:raises: ~azure.core.exceptions.HttpResponseError
"""
cls = kwargs.pop("cls", None) # type: ClsType["_models.ProductResult"]
error_map = {401: ClientAuthenticationError, 404: ResourceNotFoundError, 409: ResourceExistsError}
error_map.update(kwargs.pop("error_map", {}))

def prepare_request(next_link=None):
if not next_link:

request = build_duplicate_params_request(
filter=filter,
template_url=self.duplicate_params.metadata["url"],
)
request = _convert_request(request)
request.url = self._client.format_url(request.url)

else:

request = build_duplicate_params_request(
filter=filter,
template_url=next_link,
)
request = _convert_request(request)
request.url = self._client.format_url(request.url)
request.method = "GET"
return request

async def extract_data(pipeline_response):
deserialized = self._deserialize("ProductResult", pipeline_response)
list_of_elem = deserialized.values
if cls:
list_of_elem = cls(list_of_elem)
return deserialized.next_link or None, AsyncList(list_of_elem)

async def get_next(next_link=None):
request = prepare_request(next_link)

pipeline_response = await self._client._pipeline.run( # pylint: disable=protected-access
request, stream=False, **kwargs
)
response = pipeline_response.http_response

if response.status_code not in [200]:
map_error(status_code=response.status_code, response=response, error_map=error_map)
raise HttpResponseError(response=response)

return pipeline_response

return AsyncItemPaged(get_next, extract_data)

duplicate_params.metadata = {"url": "/paging/multiple/duplicateParams/1"} # type: ignore

@distributed_trace
def get_odata_multiple_pages(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,34 @@ def build_get_with_query_params_request(
)


def build_duplicate_params_request(
**kwargs # type: Any
):
# type: (...) -> HttpRequest
filter = kwargs.pop('filter', None) # type: Optional[str]

accept = "application/json"
# Construct URL
_url = kwargs.pop("template_url", "/paging/multiple/duplicateParams/1")

# Construct parameters
_query_parameters = kwargs.pop("params", {}) # type: Dict[str, Any]
if filter is not None:
_query_parameters['$filter'] = _SERIALIZER.query("filter", filter, 'str')

# Construct headers
_header_parameters = kwargs.pop("headers", {}) # type: Dict[str, Any]
_header_parameters['Accept'] = _SERIALIZER.header("accept", accept, 'str')

return HttpRequest(
method="GET",
url=_url,
params=_query_parameters,
headers=_header_parameters,
**kwargs
)


def build_next_operation_with_query_params_request(
**kwargs # type: Any
):
Expand Down Expand Up @@ -978,6 +1006,74 @@ def get_next(next_link=None):

get_with_query_params.metadata = {"url": "/paging/multiple/getWithQueryParams"} # type: ignore

@distributed_trace
def duplicate_params(
self,
filter=None, # type: Optional[str]
**kwargs # type: Any
):
# type: (...) -> Iterable["_models.ProductResult"]
"""Define ``filter`` as a query param for all calls. However, the returned next link will also
include the ``filter`` as part of it. Make sure you don't end up duplicating the ``filter``
param in the url sent.

:param filter: OData filter options. Pass in 'foo'.
:type filter: str
:keyword callable cls: A custom type or function that will be passed the direct response
:return: An iterator like instance of either ProductResult or the result of cls(response)
:rtype: ~azure.core.paging.ItemPaged[~paging.models.ProductResult]
:raises: ~azure.core.exceptions.HttpResponseError
"""
cls = kwargs.pop("cls", None) # type: ClsType["_models.ProductResult"]
error_map = {401: ClientAuthenticationError, 404: ResourceNotFoundError, 409: ResourceExistsError}
error_map.update(kwargs.pop("error_map", {}))

def prepare_request(next_link=None):
if not next_link:

request = build_duplicate_params_request(
filter=filter,
template_url=self.duplicate_params.metadata["url"],
)
request = _convert_request(request)
request.url = self._client.format_url(request.url)

else:

request = build_duplicate_params_request(
filter=filter,
template_url=next_link,
)
request = _convert_request(request)
request.url = self._client.format_url(request.url)
request.method = "GET"
return request

def extract_data(pipeline_response):
deserialized = self._deserialize("ProductResult", pipeline_response)
list_of_elem = deserialized.values
if cls:
list_of_elem = cls(list_of_elem)
return deserialized.next_link or None, iter(list_of_elem)

def get_next(next_link=None):
request = prepare_request(next_link)

pipeline_response = self._client._pipeline.run( # pylint: disable=protected-access
request, stream=False, **kwargs
)
response = pipeline_response.http_response

if response.status_code not in [200]:
map_error(status_code=response.status_code, response=response, error_map=error_map)
raise HttpResponseError(response=response)

return pipeline_response

return ItemPaged(get_next, extract_data)

duplicate_params.metadata = {"url": "/paging/multiple/duplicateParams/1"} # type: ignore

@distributed_trace
def get_odata_multiple_pages(
self,
Expand Down
Loading