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:
Rushi Agrawal
2021-09-16 21:23:38 +05:30
parent f4891be3c3
commit ef0c23c0e4
129 changed files with 237 additions and 237 deletions
+25
View File
@@ -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