Rename module to opensearchpy
To avoid conflict with an existing package by name 'opensearch' being present Signed-off-by: Rushi Agrawal <[email protected]>
This commit is contained in:
@@ -0,0 +1,25 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
@@ -0,0 +1,421 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
|
||||
import gzip
|
||||
import io
|
||||
import json
|
||||
import ssl
|
||||
import warnings
|
||||
from platform import python_version
|
||||
|
||||
import aiohttp
|
||||
import pytest
|
||||
from mock import patch
|
||||
from multidict import CIMultiDict
|
||||
|
||||
from opensearchpy import AIOHttpConnection, __versionstr__
|
||||
from opensearchpy.compat import reraise_exceptions
|
||||
from opensearchpy.exceptions import ConnectionError
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
def gzip_decompress(data):
|
||||
buf = gzip.GzipFile(fileobj=io.BytesIO(data), mode="rb")
|
||||
return buf.read()
|
||||
|
||||
|
||||
class TestAIOHttpConnection:
|
||||
async def _get_mock_connection(self, connection_params={}, response_body=b"{}"):
|
||||
con = AIOHttpConnection(**connection_params)
|
||||
await con._create_aiohttp_session()
|
||||
|
||||
def _dummy_request(*args, **kwargs):
|
||||
class DummyResponse:
|
||||
async def __aenter__(self, *_, **__):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_, **__):
|
||||
pass
|
||||
|
||||
async def text(self):
|
||||
return response_body.decode("utf-8", "surrogatepass")
|
||||
|
||||
dummy_response = DummyResponse()
|
||||
dummy_response.headers = CIMultiDict()
|
||||
dummy_response.status = 200
|
||||
_dummy_request.call_args = (args, kwargs)
|
||||
return dummy_response
|
||||
|
||||
con.session.request = _dummy_request
|
||||
return con
|
||||
|
||||
async def test_ssl_context(self):
|
||||
try:
|
||||
context = ssl.create_default_context()
|
||||
except AttributeError:
|
||||
# if create_default_context raises an AttributeError Exception
|
||||
# it means SSLContext is not available for that version of python
|
||||
# and we should skip this test.
|
||||
pytest.skip(
|
||||
"Test test_ssl_context is skipped cause SSLContext is not available for this version of Python"
|
||||
)
|
||||
|
||||
con = AIOHttpConnection(use_ssl=True, ssl_context=context)
|
||||
await con._create_aiohttp_session()
|
||||
assert con.use_ssl
|
||||
assert con.session.connector._ssl == context
|
||||
|
||||
def test_opaque_id(self):
|
||||
con = AIOHttpConnection(opaque_id="app-1")
|
||||
assert con.headers["x-opaque-id"] == "app-1"
|
||||
|
||||
def test_http_cloud_id(self):
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng=="
|
||||
)
|
||||
assert con.use_ssl
|
||||
assert (
|
||||
con.host
|
||||
== "https://4fa8821e75634032bed1cf22110e2f97.us-east-1.aws.found.io"
|
||||
)
|
||||
assert con.port is None
|
||||
assert con.hostname == "4fa8821e75634032bed1cf22110e2f97.us-east-1.aws.found.io"
|
||||
assert con.http_compress
|
||||
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
port=9243,
|
||||
)
|
||||
assert (
|
||||
con.host
|
||||
== "https://4fa8821e75634032bed1cf22110e2f97.us-east-1.aws.found.io:9243"
|
||||
)
|
||||
assert con.port == 9243
|
||||
assert con.hostname == "4fa8821e75634032bed1cf22110e2f97.us-east-1.aws.found.io"
|
||||
|
||||
def test_api_key_auth(self):
|
||||
# test with tuple
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
api_key=("elastic", "changeme1"),
|
||||
)
|
||||
assert con.headers["authorization"] == "ApiKey ZWxhc3RpYzpjaGFuZ2VtZTE="
|
||||
assert (
|
||||
con.host
|
||||
== "https://4fa8821e75634032bed1cf22110e2f97.us-east-1.aws.found.io"
|
||||
)
|
||||
|
||||
# test with base64 encoded string
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
api_key="ZWxhc3RpYzpjaGFuZ2VtZTI=",
|
||||
)
|
||||
assert con.headers["authorization"] == "ApiKey ZWxhc3RpYzpjaGFuZ2VtZTI="
|
||||
assert (
|
||||
con.host
|
||||
== "https://4fa8821e75634032bed1cf22110e2f97.us-east-1.aws.found.io"
|
||||
)
|
||||
|
||||
async def test_no_http_compression(self):
|
||||
con = await self._get_mock_connection()
|
||||
assert not con.http_compress
|
||||
assert "accept-encoding" not in con.headers
|
||||
|
||||
await con.perform_request("GET", "/")
|
||||
|
||||
_, kwargs = con.session.request.call_args
|
||||
|
||||
assert not kwargs["data"]
|
||||
assert "accept-encoding" not in kwargs["headers"]
|
||||
assert "content-encoding" not in kwargs["headers"]
|
||||
|
||||
async def test_http_compression(self):
|
||||
con = await self._get_mock_connection({"http_compress": True})
|
||||
assert con.http_compress
|
||||
assert con.headers["accept-encoding"] == "gzip,deflate"
|
||||
|
||||
# 'content-encoding' shouldn't be set at a connection level.
|
||||
# Should be applied only if the request is sent with a body.
|
||||
assert "content-encoding" not in con.headers
|
||||
|
||||
await con.perform_request("GET", "/", body=b"{}")
|
||||
|
||||
_, kwargs = con.session.request.call_args
|
||||
|
||||
assert gzip_decompress(kwargs["data"]) == b"{}"
|
||||
assert kwargs["headers"]["accept-encoding"] == "gzip,deflate"
|
||||
assert kwargs["headers"]["content-encoding"] == "gzip"
|
||||
|
||||
await con.perform_request("GET", "/")
|
||||
|
||||
_, kwargs = con.session.request.call_args
|
||||
|
||||
assert not kwargs["data"]
|
||||
assert kwargs["headers"]["accept-encoding"] == "gzip,deflate"
|
||||
assert "content-encoding" not in kwargs["headers"]
|
||||
|
||||
def test_cloud_id_http_compress_override(self):
|
||||
# 'http_compress' will be 'True' by default for connections with
|
||||
# 'cloud_id' set but should prioritize user-defined values.
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
)
|
||||
assert con.http_compress is True
|
||||
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
http_compress=False,
|
||||
)
|
||||
assert con.http_compress is False
|
||||
|
||||
con = AIOHttpConnection(
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
http_compress=True,
|
||||
)
|
||||
assert con.http_compress is True
|
||||
|
||||
async def test_url_prefix(self):
|
||||
con = await self._get_mock_connection(
|
||||
connection_params={"url_prefix": "/_search/"}
|
||||
)
|
||||
assert con.url_prefix == "/_search"
|
||||
|
||||
await con.perform_request("GET", "/")
|
||||
|
||||
# Need to convert the yarl URL to a string to compare.
|
||||
method, yarl_url = con.session.request.call_args[0]
|
||||
assert method == "GET" and str(yarl_url) == "http://localhost:9200/_search/"
|
||||
|
||||
def test_default_user_agent(self):
|
||||
con = AIOHttpConnection()
|
||||
assert con._get_default_user_agent() == "opensearch-py/%s (Python %s)" % (
|
||||
__versionstr__,
|
||||
python_version(),
|
||||
)
|
||||
|
||||
def test_timeout_set(self):
|
||||
con = AIOHttpConnection(timeout=42)
|
||||
assert 42 == con.timeout
|
||||
|
||||
def test_keep_alive_is_on_by_default(self):
|
||||
con = AIOHttpConnection()
|
||||
assert {
|
||||
"connection": "keep-alive",
|
||||
"content-type": "application/json",
|
||||
"user-agent": con._get_default_user_agent(),
|
||||
} == con.headers
|
||||
|
||||
def test_http_auth(self):
|
||||
con = AIOHttpConnection(http_auth="username:secret")
|
||||
assert {
|
||||
"authorization": "Basic dXNlcm5hbWU6c2VjcmV0",
|
||||
"connection": "keep-alive",
|
||||
"content-type": "application/json",
|
||||
"user-agent": con._get_default_user_agent(),
|
||||
} == con.headers
|
||||
|
||||
def test_http_auth_tuple(self):
|
||||
con = AIOHttpConnection(http_auth=("username", "secret"))
|
||||
assert {
|
||||
"authorization": "Basic dXNlcm5hbWU6c2VjcmV0",
|
||||
"content-type": "application/json",
|
||||
"connection": "keep-alive",
|
||||
"user-agent": con._get_default_user_agent(),
|
||||
} == con.headers
|
||||
|
||||
def test_http_auth_list(self):
|
||||
con = AIOHttpConnection(http_auth=["username", "secret"])
|
||||
assert {
|
||||
"authorization": "Basic dXNlcm5hbWU6c2VjcmV0",
|
||||
"content-type": "application/json",
|
||||
"connection": "keep-alive",
|
||||
"user-agent": con._get_default_user_agent(),
|
||||
} == con.headers
|
||||
|
||||
def test_uses_https_if_verify_certs_is_off(self):
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
con = AIOHttpConnection(use_ssl=True, verify_certs=False)
|
||||
assert 1 == len(w)
|
||||
assert (
|
||||
"Connecting to https://localhost:9200 using SSL with verify_certs=False is insecure."
|
||||
== str(w[0].message)
|
||||
)
|
||||
|
||||
assert con.use_ssl
|
||||
assert con.scheme == "https"
|
||||
assert con.host == "https://localhost:9200"
|
||||
|
||||
async def test_nowarn_when_test_uses_https_if_verify_certs_is_off(self):
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
con = AIOHttpConnection(
|
||||
use_ssl=True, verify_certs=False, ssl_show_warn=False
|
||||
)
|
||||
await con._create_aiohttp_session()
|
||||
assert w == []
|
||||
|
||||
assert isinstance(con.session, aiohttp.ClientSession)
|
||||
|
||||
def test_doesnt_use_https_if_not_specified(self):
|
||||
con = AIOHttpConnection()
|
||||
assert not con.use_ssl
|
||||
|
||||
def test_no_warning_when_using_ssl_context(self):
|
||||
ctx = ssl.create_default_context()
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
AIOHttpConnection(ssl_context=ctx)
|
||||
assert w == [], str([x.message for x in w])
|
||||
|
||||
def test_warns_if_using_non_default_ssl_kwargs_with_ssl_context(self):
|
||||
for kwargs in (
|
||||
{"ssl_show_warn": False},
|
||||
{"ssl_show_warn": True},
|
||||
{"verify_certs": True},
|
||||
{"verify_certs": False},
|
||||
{"ca_certs": "/path/to/certs"},
|
||||
{"ssl_show_warn": True, "ca_certs": "/path/to/certs"},
|
||||
):
|
||||
kwargs["ssl_context"] = ssl.create_default_context()
|
||||
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
warnings.simplefilter("always")
|
||||
|
||||
AIOHttpConnection(**kwargs)
|
||||
|
||||
assert 1 == len(w)
|
||||
assert (
|
||||
"When using `ssl_context`, all other SSL related kwargs are ignored"
|
||||
== str(w[0].message)
|
||||
)
|
||||
|
||||
@patch("opensearchpy.connection.base.logger")
|
||||
async def test_uncompressed_body_logged(self, logger):
|
||||
con = await self._get_mock_connection(connection_params={"http_compress": True})
|
||||
await con.perform_request("GET", "/", body=b'{"example": "body"}')
|
||||
|
||||
assert 2 == logger.debug.call_count
|
||||
req, resp = logger.debug.call_args_list
|
||||
|
||||
assert '> {"example": "body"}' == req[0][0] % req[0][1:]
|
||||
assert "< {}" == resp[0][0] % resp[0][1:]
|
||||
|
||||
async def test_surrogatepass_into_bytes(self):
|
||||
buf = b"\xe4\xbd\xa0\xe5\xa5\xbd\xed\xa9\xaa"
|
||||
con = await self._get_mock_connection(response_body=buf)
|
||||
status, headers, data = await con.perform_request("GET", "/")
|
||||
assert u"你好\uda6a" == data
|
||||
|
||||
@pytest.mark.parametrize("exception_cls", reraise_exceptions)
|
||||
async def test_recursion_error_reraised(self, exception_cls):
|
||||
conn = AIOHttpConnection()
|
||||
|
||||
def request_raise(*_, **__):
|
||||
raise exception_cls("Wasn't modified!")
|
||||
|
||||
await conn._create_aiohttp_session()
|
||||
conn.session.request = request_raise
|
||||
|
||||
with pytest.raises(exception_cls) as e:
|
||||
await conn.perform_request("GET", "/")
|
||||
assert str(e.value) == "Wasn't modified!"
|
||||
|
||||
|
||||
class TestConnectionHttpbin:
|
||||
"""Tests the HTTP connection implementations against a live server E2E"""
|
||||
|
||||
async def httpbin_anything(self, conn, **kwargs):
|
||||
status, headers, data = await conn.perform_request("GET", "/anything", **kwargs)
|
||||
data = json.loads(data)
|
||||
data["headers"].pop(
|
||||
"X-Amzn-Trace-Id", None
|
||||
) # Remove this header as it's put there by AWS.
|
||||
return (status, data)
|
||||
|
||||
async def test_aiohttp_connection(self):
|
||||
# Defaults
|
||||
conn = AIOHttpConnection("httpbin.org", port=443, use_ssl=True)
|
||||
user_agent = conn._get_default_user_agent()
|
||||
status, data = await self.httpbin_anything(conn)
|
||||
assert status == 200
|
||||
assert data["method"] == "GET"
|
||||
assert data["headers"] == {
|
||||
"Content-Type": "application/json",
|
||||
"Host": "httpbin.org",
|
||||
"User-Agent": user_agent,
|
||||
}
|
||||
|
||||
# http_compress=False
|
||||
conn = AIOHttpConnection(
|
||||
"httpbin.org", port=443, use_ssl=True, http_compress=False
|
||||
)
|
||||
status, data = await self.httpbin_anything(conn)
|
||||
assert status == 200
|
||||
assert data["method"] == "GET"
|
||||
assert data["headers"] == {
|
||||
"Content-Type": "application/json",
|
||||
"Host": "httpbin.org",
|
||||
"User-Agent": user_agent,
|
||||
}
|
||||
|
||||
# http_compress=True
|
||||
conn = AIOHttpConnection(
|
||||
"httpbin.org", port=443, use_ssl=True, http_compress=True
|
||||
)
|
||||
status, data = await self.httpbin_anything(conn)
|
||||
assert status == 200
|
||||
assert data["headers"] == {
|
||||
"Accept-Encoding": "gzip,deflate",
|
||||
"Content-Type": "application/json",
|
||||
"Host": "httpbin.org",
|
||||
"User-Agent": user_agent,
|
||||
}
|
||||
|
||||
# Headers
|
||||
conn = AIOHttpConnection(
|
||||
"httpbin.org",
|
||||
port=443,
|
||||
use_ssl=True,
|
||||
http_compress=True,
|
||||
headers={"header1": "value1"},
|
||||
)
|
||||
status, data = await self.httpbin_anything(
|
||||
conn, headers={"header2": "value2", "header1": "override!"}
|
||||
)
|
||||
assert status == 200
|
||||
assert data["headers"] == {
|
||||
"Accept-Encoding": "gzip,deflate",
|
||||
"Content-Type": "application/json",
|
||||
"Host": "httpbin.org",
|
||||
"Header1": "override!",
|
||||
"Header2": "value2",
|
||||
"User-Agent": user_agent,
|
||||
}
|
||||
|
||||
async def test_aiohttp_connection_error(self):
|
||||
conn = AIOHttpConnection("not.a.host.name")
|
||||
with pytest.raises(ConnectionError):
|
||||
await conn.perform_request("GET", "/")
|
||||
@@ -0,0 +1,25 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
import opensearchpy
|
||||
from opensearchpy.helpers.test import CA_CERTS, OPENSEARCH_URL
|
||||
|
||||
from ...utils import wipe_cluster
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
async def async_client():
|
||||
client = None
|
||||
try:
|
||||
if not hasattr(opensearchpy, "AsyncOpenSearch"):
|
||||
pytest.skip("test requires 'AsyncOpenSearch'")
|
||||
|
||||
kw = {"timeout": 3, "ca_certs": CA_CERTS}
|
||||
client = opensearchpy.AsyncOpenSearch(OPENSEARCH_URL, **kw)
|
||||
|
||||
# wait for yellow status
|
||||
for _ in range(100):
|
||||
try:
|
||||
await client.cluster.health(wait_for_status="yellow")
|
||||
break
|
||||
except ConnectionError:
|
||||
await asyncio.sleep(0.1)
|
||||
else:
|
||||
# timeout
|
||||
pytest.skip("OpenSearch failed to start.")
|
||||
|
||||
yield client
|
||||
|
||||
finally:
|
||||
if client:
|
||||
wipe_cluster(client)
|
||||
await client.close()
|
||||
@@ -0,0 +1,66 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
|
||||
from __future__ import unicode_literals
|
||||
|
||||
import pytest
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
class TestUnicode:
|
||||
async def test_indices_analyze(self, async_client):
|
||||
await async_client.indices.analyze(body='{"text": "привет"}')
|
||||
|
||||
|
||||
class TestBulk:
|
||||
async def test_bulk_works_with_string_body(self, async_client):
|
||||
docs = '{ "index" : { "_index" : "bulk_test_index", "_id" : "1" } }\n{"answer": 42}'
|
||||
response = await async_client.bulk(body=docs)
|
||||
|
||||
assert response["errors"] is False
|
||||
assert len(response["items"]) == 1
|
||||
|
||||
async def test_bulk_works_with_bytestring_body(self, async_client):
|
||||
docs = b'{ "index" : { "_index" : "bulk_test_index", "_id" : "2" } }\n{"answer": 42}'
|
||||
response = await async_client.bulk(body=docs)
|
||||
|
||||
assert response["errors"] is False
|
||||
assert len(response["items"]) == 1
|
||||
|
||||
|
||||
class TestYarlMissing:
|
||||
async def test_aiohttp_connection_works_without_yarl(
|
||||
self, async_client, monkeypatch
|
||||
):
|
||||
# This is a defensive test case for if aiohttp suddenly stops using yarl.
|
||||
from opensearchpy._async import http_aiohttp
|
||||
|
||||
monkeypatch.setattr(http_aiohttp, "yarl", False)
|
||||
|
||||
resp = await async_client.info(pretty=True)
|
||||
assert isinstance(resp, dict)
|
||||
@@ -0,0 +1,901 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
|
||||
# Licensed to Elasticsearch B.V.under one or more agreements.
|
||||
# Elasticsearch B.V.licenses this file to you under the Apache 2.0 License.
|
||||
# See the LICENSE file in the project root for more information
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from mock import MagicMock, patch
|
||||
|
||||
from opensearchpy import TransportError, helpers
|
||||
from opensearchpy.helpers import ScanError
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
class AsyncMock(MagicMock):
|
||||
async def __call__(self, *args, **kwargs):
|
||||
return super(AsyncMock, self).__call__(*args, **kwargs)
|
||||
|
||||
def __await__(self):
|
||||
return self().__await__()
|
||||
|
||||
|
||||
class FailingBulkClient(object):
|
||||
def __init__(
|
||||
self, client, fail_at=(2,), fail_with=TransportError(599, "Error!", {})
|
||||
):
|
||||
self.client = client
|
||||
self._called = 0
|
||||
self._fail_at = fail_at
|
||||
self.transport = client.transport
|
||||
self._fail_with = fail_with
|
||||
|
||||
async def bulk(self, *args, **kwargs):
|
||||
self._called += 1
|
||||
if self._called in self._fail_at:
|
||||
raise self._fail_with
|
||||
return await self.client.bulk(*args, **kwargs)
|
||||
|
||||
|
||||
class TestStreamingBulk(object):
|
||||
async def test_actions_remain_unchanged(self, async_client):
|
||||
actions = [{"_id": 1}, {"_id": 2}]
|
||||
async for ok, item in helpers.async_streaming_bulk(
|
||||
async_client, actions, index="test-index"
|
||||
):
|
||||
assert ok
|
||||
assert [{"_id": 1}, {"_id": 2}] == actions
|
||||
|
||||
async def test_all_documents_get_inserted(self, async_client):
|
||||
docs = [{"answer": x, "_id": x} for x in range(100)]
|
||||
async for ok, item in helpers.async_streaming_bulk(
|
||||
async_client, docs, index="test-index", refresh=True
|
||||
):
|
||||
assert ok
|
||||
|
||||
assert 100 == (await async_client.count(index="test-index"))["count"]
|
||||
assert {"answer": 42} == (await async_client.get(index="test-index", id=42))[
|
||||
"_source"
|
||||
]
|
||||
|
||||
async def test_documents_data_types(self, async_client):
|
||||
async def async_gen():
|
||||
for x in range(100):
|
||||
await asyncio.sleep(0)
|
||||
yield {"answer": x, "_id": x}
|
||||
|
||||
def sync_gen():
|
||||
for x in range(100):
|
||||
yield {"answer": x, "_id": x}
|
||||
|
||||
async for ok, item in helpers.async_streaming_bulk(
|
||||
async_client, async_gen(), index="test-index", refresh=True
|
||||
):
|
||||
assert ok
|
||||
|
||||
assert 100 == (await async_client.count(index="test-index"))["count"]
|
||||
assert {"answer": 42} == (await async_client.get(index="test-index", id=42))[
|
||||
"_source"
|
||||
]
|
||||
|
||||
await async_client.delete_by_query(
|
||||
index="test-index", body={"query": {"match_all": {}}}
|
||||
)
|
||||
|
||||
async for ok, item in helpers.async_streaming_bulk(
|
||||
async_client, sync_gen(), index="test-index", refresh=True
|
||||
):
|
||||
assert ok
|
||||
|
||||
assert 100 == (await async_client.count(index="test-index"))["count"]
|
||||
assert {"answer": 42} == (await async_client.get(index="test-index", id=42))[
|
||||
"_source"
|
||||
]
|
||||
|
||||
async def test_all_errors_from_chunk_are_raised_on_failure(self, async_client):
|
||||
await async_client.indices.create(
|
||||
"i",
|
||||
{
|
||||
"mappings": {"properties": {"a": {"type": "integer"}}},
|
||||
"settings": {"number_of_shards": 1, "number_of_replicas": 0},
|
||||
},
|
||||
)
|
||||
await async_client.cluster.health(wait_for_status="yellow")
|
||||
|
||||
try:
|
||||
async for ok, item in helpers.async_streaming_bulk(
|
||||
async_client, [{"a": "b"}, {"a": "c"}], index="i", raise_on_error=True
|
||||
):
|
||||
assert ok
|
||||
except helpers.BulkIndexError as e:
|
||||
assert 2 == len(e.errors)
|
||||
else:
|
||||
assert False, "exception should have been raised"
|
||||
|
||||
async def test_different_op_types(self, async_client):
|
||||
await async_client.index(index="i", id=45, body={})
|
||||
await async_client.index(index="i", id=42, body={})
|
||||
docs = [
|
||||
{"_index": "i", "_id": 47, "f": "v"},
|
||||
{"_op_type": "delete", "_index": "i", "_id": 45},
|
||||
{"_op_type": "update", "_index": "i", "_id": 42, "doc": {"answer": 42}},
|
||||
]
|
||||
async for ok, item in helpers.async_streaming_bulk(async_client, docs):
|
||||
assert ok
|
||||
|
||||
assert not await async_client.exists(index="i", id=45)
|
||||
assert {"answer": 42} == (await async_client.get(index="i", id=42))["_source"]
|
||||
assert {"f": "v"} == (await async_client.get(index="i", id=47))["_source"]
|
||||
|
||||
async def test_transport_error_can_becaught(self, async_client):
|
||||
failing_client = FailingBulkClient(async_client)
|
||||
docs = [
|
||||
{"_index": "i", "_id": 47, "f": "v"},
|
||||
{"_index": "i", "_id": 45, "f": "v"},
|
||||
{"_index": "i", "_id": 42, "f": "v"},
|
||||
]
|
||||
|
||||
results = [
|
||||
x
|
||||
async for x in helpers.async_streaming_bulk(
|
||||
failing_client,
|
||||
docs,
|
||||
raise_on_exception=False,
|
||||
raise_on_error=False,
|
||||
chunk_size=1,
|
||||
)
|
||||
]
|
||||
assert 3 == len(results)
|
||||
assert [True, False, True] == [r[0] for r in results]
|
||||
|
||||
exc = results[1][1]["index"].pop("exception")
|
||||
assert isinstance(exc, TransportError)
|
||||
assert 599 == exc.status_code
|
||||
assert {
|
||||
"index": {
|
||||
"_index": "i",
|
||||
"_id": 45,
|
||||
"data": {"f": "v"},
|
||||
"error": "TransportError(599, 'Error!')",
|
||||
"status": 599,
|
||||
}
|
||||
} == results[1][1]
|
||||
|
||||
async def test_rejected_documents_are_retried(self, async_client):
|
||||
failing_client = FailingBulkClient(
|
||||
async_client, fail_with=TransportError(429, "Rejected!", {})
|
||||
)
|
||||
docs = [
|
||||
{"_index": "i", "_id": 47, "f": "v"},
|
||||
{"_index": "i", "_id": 45, "f": "v"},
|
||||
{"_index": "i", "_id": 42, "f": "v"},
|
||||
]
|
||||
results = [
|
||||
x
|
||||
async for x in helpers.async_streaming_bulk(
|
||||
failing_client,
|
||||
docs,
|
||||
raise_on_exception=False,
|
||||
raise_on_error=False,
|
||||
chunk_size=1,
|
||||
max_retries=1,
|
||||
initial_backoff=0,
|
||||
)
|
||||
]
|
||||
assert 3 == len(results)
|
||||
assert [True, True, True] == [r[0] for r in results]
|
||||
await async_client.indices.refresh(index="i")
|
||||
res = await async_client.search(index="i")
|
||||
assert {"value": 3, "relation": "eq"} == res["hits"]["total"]
|
||||
assert 4 == failing_client._called
|
||||
|
||||
async def test_rejected_documents_are_retried_at_most_max_retries_times(
|
||||
self, async_client
|
||||
):
|
||||
failing_client = FailingBulkClient(
|
||||
async_client, fail_at=(1, 2), fail_with=TransportError(429, "Rejected!", {})
|
||||
)
|
||||
|
||||
docs = [
|
||||
{"_index": "i", "_id": 47, "f": "v"},
|
||||
{"_index": "i", "_id": 45, "f": "v"},
|
||||
{"_index": "i", "_id": 42, "f": "v"},
|
||||
]
|
||||
results = [
|
||||
x
|
||||
async for x in helpers.async_streaming_bulk(
|
||||
failing_client,
|
||||
docs,
|
||||
raise_on_exception=False,
|
||||
raise_on_error=False,
|
||||
chunk_size=1,
|
||||
max_retries=1,
|
||||
initial_backoff=0,
|
||||
)
|
||||
]
|
||||
assert 3 == len(results)
|
||||
assert [False, True, True] == [r[0] for r in results]
|
||||
await async_client.indices.refresh(index="i")
|
||||
res = await async_client.search(index="i")
|
||||
assert {"value": 2, "relation": "eq"} == res["hits"]["total"]
|
||||
assert 4 == failing_client._called
|
||||
|
||||
async def test_transport_error_is_raised_with_max_retries(self, async_client):
|
||||
failing_client = FailingBulkClient(
|
||||
async_client,
|
||||
fail_at=(1, 2, 3, 4),
|
||||
fail_with=TransportError(429, "Rejected!", {}),
|
||||
)
|
||||
|
||||
async def streaming_bulk():
|
||||
results = [
|
||||
x
|
||||
async for x in helpers.async_streaming_bulk(
|
||||
failing_client,
|
||||
[{"a": 42}, {"a": 39}],
|
||||
raise_on_exception=True,
|
||||
max_retries=3,
|
||||
initial_backoff=0,
|
||||
)
|
||||
]
|
||||
return results
|
||||
|
||||
with pytest.raises(TransportError):
|
||||
await streaming_bulk()
|
||||
assert 4 == failing_client._called
|
||||
|
||||
|
||||
class TestBulk(object):
|
||||
async def test_bulk_works_with_single_item(self, async_client):
|
||||
docs = [{"answer": 42, "_id": 1}]
|
||||
success, failed = await helpers.async_bulk(
|
||||
async_client, docs, index="test-index", refresh=True
|
||||
)
|
||||
|
||||
assert 1 == success
|
||||
assert not failed
|
||||
assert 1 == (await async_client.count(index="test-index"))["count"]
|
||||
assert {"answer": 42} == (await async_client.get(index="test-index", id=1))[
|
||||
"_source"
|
||||
]
|
||||
|
||||
async def test_all_documents_get_inserted(self, async_client):
|
||||
docs = [{"answer": x, "_id": x} for x in range(100)]
|
||||
success, failed = await helpers.async_bulk(
|
||||
async_client, docs, index="test-index", refresh=True
|
||||
)
|
||||
|
||||
assert 100 == success
|
||||
assert not failed
|
||||
assert 100 == (await async_client.count(index="test-index"))["count"]
|
||||
assert {"answer": 42} == (await async_client.get(index="test-index", id=42))[
|
||||
"_source"
|
||||
]
|
||||
|
||||
async def test_stats_only_reports_numbers(self, async_client):
|
||||
docs = [{"answer": x} for x in range(100)]
|
||||
success, failed = await helpers.async_bulk(
|
||||
async_client, docs, index="test-index", refresh=True, stats_only=True
|
||||
)
|
||||
|
||||
assert 100 == success
|
||||
assert 0 == failed
|
||||
assert 100 == (await async_client.count(index="test-index"))["count"]
|
||||
|
||||
async def test_errors_are_reported_correctly(self, async_client):
|
||||
await async_client.indices.create(
|
||||
"i",
|
||||
{
|
||||
"mappings": {"properties": {"a": {"type": "integer"}}},
|
||||
"settings": {"number_of_shards": 1, "number_of_replicas": 0},
|
||||
},
|
||||
)
|
||||
await async_client.cluster.health(wait_for_status="yellow")
|
||||
|
||||
success, failed = await helpers.async_bulk(
|
||||
async_client,
|
||||
[{"a": 42}, {"a": "c", "_id": 42}],
|
||||
index="i",
|
||||
raise_on_error=False,
|
||||
)
|
||||
assert 1 == success
|
||||
assert 1 == len(failed)
|
||||
error = failed[0]
|
||||
assert "42" == error["index"]["_id"]
|
||||
assert "i" == error["index"]["_index"]
|
||||
print(error["index"]["error"])
|
||||
assert "MapperParsingException" in repr(
|
||||
error["index"]["error"]
|
||||
) or "mapper_parsing_exception" in repr(error["index"]["error"])
|
||||
|
||||
async def test_error_is_raised(self, async_client):
|
||||
await async_client.indices.create(
|
||||
"i",
|
||||
{
|
||||
"mappings": {"properties": {"a": {"type": "integer"}}},
|
||||
"settings": {"number_of_shards": 1, "number_of_replicas": 0},
|
||||
},
|
||||
)
|
||||
await async_client.cluster.health(wait_for_status="yellow")
|
||||
|
||||
with pytest.raises(helpers.BulkIndexError):
|
||||
await helpers.async_bulk(async_client, [{"a": 42}, {"a": "c"}], index="i")
|
||||
|
||||
async def test_ignore_error_if_raised(self, async_client):
|
||||
# ignore the status code 400 in tuple
|
||||
await helpers.async_bulk(
|
||||
async_client, [{"a": 42}, {"a": "c"}], index="i", ignore_status=(400,)
|
||||
)
|
||||
|
||||
# ignore the status code 400 in list
|
||||
await helpers.async_bulk(
|
||||
async_client,
|
||||
[{"a": 42}, {"a": "c"}],
|
||||
index="i",
|
||||
ignore_status=[
|
||||
400,
|
||||
],
|
||||
)
|
||||
|
||||
# ignore the status code 400
|
||||
await helpers.async_bulk(
|
||||
async_client, [{"a": 42}, {"a": "c"}], index="i", ignore_status=400
|
||||
)
|
||||
|
||||
# ignore only the status code in the `ignore_status` argument
|
||||
with pytest.raises(helpers.BulkIndexError):
|
||||
await helpers.async_bulk(
|
||||
async_client, [{"a": 42}, {"a": "c"}], index="i", ignore_status=(444,)
|
||||
)
|
||||
|
||||
# ignore transport error exception
|
||||
failing_client = FailingBulkClient(async_client)
|
||||
await helpers.async_bulk(
|
||||
failing_client, [{"a": 42}], index="i", ignore_status=(599,)
|
||||
)
|
||||
|
||||
async def test_errors_are_collected_properly(self, async_client):
|
||||
await async_client.indices.create(
|
||||
"i",
|
||||
{
|
||||
"mappings": {"properties": {"a": {"type": "integer"}}},
|
||||
"settings": {"number_of_shards": 1, "number_of_replicas": 0},
|
||||
},
|
||||
)
|
||||
await async_client.cluster.health(wait_for_status="yellow")
|
||||
|
||||
success, failed = await helpers.async_bulk(
|
||||
async_client,
|
||||
[{"a": 42}, {"a": "c"}],
|
||||
index="i",
|
||||
stats_only=True,
|
||||
raise_on_error=False,
|
||||
)
|
||||
assert 1 == success
|
||||
assert 1 == failed
|
||||
|
||||
|
||||
class MockScroll:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
async def __call__(self, *args, **kwargs):
|
||||
self.calls.append((args, kwargs))
|
||||
if len(self.calls) == 1:
|
||||
return {
|
||||
"_scroll_id": "dummy_id",
|
||||
"_shards": {"successful": 4, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": [{"scroll_data": 42}]},
|
||||
}
|
||||
elif len(self.calls) == 2:
|
||||
return {
|
||||
"_scroll_id": "dummy_id",
|
||||
"_shards": {"successful": 4, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": []},
|
||||
}
|
||||
else:
|
||||
raise Exception("no more responses")
|
||||
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self, resp):
|
||||
self.resp = resp
|
||||
|
||||
async def __call__(self, *args, **kwargs):
|
||||
return self.resp
|
||||
|
||||
def __await__(self):
|
||||
return self().__await__()
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
async def scan_teardown(async_client):
|
||||
yield
|
||||
await async_client.clear_scroll(scroll_id="_all")
|
||||
|
||||
|
||||
class TestScan(object):
|
||||
async def test_order_can_be_preserved(self, async_client, scan_teardown):
|
||||
bulk = []
|
||||
for x in range(100):
|
||||
bulk.append({"index": {"_index": "test_index", "_id": x}})
|
||||
bulk.append({"answer": x, "correct": x == 42})
|
||||
await async_client.bulk(bulk, refresh=True)
|
||||
|
||||
docs = [
|
||||
doc
|
||||
async for doc in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
query={"sort": "answer"},
|
||||
preserve_order=True,
|
||||
)
|
||||
]
|
||||
|
||||
assert 100 == len(docs)
|
||||
assert list(map(str, range(100))) == list(d["_id"] for d in docs)
|
||||
assert list(range(100)) == list(d["_source"]["answer"] for d in docs)
|
||||
|
||||
async def test_all_documents_are_read(self, async_client, scan_teardown):
|
||||
bulk = []
|
||||
for x in range(100):
|
||||
bulk.append({"index": {"_index": "test_index", "_id": x}})
|
||||
bulk.append({"answer": x, "correct": x == 42})
|
||||
await async_client.bulk(bulk, refresh=True)
|
||||
|
||||
docs = [
|
||||
x
|
||||
async for x in helpers.async_scan(async_client, index="test_index", size=2)
|
||||
]
|
||||
|
||||
assert 100 == len(docs)
|
||||
assert set(map(str, range(100))) == set(d["_id"] for d in docs)
|
||||
assert set(range(100)) == set(d["_source"]["answer"] for d in docs)
|
||||
|
||||
async def test_scroll_error(self, async_client, scan_teardown):
|
||||
bulk = []
|
||||
for x in range(4):
|
||||
bulk.append({"index": {"_index": "test_index"}})
|
||||
bulk.append({"value": x})
|
||||
await async_client.bulk(bulk, refresh=True)
|
||||
|
||||
with patch.object(async_client, "scroll", MockScroll()):
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
size=2,
|
||||
raise_on_error=False,
|
||||
clear_scroll=False,
|
||||
)
|
||||
]
|
||||
assert len(data) == 3
|
||||
assert data[-1] == {"scroll_data": 42}
|
||||
|
||||
with patch.object(async_client, "scroll", MockScroll()):
|
||||
with pytest.raises(ScanError):
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
size=2,
|
||||
raise_on_error=True,
|
||||
clear_scroll=False,
|
||||
)
|
||||
]
|
||||
assert len(data) == 3
|
||||
assert data[-1] == {"scroll_data": 42}
|
||||
|
||||
async def test_initial_search_error(self, async_client, scan_teardown):
|
||||
with patch.object(async_client, "clear_scroll", new_callable=AsyncMock):
|
||||
with patch.object(
|
||||
async_client,
|
||||
"search",
|
||||
MockResponse(
|
||||
{
|
||||
"_scroll_id": "dummy_id",
|
||||
"_shards": {"successful": 4, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": [{"search_data": 1}]},
|
||||
}
|
||||
),
|
||||
):
|
||||
with patch.object(async_client, "scroll", MockScroll()):
|
||||
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
size=2,
|
||||
raise_on_error=False,
|
||||
)
|
||||
]
|
||||
assert data == [{"search_data": 1}, {"scroll_data": 42}]
|
||||
|
||||
with patch.object(
|
||||
async_client,
|
||||
"search",
|
||||
MockResponse(
|
||||
{
|
||||
"_scroll_id": "dummy_id",
|
||||
"_shards": {"successful": 4, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": [{"search_data": 1}]},
|
||||
}
|
||||
),
|
||||
):
|
||||
with patch.object(async_client, "scroll", MockScroll()) as mock_scroll:
|
||||
|
||||
with pytest.raises(ScanError):
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
size=2,
|
||||
raise_on_error=True,
|
||||
)
|
||||
]
|
||||
assert data == [{"search_data": 1}]
|
||||
assert mock_scroll.calls == []
|
||||
|
||||
async def test_no_scroll_id_fast_route(self, async_client, scan_teardown):
|
||||
with patch.object(async_client, "search", MockResponse({"no": "_scroll_id"})):
|
||||
with patch.object(async_client, "scroll") as scroll_mock:
|
||||
with patch.object(async_client, "clear_scroll") as clear_mock:
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client, index="test_index"
|
||||
)
|
||||
]
|
||||
|
||||
assert data == []
|
||||
scroll_mock.assert_not_called()
|
||||
clear_mock.assert_not_called()
|
||||
|
||||
@patch("opensearchpy._async.helpers.logger")
|
||||
async def test_logger(self, logger_mock, async_client, scan_teardown):
|
||||
bulk = []
|
||||
for x in range(4):
|
||||
bulk.append({"index": {"_index": "test_index"}})
|
||||
bulk.append({"value": x})
|
||||
await async_client.bulk(bulk, refresh=True)
|
||||
|
||||
with patch.object(async_client, "scroll", MockScroll()):
|
||||
_ = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
size=2,
|
||||
raise_on_error=False,
|
||||
clear_scroll=False,
|
||||
)
|
||||
]
|
||||
logger_mock.warning.assert_called()
|
||||
|
||||
with patch.object(async_client, "scroll", MockScroll()):
|
||||
try:
|
||||
_ = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
size=2,
|
||||
raise_on_error=True,
|
||||
clear_scroll=False,
|
||||
)
|
||||
]
|
||||
except ScanError:
|
||||
pass
|
||||
logger_mock.warning.assert_called_with(
|
||||
"Scroll request has only succeeded on %d (+%d skipped) shards out of %d.",
|
||||
4,
|
||||
0,
|
||||
5,
|
||||
)
|
||||
|
||||
async def test_clear_scroll(self, async_client, scan_teardown):
|
||||
bulk = []
|
||||
for x in range(4):
|
||||
bulk.append({"index": {"_index": "test_index"}})
|
||||
bulk.append({"value": x})
|
||||
await async_client.bulk(bulk, refresh=True)
|
||||
|
||||
with patch.object(
|
||||
async_client, "clear_scroll", wraps=async_client.clear_scroll
|
||||
) as spy:
|
||||
_ = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client, index="test_index", size=2
|
||||
)
|
||||
]
|
||||
spy.assert_called_once()
|
||||
|
||||
spy.reset_mock()
|
||||
_ = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client, index="test_index", size=2, clear_scroll=True
|
||||
)
|
||||
]
|
||||
spy.assert_called_once()
|
||||
|
||||
spy.reset_mock()
|
||||
_ = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client, index="test_index", size=2, clear_scroll=False
|
||||
)
|
||||
]
|
||||
spy.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs",
|
||||
[
|
||||
{"api_key": ("name", "value")},
|
||||
{"http_auth": ("username", "password")},
|
||||
{"headers": {"custom", "header"}},
|
||||
],
|
||||
)
|
||||
async def test_scan_auth_kwargs_forwarded(
|
||||
self, async_client, scan_teardown, kwargs
|
||||
):
|
||||
((key, val),) = kwargs.items()
|
||||
|
||||
with patch.object(
|
||||
async_client,
|
||||
"search",
|
||||
return_value=MockResponse(
|
||||
{
|
||||
"_scroll_id": "scroll_id",
|
||||
"_shards": {"successful": 5, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": [{"search_data": 1}]},
|
||||
}
|
||||
),
|
||||
) as search_mock:
|
||||
with patch.object(
|
||||
async_client,
|
||||
"scroll",
|
||||
return_value=MockResponse(
|
||||
{
|
||||
"_scroll_id": "scroll_id",
|
||||
"_shards": {"successful": 5, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": []},
|
||||
}
|
||||
),
|
||||
) as scroll_mock:
|
||||
with patch.object(
|
||||
async_client, "clear_scroll", return_value=MockResponse({})
|
||||
) as clear_mock:
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client, index="test_index", **kwargs
|
||||
)
|
||||
]
|
||||
|
||||
assert data == [{"search_data": 1}]
|
||||
|
||||
for api_mock in (search_mock, scroll_mock, clear_mock):
|
||||
assert api_mock.call_args[1][key] == val
|
||||
|
||||
async def test_scan_auth_kwargs_favor_scroll_kwargs_option(
|
||||
self, async_client, scan_teardown
|
||||
):
|
||||
with patch.object(
|
||||
async_client,
|
||||
"search",
|
||||
return_value=MockResponse(
|
||||
{
|
||||
"_scroll_id": "scroll_id",
|
||||
"_shards": {"successful": 5, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": [{"search_data": 1}]},
|
||||
}
|
||||
),
|
||||
):
|
||||
with patch.object(
|
||||
async_client,
|
||||
"scroll",
|
||||
return_value=MockResponse(
|
||||
{
|
||||
"_scroll_id": "scroll_id",
|
||||
"_shards": {"successful": 5, "total": 5, "skipped": 0},
|
||||
"hits": {"hits": []},
|
||||
}
|
||||
),
|
||||
):
|
||||
with patch.object(
|
||||
async_client, "clear_scroll", return_value=MockResponse({})
|
||||
):
|
||||
data = [
|
||||
x
|
||||
async for x in helpers.async_scan(
|
||||
async_client,
|
||||
index="test_index",
|
||||
headers={"not scroll": "kwargs"},
|
||||
scroll_kwargs={
|
||||
"headers": {"scroll": "kwargs"},
|
||||
"sort": "asc",
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
assert data == [{"search_data": 1}]
|
||||
|
||||
# Assert that we see 'scroll_kwargs' options used instead of 'kwargs'
|
||||
assert async_client.scroll.call_args[1]["headers"] == {
|
||||
"scroll": "kwargs"
|
||||
}
|
||||
assert async_client.scroll.call_args[1]["sort"] == "asc"
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
async def reindex_setup(async_client):
|
||||
bulk = []
|
||||
for x in range(100):
|
||||
bulk.append({"index": {"_index": "test_index", "_id": x}})
|
||||
bulk.append(
|
||||
{
|
||||
"answer": x,
|
||||
"correct": x == 42,
|
||||
"type": "answers" if x % 2 == 0 else "questions",
|
||||
}
|
||||
)
|
||||
await async_client.bulk(bulk, refresh=True)
|
||||
yield
|
||||
|
||||
|
||||
class TestReindex(object):
|
||||
async def test_reindex_passes_kwargs_to_scan_and_bulk(
|
||||
self, async_client, reindex_setup
|
||||
):
|
||||
await helpers.async_reindex(
|
||||
async_client,
|
||||
"test_index",
|
||||
"prod_index",
|
||||
scan_kwargs={"q": "type:answers"},
|
||||
bulk_kwargs={"refresh": True},
|
||||
)
|
||||
|
||||
assert await async_client.indices.exists("prod_index")
|
||||
assert (
|
||||
50
|
||||
== (await async_client.count(index="prod_index", q="type:answers"))["count"]
|
||||
)
|
||||
|
||||
assert {"answer": 42, "correct": True, "type": "answers"} == (
|
||||
await async_client.get(index="prod_index", id=42)
|
||||
)["_source"]
|
||||
|
||||
async def test_reindex_accepts_a_query(self, async_client, reindex_setup):
|
||||
await helpers.async_reindex(
|
||||
async_client,
|
||||
"test_index",
|
||||
"prod_index",
|
||||
query={"query": {"bool": {"filter": {"term": {"type": "answers"}}}}},
|
||||
)
|
||||
await async_client.indices.refresh()
|
||||
|
||||
assert await async_client.indices.exists("prod_index")
|
||||
assert (
|
||||
50
|
||||
== (await async_client.count(index="prod_index", q="type:answers"))["count"]
|
||||
)
|
||||
|
||||
assert {"answer": 42, "correct": True, "type": "answers"} == (
|
||||
await async_client.get(index="prod_index", id=42)
|
||||
)["_source"]
|
||||
|
||||
async def test_all_documents_get_moved(self, async_client, reindex_setup):
|
||||
await helpers.async_reindex(async_client, "test_index", "prod_index")
|
||||
await async_client.indices.refresh()
|
||||
|
||||
assert await async_client.indices.exists("prod_index")
|
||||
assert (
|
||||
50
|
||||
== (await async_client.count(index="prod_index", q="type:questions"))[
|
||||
"count"
|
||||
]
|
||||
)
|
||||
assert (
|
||||
50
|
||||
== (await async_client.count(index="prod_index", q="type:answers"))["count"]
|
||||
)
|
||||
|
||||
assert {"answer": 42, "correct": True, "type": "answers"} == (
|
||||
await async_client.get(index="prod_index", id=42)
|
||||
)["_source"]
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
async def parent_reindex_setup(async_client):
|
||||
body = {
|
||||
"settings": {"number_of_shards": 1, "number_of_replicas": 0},
|
||||
"mappings": {
|
||||
"properties": {
|
||||
"question_answer": {
|
||||
"type": "join",
|
||||
"relations": {"question": "answer"},
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
await async_client.indices.create(index="test-index", body=body)
|
||||
await async_client.indices.create(index="real-index", body=body)
|
||||
|
||||
await async_client.index(
|
||||
index="test-index", id=42, body={"question_answer": "question"}
|
||||
)
|
||||
await async_client.index(
|
||||
index="test-index",
|
||||
id=47,
|
||||
routing=42,
|
||||
body={"some": "data", "question_answer": {"name": "answer", "parent": 42}},
|
||||
)
|
||||
await async_client.indices.refresh(index="test-index")
|
||||
|
||||
|
||||
class TestParentChildReindex:
|
||||
async def test_children_are_reindexed_correctly(
|
||||
self, async_client, parent_reindex_setup
|
||||
):
|
||||
await helpers.async_reindex(async_client, "test-index", "real-index")
|
||||
|
||||
q = await async_client.get(index="real-index", id=42)
|
||||
assert {
|
||||
"_id": "42",
|
||||
"_index": "real-index",
|
||||
"_primary_term": 1,
|
||||
"_seq_no": 0,
|
||||
"_source": {"question_answer": "question"},
|
||||
"_type": "_doc",
|
||||
"_version": 1,
|
||||
"found": True,
|
||||
} == q
|
||||
|
||||
q = await async_client.get(index="test-index", id=47, routing=42)
|
||||
assert {
|
||||
"_routing": "42",
|
||||
"_id": "47",
|
||||
"_index": "test-index",
|
||||
"_primary_term": 1,
|
||||
"_seq_no": 1,
|
||||
"_source": {
|
||||
"some": "data",
|
||||
"question_answer": {"name": "answer", "parent": 42},
|
||||
},
|
||||
"_type": "_doc",
|
||||
"_version": 1,
|
||||
"found": True,
|
||||
} == q
|
||||
@@ -0,0 +1,233 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
|
||||
"""
|
||||
Dynamically generated set of TestCases based on set of yaml files decribing
|
||||
some integration tests. These files are shared among all official OpenSearch
|
||||
clients.
|
||||
"""
|
||||
import inspect
|
||||
import warnings
|
||||
|
||||
import pytest
|
||||
|
||||
from opensearchpy import OpenSearchWarning
|
||||
from opensearchpy.helpers.test import _get_version
|
||||
|
||||
from ...test_server.test_rest_api_spec import (
|
||||
IMPLEMENTED_FEATURES,
|
||||
PARAMS_RENAMES,
|
||||
RUN_ASYNC_REST_API_TESTS,
|
||||
YAML_TEST_SPECS,
|
||||
YamlRunner,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
OPENSEARCH_VERSION = None
|
||||
|
||||
|
||||
async def await_if_coro(x):
|
||||
if inspect.iscoroutine(x):
|
||||
return await x
|
||||
return x
|
||||
|
||||
|
||||
class AsyncYamlRunner(YamlRunner):
|
||||
async def setup(self):
|
||||
# Pull skips from individual tests to not do unnecessary setup.
|
||||
skip_code = []
|
||||
for action in self._run_code:
|
||||
assert len(action) == 1
|
||||
action_type, _ = list(action.items())[0]
|
||||
if action_type == "skip":
|
||||
skip_code.append(action)
|
||||
else:
|
||||
break
|
||||
|
||||
if self._setup_code or skip_code:
|
||||
self.section("setup")
|
||||
if skip_code:
|
||||
await self.run_code(skip_code)
|
||||
if self._setup_code:
|
||||
await self.run_code(self._setup_code)
|
||||
|
||||
async def teardown(self):
|
||||
if self._teardown_code:
|
||||
self.section("teardown")
|
||||
await self.run_code(self._teardown_code)
|
||||
|
||||
async def opensearch_version(self):
|
||||
global OPENSEARCH_VERSION
|
||||
if OPENSEARCH_VERSION is None:
|
||||
version_string = (await self.client.info())["version"]["number"]
|
||||
if "." not in version_string:
|
||||
return ()
|
||||
version = version_string.strip().split(".")
|
||||
OPENSEARCH_VERSION = tuple(int(v) if v.isdigit() else 999 for v in version)
|
||||
return OPENSEARCH_VERSION
|
||||
|
||||
def section(self, name):
|
||||
print(("=" * 10) + " " + name + " " + ("=" * 10))
|
||||
|
||||
async def run(self):
|
||||
try:
|
||||
await self.setup()
|
||||
self.section("test")
|
||||
await self.run_code(self._run_code)
|
||||
finally:
|
||||
try:
|
||||
await self.teardown()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
async def run_code(self, test):
|
||||
"""Execute an instruction based on it's type."""
|
||||
for action in test:
|
||||
assert len(action) == 1
|
||||
action_type, action = list(action.items())[0]
|
||||
print(action_type, action)
|
||||
|
||||
if hasattr(self, "run_" + action_type):
|
||||
await await_if_coro(getattr(self, "run_" + action_type)(action))
|
||||
else:
|
||||
raise RuntimeError("Invalid action type %r" % (action_type,))
|
||||
|
||||
async def run_do(self, action):
|
||||
api = self.client
|
||||
headers = action.pop("headers", None)
|
||||
catch = action.pop("catch", None)
|
||||
warn = action.pop("warnings", ())
|
||||
allowed_warnings = action.pop("allowed_warnings", ())
|
||||
assert len(action) == 1
|
||||
|
||||
# Remove the x_pack_rest_user authentication
|
||||
# if it's given via headers. We're already authenticated
|
||||
# via the 'elastic' user.
|
||||
if (
|
||||
headers
|
||||
and headers.get("Authorization", None)
|
||||
== "Basic eF9wYWNrX3Jlc3RfdXNlcjp4LXBhY2stdGVzdC1wYXNzd29yZA=="
|
||||
):
|
||||
headers.pop("Authorization")
|
||||
|
||||
method, args = list(action.items())[0]
|
||||
args["headers"] = headers
|
||||
|
||||
# locate api endpoint
|
||||
for m in method.split("."):
|
||||
assert hasattr(api, m)
|
||||
api = getattr(api, m)
|
||||
|
||||
# some parameters had to be renamed to not clash with python builtins,
|
||||
# compensate
|
||||
for k in PARAMS_RENAMES:
|
||||
if k in args:
|
||||
args[PARAMS_RENAMES[k]] = args.pop(k)
|
||||
|
||||
# resolve vars
|
||||
for k in args:
|
||||
args[k] = self._resolve(args[k])
|
||||
|
||||
warnings.simplefilter("always", category=OpenSearchWarning)
|
||||
with warnings.catch_warnings(record=True) as caught_warnings:
|
||||
try:
|
||||
self.last_response = await api(**args)
|
||||
except Exception as e:
|
||||
if not catch:
|
||||
raise
|
||||
self.run_catch(catch, e)
|
||||
else:
|
||||
if catch:
|
||||
raise AssertionError(
|
||||
"Failed to catch %r in %r." % (catch, self.last_response)
|
||||
)
|
||||
|
||||
# Filter out warnings raised by other components.
|
||||
caught_warnings = [
|
||||
str(w.message)
|
||||
for w in caught_warnings
|
||||
if w.category == OpenSearchWarning
|
||||
and str(w.message) not in allowed_warnings
|
||||
]
|
||||
|
||||
# This warning can show up in many places but isn't accounted for
|
||||
# in tests, so we remove it to make sure things pass.
|
||||
include_type_name_warning = (
|
||||
"[types removal] Using include_type_name in create index requests is deprecated. "
|
||||
"The parameter will be removed in the next major version."
|
||||
)
|
||||
if (
|
||||
include_type_name_warning in caught_warnings
|
||||
and include_type_name_warning not in warn
|
||||
):
|
||||
caught_warnings.remove(include_type_name_warning)
|
||||
|
||||
# Sorting removes the issue with order raised. We only care about
|
||||
# if all warnings are raised in the single API call.
|
||||
if warn and sorted(warn) != sorted(caught_warnings):
|
||||
raise AssertionError(
|
||||
"Expected warnings not equal to actual warnings: expected=%r actual=%r"
|
||||
% (warn, caught_warnings)
|
||||
)
|
||||
|
||||
async def run_skip(self, skip):
|
||||
if "features" in skip:
|
||||
features = skip["features"]
|
||||
if not isinstance(features, (tuple, list)):
|
||||
features = [features]
|
||||
for feature in features:
|
||||
if feature in IMPLEMENTED_FEATURES:
|
||||
continue
|
||||
pytest.skip("feature '%s' is not supported" % feature)
|
||||
|
||||
if "version" in skip:
|
||||
version, reason = skip["version"], skip["reason"]
|
||||
if version == "all":
|
||||
pytest.skip(reason)
|
||||
min_version, max_version = version.split("-")
|
||||
min_version = _get_version(min_version) or (0,)
|
||||
max_version = _get_version(max_version) or (999,)
|
||||
if min_version <= (await self.opensearch_version()) <= max_version:
|
||||
pytest.skip(reason)
|
||||
|
||||
async def _feature_enabled(self, name):
|
||||
return False
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def async_runner(async_client):
|
||||
return AsyncYamlRunner(async_client)
|
||||
|
||||
|
||||
if RUN_ASYNC_REST_API_TESTS:
|
||||
|
||||
@pytest.mark.parametrize("test_spec", YAML_TEST_SPECS)
|
||||
async def test_rest_api_spec(test_spec, async_runner):
|
||||
if test_spec.get("skip", False):
|
||||
pytest.skip("Manually skipped in 'SKIP_TESTS'")
|
||||
async_runner.use_spec(test_spec)
|
||||
await async_runner.run()
|
||||
@@ -0,0 +1,549 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
#
|
||||
# The OpenSearch Contributors require contributions made to
|
||||
# this file be licensed under the Apache-2.0 license or a
|
||||
# compatible open source license.
|
||||
#
|
||||
# Modifications Copyright OpenSearch Contributors. See
|
||||
# GitHub history for details.
|
||||
#
|
||||
# Licensed to Elasticsearch B.V. under one or more contributor
|
||||
# license agreements. See the NOTICE file distributed with
|
||||
# this work for additional information regarding copyright
|
||||
# ownership. Elasticsearch B.V. licenses this file to you under
|
||||
# the Apache License, Version 2.0 (the "License"); you may
|
||||
# not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing,
|
||||
# software distributed under the License is distributed on an
|
||||
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, either express or implied. See the License for the
|
||||
# specific language governing permissions and limitations
|
||||
# under the License.
|
||||
|
||||
from __future__ import unicode_literals
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from mock import patch
|
||||
|
||||
from opensearchpy import AsyncTransport
|
||||
from opensearchpy.connection import Connection
|
||||
from opensearchpy.connection_pool import DummyConnectionPool
|
||||
from opensearchpy.exceptions import ConnectionError, TransportError
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
class DummyConnection(Connection):
|
||||
def __init__(self, **kwargs):
|
||||
self.exception = kwargs.pop("exception", None)
|
||||
self.status, self.data = kwargs.pop("status", 200), kwargs.pop("data", "{}")
|
||||
self.headers = kwargs.pop("headers", {})
|
||||
self.delay = kwargs.pop("delay", 0)
|
||||
self.calls = []
|
||||
self.closed = False
|
||||
super(DummyConnection, self).__init__(**kwargs)
|
||||
|
||||
async def perform_request(self, *args, **kwargs):
|
||||
if self.closed:
|
||||
raise RuntimeError("This connection is closed")
|
||||
if self.delay:
|
||||
await asyncio.sleep(self.delay)
|
||||
self.calls.append((args, kwargs))
|
||||
if self.exception:
|
||||
raise self.exception
|
||||
return self.status, self.headers, self.data
|
||||
|
||||
async def close(self):
|
||||
if self.closed:
|
||||
raise RuntimeError("This connection is already closed")
|
||||
self.closed = True
|
||||
|
||||
|
||||
CLUSTER_NODES = """{
|
||||
"_nodes" : {
|
||||
"total" : 1,
|
||||
"successful" : 1,
|
||||
"failed" : 0
|
||||
},
|
||||
"cluster_name" : "opensearch",
|
||||
"nodes" : {
|
||||
"SRZpKFZdQguhhvifmN6UVA" : {
|
||||
"name" : "SRZpKFZ",
|
||||
"transport_address" : "127.0.0.1:9300",
|
||||
"host" : "127.0.0.1",
|
||||
"ip" : "127.0.0.1",
|
||||
"version" : "5.0.0",
|
||||
"build_hash" : "253032b",
|
||||
"roles" : [ "master", "data", "ingest" ],
|
||||
"http" : {
|
||||
"bound_address" : [ "[fe80::1]:9200", "[::1]:9200", "127.0.0.1:9200" ],
|
||||
"publish_address" : "1.1.1.1:123",
|
||||
"max_content_length_in_bytes" : 104857600
|
||||
}
|
||||
}
|
||||
}
|
||||
}"""
|
||||
|
||||
CLUSTER_NODES_7x_PUBLISH_HOST = """{
|
||||
"_nodes" : {
|
||||
"total" : 1,
|
||||
"successful" : 1,
|
||||
"failed" : 0
|
||||
},
|
||||
"cluster_name" : "opensearch",
|
||||
"nodes" : {
|
||||
"SRZpKFZdQguhhvifmN6UVA" : {
|
||||
"name" : "SRZpKFZ",
|
||||
"transport_address" : "127.0.0.1:9300",
|
||||
"host" : "127.0.0.1",
|
||||
"ip" : "127.0.0.1",
|
||||
"version" : "5.0.0",
|
||||
"build_hash" : "253032b",
|
||||
"roles" : [ "master", "data", "ingest" ],
|
||||
"http" : {
|
||||
"bound_address" : [ "[fe80::1]:9200", "[::1]:9200", "127.0.0.1:9200" ],
|
||||
"publish_address" : "somehost.tld/1.1.1.1:123",
|
||||
"max_content_length_in_bytes" : 104857600
|
||||
}
|
||||
}
|
||||
}
|
||||
}"""
|
||||
|
||||
|
||||
class TestTransport:
|
||||
async def test_single_connection_uses_dummy_connection_pool(self):
|
||||
t = AsyncTransport([{}])
|
||||
await t._async_call()
|
||||
assert isinstance(t.connection_pool, DummyConnectionPool)
|
||||
t = AsyncTransport([{"host": "localhost"}])
|
||||
await t._async_call()
|
||||
assert isinstance(t.connection_pool, DummyConnectionPool)
|
||||
|
||||
async def test_request_timeout_extracted_from_params_and_passed(self):
|
||||
t = AsyncTransport([{}], connection_class=DummyConnection)
|
||||
|
||||
await t.perform_request("GET", "/", params={"request_timeout": 42})
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert ("GET", "/", {}, None) == t.get_connection().calls[0][0]
|
||||
assert {
|
||||
"timeout": 42,
|
||||
"ignore": (),
|
||||
"headers": None,
|
||||
} == t.get_connection().calls[0][1]
|
||||
|
||||
async def test_opaque_id(self):
|
||||
t = AsyncTransport([{}], opaque_id="app-1", connection_class=DummyConnection)
|
||||
|
||||
await t.perform_request("GET", "/")
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert ("GET", "/", None, None) == t.get_connection().calls[0][0]
|
||||
assert {
|
||||
"timeout": None,
|
||||
"ignore": (),
|
||||
"headers": None,
|
||||
} == t.get_connection().calls[0][1]
|
||||
|
||||
# Now try with an 'x-opaque-id' set on perform_request().
|
||||
await t.perform_request("GET", "/", headers={"x-opaque-id": "request-1"})
|
||||
assert 2 == len(t.get_connection().calls)
|
||||
assert ("GET", "/", None, None) == t.get_connection().calls[1][0]
|
||||
assert {
|
||||
"timeout": None,
|
||||
"ignore": (),
|
||||
"headers": {"x-opaque-id": "request-1"},
|
||||
} == t.get_connection().calls[1][1]
|
||||
|
||||
async def test_request_with_custom_user_agent_header(self):
|
||||
t = AsyncTransport([{}], connection_class=DummyConnection)
|
||||
|
||||
await t.perform_request(
|
||||
"GET", "/", headers={"user-agent": "my-custom-value/1.2.3"}
|
||||
)
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert {
|
||||
"timeout": None,
|
||||
"ignore": (),
|
||||
"headers": {"user-agent": "my-custom-value/1.2.3"},
|
||||
} == t.get_connection().calls[0][1]
|
||||
|
||||
async def test_send_get_body_as_source(self):
|
||||
t = AsyncTransport(
|
||||
[{}], send_get_body_as="source", connection_class=DummyConnection
|
||||
)
|
||||
|
||||
await t.perform_request("GET", "/", body={})
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert ("GET", "/", {"source": "{}"}, None) == t.get_connection().calls[0][0]
|
||||
|
||||
async def test_send_get_body_as_post(self):
|
||||
t = AsyncTransport(
|
||||
[{}], send_get_body_as="POST", connection_class=DummyConnection
|
||||
)
|
||||
|
||||
await t.perform_request("GET", "/", body={})
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert ("POST", "/", None, b"{}") == t.get_connection().calls[0][0]
|
||||
|
||||
async def test_body_gets_encoded_into_bytes(self):
|
||||
t = AsyncTransport([{}], connection_class=DummyConnection)
|
||||
|
||||
await t.perform_request("GET", "/", body="你好")
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert (
|
||||
"GET",
|
||||
"/",
|
||||
None,
|
||||
b"\xe4\xbd\xa0\xe5\xa5\xbd",
|
||||
) == t.get_connection().calls[0][0]
|
||||
|
||||
async def test_body_bytes_get_passed_untouched(self):
|
||||
t = AsyncTransport([{}], connection_class=DummyConnection)
|
||||
|
||||
body = b"\xe4\xbd\xa0\xe5\xa5\xbd"
|
||||
await t.perform_request("GET", "/", body=body)
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert ("GET", "/", None, body) == t.get_connection().calls[0][0]
|
||||
|
||||
async def test_body_surrogates_replaced_encoded_into_bytes(self):
|
||||
t = AsyncTransport([{}], connection_class=DummyConnection)
|
||||
|
||||
await t.perform_request("GET", "/", body="你好\uda6a")
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert (
|
||||
"GET",
|
||||
"/",
|
||||
None,
|
||||
b"\xe4\xbd\xa0\xe5\xa5\xbd\xed\xa9\xaa",
|
||||
) == t.get_connection().calls[0][0]
|
||||
|
||||
async def test_kwargs_passed_on_to_connections(self):
|
||||
t = AsyncTransport([{"host": "google.com"}], port=123)
|
||||
await t._async_call()
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert "http://google.com:123" == t.connection_pool.connections[0].host
|
||||
|
||||
async def test_kwargs_passed_on_to_connection_pool(self):
|
||||
dt = object()
|
||||
t = AsyncTransport([{}, {}], dead_timeout=dt)
|
||||
await t._async_call()
|
||||
assert dt is t.connection_pool.dead_timeout
|
||||
|
||||
async def test_custom_connection_class(self):
|
||||
class MyConnection(object):
|
||||
def __init__(self, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
|
||||
t = AsyncTransport([{}], connection_class=MyConnection)
|
||||
await t._async_call()
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert isinstance(t.connection_pool.connections[0], MyConnection)
|
||||
|
||||
def test_add_connection(self):
|
||||
t = AsyncTransport([{}], randomize_hosts=False)
|
||||
t.add_connection({"host": "google.com", "port": 1234})
|
||||
|
||||
assert 2 == len(t.connection_pool.connections)
|
||||
assert "http://google.com:1234" == t.connection_pool.connections[1].host
|
||||
|
||||
async def test_request_will_fail_after_X_retries(self):
|
||||
t = AsyncTransport(
|
||||
[{"exception": ConnectionError("abandon ship")}],
|
||||
connection_class=DummyConnection,
|
||||
)
|
||||
|
||||
connection_error = False
|
||||
try:
|
||||
await t.perform_request("GET", "/")
|
||||
except ConnectionError:
|
||||
connection_error = True
|
||||
|
||||
assert connection_error
|
||||
assert 4 == len(t.get_connection().calls)
|
||||
|
||||
async def test_failed_connection_will_be_marked_as_dead(self):
|
||||
t = AsyncTransport(
|
||||
[{"exception": ConnectionError("abandon ship")}] * 2,
|
||||
connection_class=DummyConnection,
|
||||
)
|
||||
|
||||
connection_error = False
|
||||
try:
|
||||
await t.perform_request("GET", "/")
|
||||
except ConnectionError:
|
||||
connection_error = True
|
||||
|
||||
assert connection_error
|
||||
assert 0 == len(t.connection_pool.connections)
|
||||
|
||||
async def test_resurrected_connection_will_be_marked_as_live_on_success(self):
|
||||
for method in ("GET", "HEAD"):
|
||||
t = AsyncTransport([{}, {}], connection_class=DummyConnection)
|
||||
await t._async_call()
|
||||
con1 = t.connection_pool.get_connection()
|
||||
con2 = t.connection_pool.get_connection()
|
||||
t.connection_pool.mark_dead(con1)
|
||||
t.connection_pool.mark_dead(con2)
|
||||
|
||||
await t.perform_request(method, "/")
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert 1 == len(t.connection_pool.dead_count)
|
||||
|
||||
async def test_sniff_will_use_seed_connections(self):
|
||||
t = AsyncTransport([{"data": CLUSTER_NODES}], connection_class=DummyConnection)
|
||||
await t._async_call()
|
||||
t.set_connections([{"data": "invalid"}])
|
||||
|
||||
await t.sniff_hosts()
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert "http://1.1.1.1:123" == t.get_connection().host
|
||||
|
||||
async def test_sniff_on_start_fetches_and_uses_nodes_list(self):
|
||||
t = AsyncTransport(
|
||||
[{"data": CLUSTER_NODES}],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_start=True,
|
||||
)
|
||||
await t._async_call()
|
||||
await t.sniffing_task # Need to wait for the sniffing task to complete
|
||||
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert "http://1.1.1.1:123" == t.get_connection().host
|
||||
|
||||
async def test_sniff_on_start_ignores_sniff_timeout(self):
|
||||
t = AsyncTransport(
|
||||
[{"data": CLUSTER_NODES}],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_start=True,
|
||||
sniff_timeout=12,
|
||||
)
|
||||
await t._async_call()
|
||||
await t.sniffing_task # Need to wait for the sniffing task to complete
|
||||
|
||||
assert (("GET", "/_nodes/_all/http"), {"timeout": None}) == t.seed_connections[
|
||||
0
|
||||
].calls[0]
|
||||
|
||||
async def test_sniff_uses_sniff_timeout(self):
|
||||
t = AsyncTransport(
|
||||
[{"data": CLUSTER_NODES}],
|
||||
connection_class=DummyConnection,
|
||||
sniff_timeout=42,
|
||||
)
|
||||
await t._async_call()
|
||||
await t.sniff_hosts()
|
||||
|
||||
assert (("GET", "/_nodes/_all/http"), {"timeout": 42}) == t.seed_connections[
|
||||
0
|
||||
].calls[0]
|
||||
|
||||
async def test_sniff_reuses_connection_instances_if_possible(self):
|
||||
t = AsyncTransport(
|
||||
[{"data": CLUSTER_NODES}, {"host": "1.1.1.1", "port": 123}],
|
||||
connection_class=DummyConnection,
|
||||
randomize_hosts=False,
|
||||
)
|
||||
await t._async_call()
|
||||
connection = t.connection_pool.connections[1]
|
||||
connection.delay = 3.0 # Add this delay to make the sniffing deterministic.
|
||||
|
||||
await t.sniff_hosts()
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert connection is t.get_connection()
|
||||
|
||||
async def test_sniff_on_fail_triggers_sniffing_on_fail(self):
|
||||
t = AsyncTransport(
|
||||
[{"exception": ConnectionError("abandon ship")}, {"data": CLUSTER_NODES}],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_connection_fail=True,
|
||||
max_retries=0,
|
||||
randomize_hosts=False,
|
||||
)
|
||||
await t._async_call()
|
||||
|
||||
connection_error = False
|
||||
try:
|
||||
await t.perform_request("GET", "/")
|
||||
except ConnectionError:
|
||||
connection_error = True
|
||||
|
||||
await t.sniffing_task # Need to wait for the sniffing task to complete
|
||||
|
||||
assert connection_error
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert "http://1.1.1.1:123" == t.get_connection().host
|
||||
|
||||
@patch("opensearchpy._async.transport.AsyncTransport.sniff_hosts")
|
||||
async def test_sniff_on_fail_failing_does_not_prevent_retires(self, sniff_hosts):
|
||||
sniff_hosts.side_effect = [TransportError("sniff failed")]
|
||||
t = AsyncTransport(
|
||||
[{"exception": ConnectionError("abandon ship")}, {"data": CLUSTER_NODES}],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_connection_fail=True,
|
||||
max_retries=3,
|
||||
randomize_hosts=False,
|
||||
)
|
||||
await t._async_init()
|
||||
|
||||
conn_err, conn_data = t.connection_pool.connections
|
||||
response = await t.perform_request("GET", "/")
|
||||
assert json.loads(CLUSTER_NODES) == response
|
||||
assert 1 == sniff_hosts.call_count
|
||||
assert 1 == len(conn_err.calls)
|
||||
assert 1 == len(conn_data.calls)
|
||||
|
||||
async def test_sniff_after_n_seconds(self, event_loop):
|
||||
t = AsyncTransport(
|
||||
[{"data": CLUSTER_NODES}],
|
||||
connection_class=DummyConnection,
|
||||
sniffer_timeout=5,
|
||||
)
|
||||
await t._async_call()
|
||||
|
||||
for _ in range(4):
|
||||
await t.perform_request("GET", "/")
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert isinstance(t.get_connection(), DummyConnection)
|
||||
t.last_sniff = event_loop.time() - 5.1
|
||||
|
||||
await t.perform_request("GET", "/")
|
||||
await t.sniffing_task # Need to wait for the sniffing task to complete
|
||||
|
||||
assert 1 == len(t.connection_pool.connections)
|
||||
assert "http://1.1.1.1:123" == t.get_connection().host
|
||||
assert event_loop.time() - 1 < t.last_sniff < event_loop.time() + 0.01
|
||||
|
||||
async def test_sniff_7x_publish_host(self):
|
||||
# Test the response shaped when a 7.x node has publish_host set
|
||||
# and the returend data is shaped in the fqdn/ip:port format.
|
||||
t = AsyncTransport(
|
||||
[{"data": CLUSTER_NODES_7x_PUBLISH_HOST}],
|
||||
connection_class=DummyConnection,
|
||||
sniff_timeout=42,
|
||||
)
|
||||
await t._async_call()
|
||||
await t.sniff_hosts()
|
||||
# Ensure we parsed out the fqdn and port from the fqdn/ip:port string.
|
||||
assert t.connection_pool.connection_opts[0][1] == {
|
||||
"host": "somehost.tld",
|
||||
"port": 123,
|
||||
}
|
||||
|
||||
@patch("opensearchpy._async.transport.AsyncTransport.sniff_hosts")
|
||||
async def test_sniffing_disabled_on_cloud_instances(self, sniff_hosts):
|
||||
t = AsyncTransport(
|
||||
[{}],
|
||||
sniff_on_start=True,
|
||||
sniff_on_connection_fail=True,
|
||||
connection_class=DummyConnection,
|
||||
cloud_id="cluster:dXMtZWFzdC0xLmF3cy5mb3VuZC5pbyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5NyQ0ZmE4ODIxZTc1NjM0MDMyYmVkMWNmMjIxMTBlMmY5Ng==",
|
||||
)
|
||||
await t._async_call()
|
||||
|
||||
assert not t.sniff_on_connection_fail
|
||||
assert sniff_hosts.call_args is None # Assert not called.
|
||||
await t.perform_request("GET", "/", body={})
|
||||
assert 1 == len(t.get_connection().calls)
|
||||
assert ("GET", "/", None, b"{}") == t.get_connection().calls[0][0]
|
||||
|
||||
async def test_transport_close_closes_all_pool_connections(self):
|
||||
t = AsyncTransport([{}], connection_class=DummyConnection)
|
||||
await t._async_call()
|
||||
|
||||
assert not any([conn.closed for conn in t.connection_pool.connections])
|
||||
await t.close()
|
||||
assert all([conn.closed for conn in t.connection_pool.connections])
|
||||
|
||||
t = AsyncTransport([{}, {}], connection_class=DummyConnection)
|
||||
await t._async_call()
|
||||
|
||||
assert not any([conn.closed for conn in t.connection_pool.connections])
|
||||
await t.close()
|
||||
assert all([conn.closed for conn in t.connection_pool.connections])
|
||||
|
||||
async def test_sniff_on_start_error_if_no_sniffed_hosts(self, event_loop):
|
||||
t = AsyncTransport(
|
||||
[
|
||||
{"data": ""},
|
||||
{"data": ""},
|
||||
{"data": ""},
|
||||
],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_start=True,
|
||||
)
|
||||
|
||||
# If our initial sniffing attempt comes back
|
||||
# empty then we raise an error.
|
||||
with pytest.raises(TransportError) as e:
|
||||
await t._async_call()
|
||||
assert str(e.value) == "TransportError(N/A, 'Unable to sniff hosts.')"
|
||||
|
||||
async def test_sniff_on_start_waits_for_sniff_to_complete(self, event_loop):
|
||||
t = AsyncTransport(
|
||||
[
|
||||
{"delay": 1, "data": ""},
|
||||
{"delay": 1, "data": ""},
|
||||
{"delay": 1, "data": CLUSTER_NODES},
|
||||
],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_start=True,
|
||||
)
|
||||
|
||||
# Start the timer right before the first task
|
||||
# and have a bunch of tasks come in immediately.
|
||||
tasks = []
|
||||
start_time = event_loop.time()
|
||||
for _ in range(5):
|
||||
tasks.append(event_loop.create_task(t._async_call()))
|
||||
await asyncio.sleep(0) # Yield to the loop
|
||||
|
||||
assert t.sniffing_task is not None
|
||||
|
||||
# Tasks streaming in later.
|
||||
for _ in range(5):
|
||||
tasks.append(event_loop.create_task(t._async_call()))
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Now that all the API calls have come in we wait for
|
||||
# them all to resolve before
|
||||
await asyncio.gather(*tasks)
|
||||
end_time = event_loop.time()
|
||||
duration = end_time - start_time
|
||||
|
||||
# All the tasks blocked on the sniff of each node
|
||||
# and then resolved immediately after.
|
||||
assert 1 <= duration < 2
|
||||
|
||||
async def test_sniff_on_start_close_unlocks_async_calls(self, event_loop):
|
||||
t = AsyncTransport(
|
||||
[
|
||||
{"delay": 10, "data": CLUSTER_NODES},
|
||||
],
|
||||
connection_class=DummyConnection,
|
||||
sniff_on_start=True,
|
||||
)
|
||||
|
||||
# Start making _async_calls() before we cancel
|
||||
tasks = []
|
||||
start_time = event_loop.time()
|
||||
for _ in range(3):
|
||||
tasks.append(event_loop.create_task(t._async_call()))
|
||||
await asyncio.sleep(0)
|
||||
|
||||
# Close the transport while the sniffing task is active! :(
|
||||
await t.close()
|
||||
|
||||
# Now we start waiting on all those _async_calls()
|
||||
await asyncio.gather(*tasks)
|
||||
end_time = event_loop.time()
|
||||
duration = end_time - start_time
|
||||
|
||||
# A lot quicker than 10 seconds defined in 'delay'
|
||||
assert duration < 1
|
||||
Reference in New Issue
Block a user