方法参数支持Java的对象类型
diff --git a/dubbo/client.py b/dubbo/client.py index 6b00944..5fb2232 100644 --- a/dubbo/client.py +++ b/dubbo/client.py
@@ -43,6 +43,14 @@ 1. 对于没有参数的方法,此参数不填; 2. 对于只有一个参数的方法,直接填入该参数; 3. 对于有多个参数的方法,传入一个包含了所有参数的列表; + 4. 当前方法参数支持以下类型: + * bool + * int + * long + * float + * double + * java.lang.String + * java.lang.Object :param timeout: 请求超时时间(秒),不设置则不会超时。默认不设置,如无特殊需求不建议设置 * 不设置超时时间在某些极限情况下可能导致此连接一直阻塞; * 设置超时时间会增加远程调用的时间;
diff --git a/dubbo/codec/encoder.py b/dubbo/codec/encoder.py index cce9a43..4f1ee86 100644 --- a/dubbo/codec/encoder.py +++ b/dubbo/codec/encoder.py
@@ -1,4 +1,15 @@ # -*- coding: utf-8 -*- +""" +把Python的数据结构根据Hessian协议序列化为相应的字节数组 +当前支持的数据类型: +* bool +* int +* long +* float +* double +* java.lang.String +* java.lang.Object +""" import struct from dubbo.common.constants import DEFAULT_REQUEST_META, INT_DIRECT_MAX, INT_DIRECT_MIN, BC_INT_ZERO, INT_BYTE_MAX, \ @@ -9,167 +20,247 @@ from dubbo.common.util import double_to_long_bits, num_2_byte_list -def encode(request): +class Object(object): """ - 把请求序列化为字节数组 - :param request: - :return: + 创建一个Java对象 """ - request_body = _encode_request_body(request) - request_head = DEFAULT_REQUEST_META + _get_request_body_length(request_body) - return bytearray(request_head + request_body) + + def __init__(self, path): + """ + :param path: Java对象的路径,例如:java.lang.Object + """ + if not isinstance(path, str): + raise ValueError('Object path {} should be string type.'.format(path)) + self.__path = path + self.__values = {} + + def __getitem__(self, key): + return self.__values[key] + + def __setitem__(self, key, value): + if not isinstance(key, str): + raise ValueError('Object key {} should be string type.'.format(key)) + self.__values[key] = value + + def __delitem__(self, key): + del self.__values[key] + + def __repr__(self): + return '<{}>'.format(self.__path) + + def __contains__(self, key): + return key in self.__values + + def keys(self): + return self.__values.keys() + + def get_path(self): + return self.__path -def _encode_request_body(body): - """ - 对所有的已知的参数根据dubbo协议进行编码 - :param body: - :return: - """ - dubbo_version = body['dubbo_version'] - path = body['path'] - version = body['version'] - method = body['method'] - arguments = body['arguments'] +class Request(object): + def __init__(self, request): + self.__body = request + self.__classes = [] - parameter_types = '' - # 判断并得出参数的类型 - for argument in arguments: - if isinstance(argument, bool): # bool类型的判断必须放在int类型判断的前面 - parameter_types += 'Z' - elif isinstance(argument, int): - if MIN_INT_32 <= argument <= MAX_INT_32: - parameter_types += 'I' + def encode(self): + """ + 把请求序列化为字节数组 + :return: + """ + request_body = self._encode_request_body() + request_head = DEFAULT_REQUEST_META + get_request_body_length(request_body) + return bytearray(request_head + request_body) + + @staticmethod + def _get_parameter_types(arguments): + """ + 针对所有的参数计算得到参数类型字符串 + :param arguments: + :return: + """ + parameter_types = '' + # 判断并得出参数的类型 + for argument in arguments: + if isinstance(argument, bool): # bool类型的判断必须放在int类型判断的前面 + parameter_types += 'Z' + elif isinstance(argument, int): + if MIN_INT_32 <= argument <= MAX_INT_32: + parameter_types += 'I' + else: + parameter_types += 'J' + elif isinstance(argument, float): + parameter_types += 'D' + elif isinstance(argument, str): + parameter_types += 'Ljava/lang/String;' + elif isinstance(argument, Object): + path = argument.get_path() + path = 'L' + path.replace('.', '/') + ';' + parameter_types += path else: - parameter_types += 'J' - elif isinstance(argument, float): - parameter_types += 'D' - elif isinstance(argument, str): - parameter_types += 'Ljava/lang/String;' - else: - raise HessianTypeError('Unknown argument type: {0}'.format(argument)) + raise HessianTypeError('Unknown argument type: {0}'.format(argument)) + return parameter_types - body = [] - body.extend(_encode_single_value(dubbo_version)) - body.extend(_encode_single_value(path)) - body.extend(_encode_single_value(version)) - body.extend(_encode_single_value(method)) - body.extend(_encode_single_value(parameter_types)) - for argument in arguments: - body.extend(_encode_single_value(argument)) + def _encode_request_body(self): + """ + 对所有的已知的参数根据dubbo协议进行编码 + :return: + """ + dubbo_version = self.__body['dubbo_version'] + path = self.__body['path'] + version = self.__body['version'] + method = self.__body['method'] + arguments = self.__body['arguments'] - attachments = { - 'path': path, - 'interface': path, - 'version': version - } - # attachments参数以H开头,以Z结尾 - body.append(ord('H')) - for key in attachments.keys(): - value = attachments[key] - body.extend(_encode_single_value(key)) - body.extend(_encode_single_value(value)) - body.append(ord('Z')) + body = [] + body.extend(self._encode_single_value(dubbo_version)) + body.extend(self._encode_single_value(path)) + body.extend(self._encode_single_value(version)) + body.extend(self._encode_single_value(method)) + body.extend(self._encode_single_value(self._get_parameter_types(arguments))) + for argument in arguments: + body.extend(self._encode_single_value(argument)) - # 因为在上面的逻辑中没有对byte大小进行检测,所以在这里进行统一的处理 - for i in range(len(body)): - body[i] = body[i] & 0xff - return body + attachments = { + 'path': path, + 'interface': path, + 'version': version + } + # attachments参数以H开头,以Z结尾 + body.append(ord('H')) + for key in attachments.keys(): + value = attachments[key] + body.extend(self._encode_single_value(key)) + body.extend(self._encode_single_value(value)) + body.append(ord('Z')) + # 因为在上面的逻辑中没有对byte大小进行检测,所以在这里进行统一的处理 + for i in range(len(body)): + body[i] = body[i] & 0xff + return body -def _encode_single_value(value): - """ - 根据hessian协议对单个变量进行编码 - :param value: - :return: - """ - result = [] - if isinstance(value, bool): - if value: - result.append(ord('T')) - else: - result.append(ord('F')) - return result - elif isinstance(value, int): - if value > MAX_INT_32 or value < MIN_INT_32: - result.append(ord('L')) - result.extend(list(bytearray(struct.pack('>q', value)))) + def _encode_single_value(self, value): + """ + 根据hessian协议对单个变量进行编码 + :param value: + :return: + """ + result = [] + # 布尔类型 + if isinstance(value, bool): + if value: + result.append(ord('T')) + else: + result.append(ord('F')) return result - - if INT_DIRECT_MIN <= value <= INT_DIRECT_MAX: - result.append(value + BC_INT_ZERO) - elif INT_BYTE_MIN <= value <= INT_BYTE_MAX: - result.append(BC_INT_BYTE_ZERO + (value >> 8)) - result.append(value) - elif INT_SHORT_MIN <= value <= INT_SHORT_MAX: - result.append(BC_INT_SHORT_ZERO + (value >> 16)) - result.append(value >> 8) - result.append(value) - else: - result.append(ord('I')) - result.append(value >> 24) - result.append(value >> 16) - result.append(value >> 8) - result.append(value) - return result - elif isinstance(value, float): - int_value = int(value) - if int_value == value: - if int_value == 0: - result.append(BC_DOUBLE_ZERO) - return result - elif int_value == 1: - result.append(BC_DOUBLE_ONE) - return result - elif -0x80 <= int_value < 0x80: - result.append(BC_DOUBLE_BYTE) - result.append(int_value) - return result - elif -0x8000 <= int_value < 0x8000: - result.append(BC_DOUBLE_SHORT) - result.append(int_value >> 8) - result.append(int_value) + # 整型(包括长整型) + elif isinstance(value, int): + if value > MAX_INT_32 or value < MIN_INT_32: + result.append(ord('L')) + result.extend(list(bytearray(struct.pack('>q', value)))) return result - mills = int(value * 1000) - if 0.001 * mills == value and MIN_INT_32 <= mills <= MAX_INT_32: - result.append(BC_DOUBLE_MILL) - result.append(mills >> 24) - result.append(mills >> 16) - result.append(mills >> 8) - result.append(mills) + if INT_DIRECT_MIN <= value <= INT_DIRECT_MAX: + result.append(value + BC_INT_ZERO) + elif INT_BYTE_MIN <= value <= INT_BYTE_MAX: + result.append(BC_INT_BYTE_ZERO + (value >> 8)) + result.append(value) + elif INT_SHORT_MIN <= value <= INT_SHORT_MAX: + result.append(BC_INT_SHORT_ZERO + (value >> 16)) + result.append(value >> 8) + result.append(value) + else: + result.append(ord('I')) + result.append(value >> 24) + result.append(value >> 16) + result.append(value >> 8) + result.append(value) return result + # 浮点类型 + elif isinstance(value, float): + int_value = int(value) + if int_value == value: + if int_value == 0: + result.append(BC_DOUBLE_ZERO) + return result + elif int_value == 1: + result.append(BC_DOUBLE_ONE) + return result + elif -0x80 <= int_value < 0x80: + result.append(BC_DOUBLE_BYTE) + result.append(int_value) + return result + elif -0x8000 <= int_value < 0x8000: + result.append(BC_DOUBLE_SHORT) + result.append(int_value >> 8) + result.append(int_value) + return result - bits = double_to_long_bits(value) - result.append(ord('D')) - result.append(bits >> 56) - result.append(bits >> 48) - result.append(bits >> 40) - result.append(bits >> 32) - result.append(bits >> 24) - result.append(bits >> 16) - result.append(bits >> 8) - result.append(bits) - return result - elif isinstance(value, str): - # 根据hessian协议这里的长度必须是字符串长度而不是字节长度,所以需要Unicode类型 - length = len(value.decode('utf-8')) - if length <= STRING_DIRECT_MAX: - result.append(BC_STRING_DIRECT + length) - elif length <= STRING_SHORT_MAX: - result.append(BC_STRING_SHORT + (length >> 8)) - result.append(length) + mills = int(value * 1000) + if 0.001 * mills == value and MIN_INT_32 <= mills <= MAX_INT_32: + result.append(BC_DOUBLE_MILL) + result.append(mills >> 24) + result.append(mills >> 16) + result.append(mills >> 8) + result.append(mills) + return result + + bits = double_to_long_bits(value) + result.append(ord('D')) + result.append(bits >> 56) + result.append(bits >> 48) + result.append(bits >> 40) + result.append(bits >> 32) + result.append(bits >> 24) + result.append(bits >> 16) + result.append(bits >> 8) + result.append(bits) + return result + # 字符串类型 + elif isinstance(value, str): + # 根据hessian协议这里的长度必须是字符串长度而不是字节长度,所以需要Unicode类型 + length = len(value.decode('utf-8')) + if length <= STRING_DIRECT_MAX: + result.append(BC_STRING_DIRECT + length) + elif length <= STRING_SHORT_MAX: + result.append(BC_STRING_SHORT + (length >> 8)) + result.append(length) + else: + result.append(ord('S')) + result.append(length >> 8) + result.append(length) + result.extend(list(bytearray(value))) # 加上变量数组 + return result + # 对象类型 + elif isinstance(value, Object): + path = value.get_path() + field_names = value.keys() + + if path not in self.__classes: + result.append(ord('C')) + result.extend(self._encode_single_value(path)) + + result.extend(self._encode_single_value(len(field_names))) + + for field_name in field_names: + result.extend(self._encode_single_value(field_name)) + self.__classes.append(path) + class_id = self.__classes.index(path) + if class_id <= 0xf: + class_id += 0x60 + class_id &= 0xff + result.append(class_id) + else: + result.append(ord('O')) + result.extend(self._encode_single_value(class_id)) + for field_name in field_names: + result.extend(self._encode_single_value(value[field_name])) + return result else: - result.append(ord('S')) - result.append(length >> 8) - result.append(length) - result.extend(list(bytearray(value))) # 加上变量数组 - return result - else: - raise HessianTypeError('Unknown argument type: {0}'.format(value)) + raise HessianTypeError('Unknown argument type: {0}'.format(value)) -def _get_request_body_length(body): +def get_request_body_length(body): """ 获取body的长度,并将其转为长度为4个字节的字节数组 :param body: @@ -180,3 +271,12 @@ while len(request_body_length) < 4: request_body_length = [0] + request_body_length return request_body_length + + +if __name__ == '__main__': + o = Object('java.lang.Object') + o['name'] = '张三' + o['age'] = 20 + print o.keys() + print '111' in o + print o
diff --git a/dubbo/connection/connections.py b/dubbo/connection/connections.py index 7887904..ca9aa36 100644 --- a/dubbo/connection/connections.py +++ b/dubbo/connection/connections.py
@@ -6,7 +6,7 @@ import time from struct import unpack -from dubbo.codec.encoder import encode +from dubbo.codec.encoder import Request from dubbo.codec.decoder import Response, get_body_length from dubbo.common.constants import CLI_HEARTBEAT_RES_HEAD, CLI_HEARTBEAT_TAIL, CLI_HEARTBEAT_REQ_HEAD from dubbo.common.exceptions import DubboResponseException, DubboRequestTimeoutException @@ -36,7 +36,7 @@ def get(self, host, request_param, timeout=None): conn = self._get_connection(host) - request = encode(request_param) + request = Request(request_param).encode() conn.lock() conn.clear()
diff --git a/tests/run_test.py b/tests/run_test.py index 1b28455..6999398 100644 --- a/tests/run_test.py +++ b/tests/run_test.py
@@ -3,6 +3,7 @@ import unittest from dubbo.client import DubboClient, ZkRegister +from dubbo.codec.encoder import Object from dubbo.common.loggers import init_log @@ -10,12 +11,62 @@ def setUp(self): init_log() # 初始化日志配置,调用端需要自己配置日志属性 - zk = ZkRegister('172.21.4.98:2181') - self.dubbo = DubboClient('me.hourui.echo.provider.Echo', zk_register=zk) - # dubbo = DubboClient('me.hourui.echo.provider.Echo', host='127.0.0.1:20880') + # zk = ZkRegister('172.21.4.98:2181') + # self.dubbo = DubboClient('me.hourui.echo.provider.Echo', zk_register=zk) + self.dubbo = DubboClient('me.hourui.echo.provider.Echo', host='127.0.0.1:20880') def test_run(self): - result = self.dubbo.call('echo23') + new_user = Object('me.hourui.echo.bean.NewUser') + user1 = Object('me.hourui.echo.bean.User1') + user2 = Object('me.hourui.echo.bean.User2') + user3 = Object('me.hourui.echo.bean.User3') + user4 = Object('me.hourui.echo.bean.User4') + user5 = Object('me.hourui.echo.bean.User5') + user6 = Object('me.hourui.echo.bean.User6') + user7 = Object('me.hourui.echo.bean.User7') + user8 = Object('me.hourui.echo.bean.User8') + user9 = Object('me.hourui.echo.bean.User9') + user10 = Object('me.hourui.echo.bean.User10') + user11 = Object('me.hourui.echo.bean.User11') + user12 = Object('me.hourui.echo.bean.User12') + user13 = Object('me.hourui.echo.bean.User13') + + location = Object('me.hourui.echo.bean.Location') + location['province'] = '江苏省' + location['city'] = '南京市' + location['street'] = '软件大道' + + name = Object('me.hourui.echo.bean.Name') + name['firstName'] = '隔壁的' + name['lastName'] = '王叔叔' + + employee = Object('me.hourui.echo.bean.retail.Employee') + employee['id'] = 'A137639' + employee['name'] = '我勒个去居然不能用emoji啊' + + lock = Object('me.hourui.echo.bean.retail.Lock') + lock['lockReason'] = '加锁的原因是什么呢?' + lock['employee'] = employee + lock['locked'] = True + + new_user['user1'] = user1 + new_user['user2'] = user2 + new_user['user3'] = user3 + new_user['user4'] = user4 + new_user['user5'] = user5 + new_user['user6'] = user6 + new_user['user7'] = user7 + new_user['user8'] = user8 + new_user['user9'] = user9 + new_user['user10'] = user10 + new_user['user11'] = user11 + new_user['user12'] = user12 + new_user['user13'] = user13 + new_user['location'] = location + new_user['name'] = name + new_user['lock'] = lock + + result = self.dubbo.call('test1', [new_user, name, '一个傻傻的用于测试的字符串', location, lock]) # result = dubbo.call('echo23') pretty_print(result)