Fix AuthorizationException with AWSV4SignerAsyncAuth when the doc ID has special characters. (#848)

* Lifecycle integration tests.

Signed-off-by: dblock <[email protected]>

* Added a test that makes sure the slash is properly encoded.

Signed-off-by: dblock <[email protected]>

* Added more tests for signer and _make_path.

Signed-off-by: Nathalie Jonathan <[email protected]>

* Prevent AIOHttpConnection from encoding the url a second time.

Signed-off-by: Nathalie Jonathan <[email protected]>

---------

Signed-off-by: dblock <[email protected]>
Signed-off-by: Nathalie Jonathan <[email protected]>
Co-authored-by: dblock <[email protected]>
This commit is contained in:
nathaliellenaa
2024-11-27 17:50:22 -05:00
committed by GitHub
co-authored by dblock
parent bf9add4eed
commit b9e48dc847
13 changed files with 445 additions and 9 deletions
@@ -29,6 +29,7 @@ from typing import Any
from unittest import mock
import pytest
import yarl
from multidict import CIMultiDict
from opensearchpy._async._extra_imports import aiohttp # type: ignore
@@ -91,7 +92,7 @@ class TestAsyncHttpConnection:
await c.perform_request("post", "/test")
mock_request.assert_called_with(
"post",
"http://localhost:9200/test",
yarl.URL("http://localhost:9200/test", encoded=True),
data=None,
auth=c._http_auth,
headers={},
@@ -120,7 +121,7 @@ class TestAsyncHttpConnection:
mock_request.assert_called_with(
"post",
"http://localhost:9200/test",
yarl.URL("http://localhost:9200/test", encoded=True),
data=None,
auth=None,
headers={
@@ -30,10 +30,70 @@ from typing import Any
import pytest
from _pytest.mark.structures import MarkDecorator
from opensearchpy.exceptions import RequestError
pytestmark: MarkDecorator = pytest.mark.asyncio
class TestSpecialCharacters:
async def test_index_with_slash(self, async_client: Any) -> None:
index_name = "movies/shmovies"
with pytest.raises(RequestError) as e:
await async_client.indices.create(index=index_name)
assert (
str(e.value)
== "RequestError(400, 'invalid_index_name_exception', 'Invalid index name [movies/shmovies], must not contain the following characters [ , \", *, \\\\, <, |, ,, >, /, ?]')"
)
class TestUnicode:
async def test_indices_lifecycle_english(self, async_client: Any) -> None:
index_name = "movies"
index_create_result = await async_client.indices.create(index=index_name)
assert index_create_result["acknowledged"] is True
assert index_name == index_create_result["index"]
document = {"name": "Solaris", "director": "Andrei Tartakovsky", "year": "2011"}
id = "solaris@2011"
doc_insert_result = await async_client.index(
index=index_name, body=document, id=id, refresh=True
)
assert "created" == doc_insert_result["result"]
assert index_name == doc_insert_result["_index"]
assert id == doc_insert_result["_id"]
doc_delete_result = await async_client.delete(index=index_name, id=id)
assert "deleted" == doc_delete_result["result"]
assert index_name == doc_delete_result["_index"]
assert id == doc_delete_result["_id"]
index_delete_result = await async_client.indices.delete(index=index_name)
assert index_delete_result["acknowledged"] is True
async def test_indices_lifecycle_russian(self, async_client: Any) -> None:
index_name = "кино"
index_create_result = await async_client.indices.create(index=index_name)
assert index_create_result["acknowledged"] is True
assert index_name == index_create_result["index"]
document = {"название": "Солярис", "автор": "Андрей Тарковский", "год": "2011"}
id = "соларис@2011"
doc_insert_result = await async_client.index(
index=index_name, body=document, id=id, refresh=True
)
assert "created" == doc_insert_result["result"]
assert index_name == doc_insert_result["_index"]
assert id == doc_insert_result["_id"]
doc_delete_result = await async_client.delete(index=index_name, id=id)
assert "deleted" == doc_delete_result["result"]
assert index_name == doc_delete_result["_index"]
assert id == doc_delete_result["_id"]
index_delete_result = await async_client.indices.delete(index=index_name)
assert index_delete_result["acknowledged"] is True
async def test_indices_analyze(self, async_client: Any) -> None:
await async_client.indices.analyze(body='{"text": "привет"}')
@@ -8,6 +8,7 @@
# GitHub history for details.
import uuid
from typing import Any, Collection, Dict, Mapping, Optional, Tuple, Union
from unittest.mock import Mock
import pytest
@@ -103,3 +104,75 @@ class TestAsyncSignerWithFrozenCredentials(TestAsyncSigner):
assert "X-Amz-Date" in headers
assert "X-Amz-Security-Token" in headers
assert len(mock_session.get_frozen_credentials.mock_calls) == 1
class TestAsyncSignerWithSpecialCharacters:
def mock_session(self) -> Mock:
access_key = uuid.uuid4().hex
secret_key = uuid.uuid4().hex
token = uuid.uuid4().hex
dummy_session = Mock()
dummy_session.access_key = access_key
dummy_session.secret_key = secret_key
dummy_session.token = token
del dummy_session.get_frozen_credentials
return dummy_session
async def test_aws_signer_async_consitent_url(self) -> None:
region = "us-west-2"
from opensearchpy import AsyncOpenSearch
from opensearchpy.connection.http_async import AsyncHttpConnection
from opensearchpy.helpers.asyncsigner import AWSV4SignerAsyncAuth
# Store URLs for comparison
signed_url = None
sent_url = None
doc_id = "doc_id:with!special*chars%3A"
quoted_doc_id = "doc_id%3Awith%21special*chars%253A"
url = f"https://search-domain.region.es.amazonaws.com:9200/index/_doc/{quoted_doc_id}"
# Create a mock signer class to capture the signed URL
class MockSigner(AWSV4SignerAsyncAuth):
def _sign_request(
self,
method: str,
url: str,
query_string: Optional[str] = None,
body: Optional[Union[str, bytes]] = None,
) -> Dict[str, str]:
nonlocal signed_url
signed_url = url
return {}
# Create a mock connection class to capture the sent URL
class MockConnection(AsyncHttpConnection):
async def perform_request(
self: "MockConnection",
method: str,
url: str,
params: Optional[Mapping[str, Any]] = None,
body: Optional[Any] = None,
timeout: Optional[Union[int, float]] = None,
ignore: Collection[int] = (),
headers: Optional[Mapping[str, str]] = None,
) -> Tuple[int, Mapping[str, str], str]:
nonlocal sent_url
sent_url = f"{self.host}{url}"
return 200, {}, "{}"
auth = MockSigner(self.mock_session(), region)
auth("GET", url)
client = AsyncOpenSearch(
hosts=[{"host": "search-domain.region.es.amazonaws.com"}],
http_auth=auth,
use_ssl=True,
verify_certs=True,
connection_class=MockConnection,
)
await client.index("index", {"test": "data"}, id=doc_id)
assert signed_url == sent_url, "URLs don't match"
+40 -1
View File
@@ -154,9 +154,48 @@ class TestQueryParams(TestCase):
class TestMakePath(TestCase):
def test_handles_unicode(self) -> None:
from urllib.parse import quote
id = "中文"
self.assertEqual(
"/some-index/type/%E4%B8%AD%E6%96%87", _make_path("some-index", "type", id)
_make_path("some-index", "type", quote(id)),
"/some-index/type/%25E4%25B8%25AD%25E6%2596%2587",
)
def test_handles_single_arg(self) -> None:
from urllib.parse import quote
id = "idwith!char"
self.assertEqual(
_make_path("some-index", "type", quote(id)),
"/some-index/type/idwith%2521char",
)
def test_handles_multiple_args(self) -> None:
from urllib.parse import quote
ids = ["id!with@char", "another#id$here"]
quoted_ids = [quote(id) for id in ids]
self.assertEqual(
_make_path("some-index", "type", quoted_ids),
"/some-index/type/id%2521with%2540char,another%2523id%2524here",
)
def test_handles_arrays_of_args(self) -> None:
self.assertEqual(
"/index1,index2/type1,type2/doc1,doc2",
_make_path(
("index1", "index2"), ["type1", "type2"], tuple(["doc1", "doc2"])
),
)
from urllib.parse import quote
ids = [quote("$id!1"), quote("id*@2"), quote("#id3#")]
self.assertEqual(
_make_path("some-index", ids, "type"),
"/some-index/%2524id%25211,id%252A%25402,%2523id3%2523/type",
)
@@ -513,6 +513,65 @@ class TestRequestsHttpConnection(TestCase):
("GET", "http://localhost/?key1=value1&key2=value2", None),
)
def test_aws_signer_consitent_url(self) -> None:
region = "us-west-2"
from typing import Any, Collection, Mapping, Optional, Union
from opensearchpy import OpenSearch
from opensearchpy.helpers.signer import RequestsAWSV4SignerAuth
# Store URLs for comparison
signed_url = None
sent_url = None
doc_id = "doc_id:with!special*chars%3A"
quoted_doc_id = "doc_id%3Awith%21special*chars%253A"
url = f"https://search-domain.region.es.amazonaws.com:9200/index/_doc/{quoted_doc_id}"
# Create a mock signer class to capture the signed URL
class MockSigner(RequestsAWSV4SignerAuth):
def __call__(self, prepared_request): # type: ignore
nonlocal signed_url
if isinstance(prepared_request, str):
signed_url = prepared_request
else:
signed_url = prepared_request.url
return prepared_request
# Create a mock connection class to capture the sent URL
class MockConnection(RequestsHttpConnection):
def perform_request( # type: ignore
self,
method: str,
url: str,
params: Optional[Mapping[str, Any]] = None,
body: Optional[bytes] = None,
timeout: Optional[Union[int, float]] = None,
allow_redirects: Optional[bool] = True,
ignore: Collection[int] = (),
headers: Optional[Mapping[str, str]] = None,
) -> Any:
nonlocal sent_url
sent_url = f"{self.host}{url}"
return 200, {}, "{}"
auth = MockSigner(self.mock_session(), region)
client = OpenSearch(
hosts=[{"host": "search-domain.region.es.amazonaws.com"}],
http_auth=auth(url),
use_ssl=True,
verify_certs=True,
connection_class=MockConnection,
)
client.index("index", {"test": "data"}, id=doc_id)
self.assertEqual(
signed_url,
sent_url,
"URLs don't match",
)
class TestRequestsConnectionRedirect(TestCase):
server1: TestHTTPServer
@@ -25,10 +25,72 @@
# under the License.
import pytest
from opensearchpy.exceptions import RequestError
from . import OpenSearchTestCase
class TestSpecialCharacters(OpenSearchTestCase):
def test_index_with_slash(self) -> None:
index_name = "movies/shmovies"
with pytest.raises(RequestError) as e:
self.client.indices.create(index=index_name)
self.assertEqual(
str(e.value),
"RequestError(400, 'invalid_index_name_exception', 'Invalid index name [movies/shmovies], must not contain the following characters [ , \", *, \\\\, <, |, ,, >, /, ?]')",
)
class TestUnicode(OpenSearchTestCase):
def test_indices_lifecycle_english(self) -> None:
index_name = "movies"
index_create_result = self.client.indices.create(index=index_name)
self.assertTrue(index_create_result["acknowledged"])
self.assertEqual(index_name, index_create_result["index"])
document = {"name": "Solaris", "director": "Andrei Tartakovsky", "year": "2011"}
id = "solaris@2011"
doc_insert_result = self.client.index(
index=index_name, body=document, id=id, refresh=True
)
self.assertEqual("created", doc_insert_result["result"])
self.assertEqual(index_name, doc_insert_result["_index"])
self.assertEqual(id, doc_insert_result["_id"])
doc_delete_result = self.client.delete(index=index_name, id=id)
self.assertEqual("deleted", doc_delete_result["result"])
self.assertEqual(index_name, doc_delete_result["_index"])
self.assertEqual(id, doc_delete_result["_id"])
index_delete_result = self.client.indices.delete(index=index_name)
self.assertTrue(index_delete_result["acknowledged"])
def test_indices_lifecycle_russian(self) -> None:
index_name = "кино"
index_create_result = self.client.indices.create(index=index_name)
self.assertTrue(index_create_result["acknowledged"])
self.assertEqual(index_name, index_create_result["index"])
document = {"название": "Солярис", "автор": "Андрей Тарковский", "год": "2011"}
id = "соларис@2011"
doc_insert_result = self.client.index(
index=index_name, body=document, id=id, refresh=True
)
self.assertEqual("created", doc_insert_result["result"])
self.assertEqual(index_name, doc_insert_result["_index"])
self.assertEqual(id, doc_insert_result["_id"])
doc_delete_result = self.client.delete(index=index_name, id=id)
self.assertEqual("deleted", doc_delete_result["result"])
self.assertEqual(index_name, doc_delete_result["_index"])
self.assertEqual(id, doc_delete_result["_id"])
index_delete_result = self.client.indices.delete(index=index_name)
self.assertTrue(index_delete_result["acknowledged"])
def test_indices_analyze(self) -> None:
self.client.indices.analyze(body='{"text": "привет"}')