diff --git a/dev_requirements.txt b/dev_requirements.txt index ad286d1b..84675024 100644 --- a/dev_requirements.txt +++ b/dev_requirements.txt @@ -4,3 +4,4 @@ nose coverage mock pyaml +pylibmc diff --git a/elasticsearch/__init__.py b/elasticsearch/__init__.py index b0494376..914f3889 100644 --- a/elasticsearch/__init__.py +++ b/elasticsearch/__init__.py @@ -9,6 +9,7 @@ from elasticsearch.transport import Transport from elasticsearch.connection_pool import ConnectionPool, ConnectionSelector, \ RoundRobinSelector from elasticsearch.serializer import JSONSerializer -from elasticsearch.connection import Connection, RequestsHttpConnection +from elasticsearch.connection import Connection, RequestsHttpConnection, \ + Urllib3HttpConnection, MemcachedConnection from elasticsearch.exceptions import * diff --git a/elasticsearch/connection.py b/elasticsearch/connection.py index 38ba4046..0795652d 100644 --- a/elasticsearch/connection.py +++ b/elasticsearch/connection.py @@ -45,11 +45,10 @@ class Connection(object): def _pretty_json(data): # pretty JSON in tracer curl logs try: - data = json.dumps(json.loads(data), sort_keys=True, indent=2, separators=(',', ': ')) - except ValueError: + return json.dumps(json.loads(data), sort_keys=True, indent=2, separators=(',', ': ')) + except (ValueError, TypeError): # non-json data or a bulk request - pass - return data + return repr(data) logger.info( '%s %s [status:%s request:%.3fs]', method, full_url, @@ -149,3 +148,65 @@ class Urllib3HttpConnection(Connection): raw_data, duration) return response.status, raw_data + +class MemcachedConnection(Connection): + transport_schema = 'memcached' + + method_map = { + 'PUT': 'set', + 'POST': 'set', + 'DELETE': 'delete', + 'HEAD': 'get', + 'GET': 'get', + } + + def __init__(self, host='localhost', port=11211, **kwargs): + try: + import pylibmc + except ImportError: + raise ImproperlyConfigured("You need to install pylibmc to use the MemcachedConnection class.") + super(MemcachedConnection, self).__init__(host=host, port=port, **kwargs) + self.mc = pylibmc.Client(['%s:%s' % (host, port)],behaviors={"tcp_nodelay": True}) + + def perform_request(self, method, url, params=None, body=None, timeout=None): + url = self.url_prefix + url + if params: + url = '%s?%s' % (url, urlencode(params or {})) + full_url = self.host + url + + mc_method = self.method_map.get(method, 'get') + + start = time.time() + try: + status = 200 + if mc_method == 'set': + # no response from set commands + response = '' + if not json.dumps(self.mc.set(url, body)): + status = 500 + else: + response = self.mc.get(url) + + duration = time.time() - start + if response: + response = response.decode('utf-8') + except Exception as e: + self.log_request_fail(method, full_url, time.time() - start, exception=e) + raise ConnectionError('N/A', str(e), e) + + # try not to load the json every time + if response and response[0] == '{' and ('"status"' in response or '"error"' in response): + data = json.loads(response) + if 'status' in data: + status = data['status'] + elif 'error' in data: + raise TransportError('N/A', data['error']) + + if not (200 <= status < 300): + self.log_request_fail(method, url, duration, status) + self._raise_error(status, response) + + self.log_request_success(method, full_url, url, body, status, + response, duration) + + return status, response diff --git a/test_elasticsearch/test_server/test_memcached.py b/test_elasticsearch/test_server/test_memcached.py new file mode 100644 index 00000000..cab84ae0 --- /dev/null +++ b/test_elasticsearch/test_server/test_memcached.py @@ -0,0 +1,38 @@ +# -*- coding: utf-8 -*- +from elasticsearch import Elasticsearch, MemcachedConnection, NotFoundError + +from elasticsearch.transport import ADDRESS_RE + +from . import ElasticTestCase + +from unittest import SkipTest + +class TestMemcachedConnection(ElasticTestCase): + def setUp(self): + super(TestMemcachedConnection, self).setUp() + nodes = self.client.cluster.node_info() + for node_id, node_info in nodes["nodes"].items(): + if 'memcached_address' in node_info: + connection_info = ADDRESS_RE.search(node_info['memcached_address']).groupdict() + self.mc_client = Elasticsearch( + [connection_info], + connection_class=MemcachedConnection + ) + break + else: + raise SkipTest("No memcached plugin.") + + def test_index(self): + self.mc_client.index("test_index", "test_type", {"answer": 42}, id=1) + self.assertTrue(self.client.exists("test_index", doc_type="test_type", id=1)) + + def test_get(self): + self.client.index("test_index", "test_type", {"answer": 42}, id=1) + self.assertEquals({"answer": 42}, self.mc_client.get("test_index", doc_type="test_type", id=1)["_source"]) + + def test_unicode(self): + self.mc_client.index("test_index", "test_type", {"answer": u"你好"}, id=u"你好") + self.assertEquals({"answer": u"你好"}, self.mc_client.get("test_index", doc_type="test_type", id=u"你好")["_source"]) + + def test_missing(self): + self.assertRaises(NotFoundError, self.mc_client.get, "test_index", doc_type="test_type", id=42)