mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
PY-18029 Move thriftpy to third-party dependencies for Python helpers
This commit is contained in:
@@ -1,5 +0,0 @@
|
||||
# PLY package
|
||||
# Author: David Beazley (dave@dabeaz.com)
|
||||
|
||||
__version__ = '3.7'
|
||||
__all__ = ['lex','yacc']
|
||||
@@ -1,908 +0,0 @@
|
||||
# -----------------------------------------------------------------------------
|
||||
# cpp.py
|
||||
#
|
||||
# Author: David Beazley (http://www.dabeaz.com)
|
||||
# Copyright (C) 2007
|
||||
# All rights reserved
|
||||
#
|
||||
# This module implements an ANSI-C style lexical preprocessor for PLY.
|
||||
# -----------------------------------------------------------------------------
|
||||
from __future__ import generators
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Default preprocessor lexer definitions. These tokens are enough to get
|
||||
# a basic preprocessor working. Other modules may import these if they want
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
tokens = (
|
||||
'CPP_ID','CPP_INTEGER', 'CPP_FLOAT', 'CPP_STRING', 'CPP_CHAR', 'CPP_WS', 'CPP_COMMENT1', 'CPP_COMMENT2', 'CPP_POUND','CPP_DPOUND'
|
||||
)
|
||||
|
||||
literals = "+-*/%|&~^<>=!?()[]{}.,;:\\\'\""
|
||||
|
||||
# Whitespace
|
||||
def t_CPP_WS(t):
|
||||
r'\s+'
|
||||
t.lexer.lineno += t.value.count("\n")
|
||||
return t
|
||||
|
||||
t_CPP_POUND = r'\#'
|
||||
t_CPP_DPOUND = r'\#\#'
|
||||
|
||||
# Identifier
|
||||
t_CPP_ID = r'[A-Za-z_][\w_]*'
|
||||
|
||||
# Integer literal
|
||||
def CPP_INTEGER(t):
|
||||
r'(((((0x)|(0X))[0-9a-fA-F]+)|(\d+))([uU][lL]|[lL][uU]|[uU]|[lL])?)'
|
||||
return t
|
||||
|
||||
t_CPP_INTEGER = CPP_INTEGER
|
||||
|
||||
# Floating literal
|
||||
t_CPP_FLOAT = r'((\d+)(\.\d+)(e(\+|-)?(\d+))? | (\d+)e(\+|-)?(\d+))([lL]|[fF])?'
|
||||
|
||||
# String literal
|
||||
def t_CPP_STRING(t):
|
||||
r'\"([^\\\n]|(\\(.|\n)))*?\"'
|
||||
t.lexer.lineno += t.value.count("\n")
|
||||
return t
|
||||
|
||||
# Character constant 'c' or L'c'
|
||||
def t_CPP_CHAR(t):
|
||||
r'(L)?\'([^\\\n]|(\\(.|\n)))*?\''
|
||||
t.lexer.lineno += t.value.count("\n")
|
||||
return t
|
||||
|
||||
# Comment
|
||||
def t_CPP_COMMENT1(t):
|
||||
r'(/\*(.|\n)*?\*/)'
|
||||
ncr = t.value.count("\n")
|
||||
t.lexer.lineno += ncr
|
||||
# replace with one space or a number of '\n'
|
||||
t.type = 'CPP_WS'; t.value = '\n' * ncr if ncr else ' '
|
||||
return t
|
||||
|
||||
# Line comment
|
||||
def t_CPP_COMMENT2(t):
|
||||
r'(//.*?(\n|$))'
|
||||
# replace with '/n'
|
||||
t.type = 'CPP_WS'; t.value = '\n'
|
||||
|
||||
def t_error(t):
|
||||
t.type = t.value[0]
|
||||
t.value = t.value[0]
|
||||
t.lexer.skip(1)
|
||||
return t
|
||||
|
||||
import re
|
||||
import copy
|
||||
import time
|
||||
import os.path
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# trigraph()
|
||||
#
|
||||
# Given an input string, this function replaces all trigraph sequences.
|
||||
# The following mapping is used:
|
||||
#
|
||||
# ??= #
|
||||
# ??/ \
|
||||
# ??' ^
|
||||
# ??( [
|
||||
# ??) ]
|
||||
# ??! |
|
||||
# ??< {
|
||||
# ??> }
|
||||
# ??- ~
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
_trigraph_pat = re.compile(r'''\?\?[=/\'\(\)\!<>\-]''')
|
||||
_trigraph_rep = {
|
||||
'=':'#',
|
||||
'/':'\\',
|
||||
"'":'^',
|
||||
'(':'[',
|
||||
')':']',
|
||||
'!':'|',
|
||||
'<':'{',
|
||||
'>':'}',
|
||||
'-':'~'
|
||||
}
|
||||
|
||||
def trigraph(input):
|
||||
return _trigraph_pat.sub(lambda g: _trigraph_rep[g.group()[-1]],input)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Macro object
|
||||
#
|
||||
# This object holds information about preprocessor macros
|
||||
#
|
||||
# .name - Macro name (string)
|
||||
# .value - Macro value (a list of tokens)
|
||||
# .arglist - List of argument names
|
||||
# .variadic - Boolean indicating whether or not variadic macro
|
||||
# .vararg - Name of the variadic parameter
|
||||
#
|
||||
# When a macro is created, the macro replacement token sequence is
|
||||
# pre-scanned and used to create patch lists that are later used
|
||||
# during macro expansion
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
class Macro(object):
|
||||
def __init__(self,name,value,arglist=None,variadic=False):
|
||||
self.name = name
|
||||
self.value = value
|
||||
self.arglist = arglist
|
||||
self.variadic = variadic
|
||||
if variadic:
|
||||
self.vararg = arglist[-1]
|
||||
self.source = None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Preprocessor object
|
||||
#
|
||||
# Object representing a preprocessor. Contains macro definitions,
|
||||
# include directories, and other information
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
class Preprocessor(object):
|
||||
def __init__(self,lexer=None):
|
||||
if lexer is None:
|
||||
lexer = lex.lexer
|
||||
self.lexer = lexer
|
||||
self.macros = { }
|
||||
self.path = []
|
||||
self.temp_path = []
|
||||
|
||||
# Probe the lexer for selected tokens
|
||||
self.lexprobe()
|
||||
|
||||
tm = time.localtime()
|
||||
self.define("__DATE__ \"%s\"" % time.strftime("%b %d %Y",tm))
|
||||
self.define("__TIME__ \"%s\"" % time.strftime("%H:%M:%S",tm))
|
||||
self.parser = None
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# tokenize()
|
||||
#
|
||||
# Utility function. Given a string of text, tokenize into a list of tokens
|
||||
# -----------------------------------------------------------------------------
|
||||
|
||||
def tokenize(self,text):
|
||||
tokens = []
|
||||
self.lexer.input(text)
|
||||
while True:
|
||||
tok = self.lexer.token()
|
||||
if not tok: break
|
||||
tokens.append(tok)
|
||||
return tokens
|
||||
|
||||
# ---------------------------------------------------------------------
|
||||
# error()
|
||||
#
|
||||
# Report a preprocessor error/warning of some kind
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def error(self,file,line,msg):
|
||||
print("%s:%d %s" % (file,line,msg))
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# lexprobe()
|
||||
#
|
||||
# This method probes the preprocessor lexer object to discover
|
||||
# the token types of symbols that are important to the preprocessor.
|
||||
# If this works right, the preprocessor will simply "work"
|
||||
# with any suitable lexer regardless of how tokens have been named.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def lexprobe(self):
|
||||
|
||||
# Determine the token type for identifiers
|
||||
self.lexer.input("identifier")
|
||||
tok = self.lexer.token()
|
||||
if not tok or tok.value != "identifier":
|
||||
print("Couldn't determine identifier type")
|
||||
else:
|
||||
self.t_ID = tok.type
|
||||
|
||||
# Determine the token type for integers
|
||||
self.lexer.input("12345")
|
||||
tok = self.lexer.token()
|
||||
if not tok or int(tok.value) != 12345:
|
||||
print("Couldn't determine integer type")
|
||||
else:
|
||||
self.t_INTEGER = tok.type
|
||||
self.t_INTEGER_TYPE = type(tok.value)
|
||||
|
||||
# Determine the token type for strings enclosed in double quotes
|
||||
self.lexer.input("\"filename\"")
|
||||
tok = self.lexer.token()
|
||||
if not tok or tok.value != "\"filename\"":
|
||||
print("Couldn't determine string type")
|
||||
else:
|
||||
self.t_STRING = tok.type
|
||||
|
||||
# Determine the token type for whitespace--if any
|
||||
self.lexer.input(" ")
|
||||
tok = self.lexer.token()
|
||||
if not tok or tok.value != " ":
|
||||
self.t_SPACE = None
|
||||
else:
|
||||
self.t_SPACE = tok.type
|
||||
|
||||
# Determine the token type for newlines
|
||||
self.lexer.input("\n")
|
||||
tok = self.lexer.token()
|
||||
if not tok or tok.value != "\n":
|
||||
self.t_NEWLINE = None
|
||||
print("Couldn't determine token for newlines")
|
||||
else:
|
||||
self.t_NEWLINE = tok.type
|
||||
|
||||
self.t_WS = (self.t_SPACE, self.t_NEWLINE)
|
||||
|
||||
# Check for other characters used by the preprocessor
|
||||
chars = [ '<','>','#','##','\\','(',')',',','.']
|
||||
for c in chars:
|
||||
self.lexer.input(c)
|
||||
tok = self.lexer.token()
|
||||
if not tok or tok.value != c:
|
||||
print("Unable to lex '%s' required for preprocessor" % c)
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# add_path()
|
||||
#
|
||||
# Adds a search path to the preprocessor.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def add_path(self,path):
|
||||
self.path.append(path)
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# group_lines()
|
||||
#
|
||||
# Given an input string, this function splits it into lines. Trailing whitespace
|
||||
# is removed. Any line ending with \ is grouped with the next line. This
|
||||
# function forms the lowest level of the preprocessor---grouping into text into
|
||||
# a line-by-line format.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def group_lines(self,input):
|
||||
lex = self.lexer.clone()
|
||||
lines = [x.rstrip() for x in input.splitlines()]
|
||||
for i in xrange(len(lines)):
|
||||
j = i+1
|
||||
while lines[i].endswith('\\') and (j < len(lines)):
|
||||
lines[i] = lines[i][:-1]+lines[j]
|
||||
lines[j] = ""
|
||||
j += 1
|
||||
|
||||
input = "\n".join(lines)
|
||||
lex.input(input)
|
||||
lex.lineno = 1
|
||||
|
||||
current_line = []
|
||||
while True:
|
||||
tok = lex.token()
|
||||
if not tok:
|
||||
break
|
||||
current_line.append(tok)
|
||||
if tok.type in self.t_WS and '\n' in tok.value:
|
||||
yield current_line
|
||||
current_line = []
|
||||
|
||||
if current_line:
|
||||
yield current_line
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# tokenstrip()
|
||||
#
|
||||
# Remove leading/trailing whitespace tokens from a token list
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def tokenstrip(self,tokens):
|
||||
i = 0
|
||||
while i < len(tokens) and tokens[i].type in self.t_WS:
|
||||
i += 1
|
||||
del tokens[:i]
|
||||
i = len(tokens)-1
|
||||
while i >= 0 and tokens[i].type in self.t_WS:
|
||||
i -= 1
|
||||
del tokens[i+1:]
|
||||
return tokens
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# collect_args()
|
||||
#
|
||||
# Collects comma separated arguments from a list of tokens. The arguments
|
||||
# must be enclosed in parenthesis. Returns a tuple (tokencount,args,positions)
|
||||
# where tokencount is the number of tokens consumed, args is a list of arguments,
|
||||
# and positions is a list of integers containing the starting index of each
|
||||
# argument. Each argument is represented by a list of tokens.
|
||||
#
|
||||
# When collecting arguments, leading and trailing whitespace is removed
|
||||
# from each argument.
|
||||
#
|
||||
# This function properly handles nested parenthesis and commas---these do not
|
||||
# define new arguments.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def collect_args(self,tokenlist):
|
||||
args = []
|
||||
positions = []
|
||||
current_arg = []
|
||||
nesting = 1
|
||||
tokenlen = len(tokenlist)
|
||||
|
||||
# Search for the opening '('.
|
||||
i = 0
|
||||
while (i < tokenlen) and (tokenlist[i].type in self.t_WS):
|
||||
i += 1
|
||||
|
||||
if (i < tokenlen) and (tokenlist[i].value == '('):
|
||||
positions.append(i+1)
|
||||
else:
|
||||
self.error(self.source,tokenlist[0].lineno,"Missing '(' in macro arguments")
|
||||
return 0, [], []
|
||||
|
||||
i += 1
|
||||
|
||||
while i < tokenlen:
|
||||
t = tokenlist[i]
|
||||
if t.value == '(':
|
||||
current_arg.append(t)
|
||||
nesting += 1
|
||||
elif t.value == ')':
|
||||
nesting -= 1
|
||||
if nesting == 0:
|
||||
if current_arg:
|
||||
args.append(self.tokenstrip(current_arg))
|
||||
positions.append(i)
|
||||
return i+1,args,positions
|
||||
current_arg.append(t)
|
||||
elif t.value == ',' and nesting == 1:
|
||||
args.append(self.tokenstrip(current_arg))
|
||||
positions.append(i+1)
|
||||
current_arg = []
|
||||
else:
|
||||
current_arg.append(t)
|
||||
i += 1
|
||||
|
||||
# Missing end argument
|
||||
self.error(self.source,tokenlist[-1].lineno,"Missing ')' in macro arguments")
|
||||
return 0, [],[]
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# macro_prescan()
|
||||
#
|
||||
# Examine the macro value (token sequence) and identify patch points
|
||||
# This is used to speed up macro expansion later on---we'll know
|
||||
# right away where to apply patches to the value to form the expansion
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def macro_prescan(self,macro):
|
||||
macro.patch = [] # Standard macro arguments
|
||||
macro.str_patch = [] # String conversion expansion
|
||||
macro.var_comma_patch = [] # Variadic macro comma patch
|
||||
i = 0
|
||||
while i < len(macro.value):
|
||||
if macro.value[i].type == self.t_ID and macro.value[i].value in macro.arglist:
|
||||
argnum = macro.arglist.index(macro.value[i].value)
|
||||
# Conversion of argument to a string
|
||||
if i > 0 and macro.value[i-1].value == '#':
|
||||
macro.value[i] = copy.copy(macro.value[i])
|
||||
macro.value[i].type = self.t_STRING
|
||||
del macro.value[i-1]
|
||||
macro.str_patch.append((argnum,i-1))
|
||||
continue
|
||||
# Concatenation
|
||||
elif (i > 0 and macro.value[i-1].value == '##'):
|
||||
macro.patch.append(('c',argnum,i-1))
|
||||
del macro.value[i-1]
|
||||
continue
|
||||
elif ((i+1) < len(macro.value) and macro.value[i+1].value == '##'):
|
||||
macro.patch.append(('c',argnum,i))
|
||||
i += 1
|
||||
continue
|
||||
# Standard expansion
|
||||
else:
|
||||
macro.patch.append(('e',argnum,i))
|
||||
elif macro.value[i].value == '##':
|
||||
if macro.variadic and (i > 0) and (macro.value[i-1].value == ',') and \
|
||||
((i+1) < len(macro.value)) and (macro.value[i+1].type == self.t_ID) and \
|
||||
(macro.value[i+1].value == macro.vararg):
|
||||
macro.var_comma_patch.append(i-1)
|
||||
i += 1
|
||||
macro.patch.sort(key=lambda x: x[2],reverse=True)
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# macro_expand_args()
|
||||
#
|
||||
# Given a Macro and list of arguments (each a token list), this method
|
||||
# returns an expanded version of a macro. The return value is a token sequence
|
||||
# representing the replacement macro tokens
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def macro_expand_args(self,macro,args):
|
||||
# Make a copy of the macro token sequence
|
||||
rep = [copy.copy(_x) for _x in macro.value]
|
||||
|
||||
# Make string expansion patches. These do not alter the length of the replacement sequence
|
||||
|
||||
str_expansion = {}
|
||||
for argnum, i in macro.str_patch:
|
||||
if argnum not in str_expansion:
|
||||
str_expansion[argnum] = ('"%s"' % "".join([x.value for x in args[argnum]])).replace("\\","\\\\")
|
||||
rep[i] = copy.copy(rep[i])
|
||||
rep[i].value = str_expansion[argnum]
|
||||
|
||||
# Make the variadic macro comma patch. If the variadic macro argument is empty, we get rid
|
||||
comma_patch = False
|
||||
if macro.variadic and not args[-1]:
|
||||
for i in macro.var_comma_patch:
|
||||
rep[i] = None
|
||||
comma_patch = True
|
||||
|
||||
# Make all other patches. The order of these matters. It is assumed that the patch list
|
||||
# has been sorted in reverse order of patch location since replacements will cause the
|
||||
# size of the replacement sequence to expand from the patch point.
|
||||
|
||||
expanded = { }
|
||||
for ptype, argnum, i in macro.patch:
|
||||
# Concatenation. Argument is left unexpanded
|
||||
if ptype == 'c':
|
||||
rep[i:i+1] = args[argnum]
|
||||
# Normal expansion. Argument is macro expanded first
|
||||
elif ptype == 'e':
|
||||
if argnum not in expanded:
|
||||
expanded[argnum] = self.expand_macros(args[argnum])
|
||||
rep[i:i+1] = expanded[argnum]
|
||||
|
||||
# Get rid of removed comma if necessary
|
||||
if comma_patch:
|
||||
rep = [_i for _i in rep if _i]
|
||||
|
||||
return rep
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# expand_macros()
|
||||
#
|
||||
# Given a list of tokens, this function performs macro expansion.
|
||||
# The expanded argument is a dictionary that contains macros already
|
||||
# expanded. This is used to prevent infinite recursion.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def expand_macros(self,tokens,expanded=None):
|
||||
if expanded is None:
|
||||
expanded = {}
|
||||
i = 0
|
||||
while i < len(tokens):
|
||||
t = tokens[i]
|
||||
if t.type == self.t_ID:
|
||||
if t.value in self.macros and t.value not in expanded:
|
||||
# Yes, we found a macro match
|
||||
expanded[t.value] = True
|
||||
|
||||
m = self.macros[t.value]
|
||||
if not m.arglist:
|
||||
# A simple macro
|
||||
ex = self.expand_macros([copy.copy(_x) for _x in m.value],expanded)
|
||||
for e in ex:
|
||||
e.lineno = t.lineno
|
||||
tokens[i:i+1] = ex
|
||||
i += len(ex)
|
||||
else:
|
||||
# A macro with arguments
|
||||
j = i + 1
|
||||
while j < len(tokens) and tokens[j].type in self.t_WS:
|
||||
j += 1
|
||||
if tokens[j].value == '(':
|
||||
tokcount,args,positions = self.collect_args(tokens[j:])
|
||||
if not m.variadic and len(args) != len(m.arglist):
|
||||
self.error(self.source,t.lineno,"Macro %s requires %d arguments" % (t.value,len(m.arglist)))
|
||||
i = j + tokcount
|
||||
elif m.variadic and len(args) < len(m.arglist)-1:
|
||||
if len(m.arglist) > 2:
|
||||
self.error(self.source,t.lineno,"Macro %s must have at least %d arguments" % (t.value, len(m.arglist)-1))
|
||||
else:
|
||||
self.error(self.source,t.lineno,"Macro %s must have at least %d argument" % (t.value, len(m.arglist)-1))
|
||||
i = j + tokcount
|
||||
else:
|
||||
if m.variadic:
|
||||
if len(args) == len(m.arglist)-1:
|
||||
args.append([])
|
||||
else:
|
||||
args[len(m.arglist)-1] = tokens[j+positions[len(m.arglist)-1]:j+tokcount-1]
|
||||
del args[len(m.arglist):]
|
||||
|
||||
# Get macro replacement text
|
||||
rep = self.macro_expand_args(m,args)
|
||||
rep = self.expand_macros(rep,expanded)
|
||||
for r in rep:
|
||||
r.lineno = t.lineno
|
||||
tokens[i:j+tokcount] = rep
|
||||
i += len(rep)
|
||||
del expanded[t.value]
|
||||
continue
|
||||
elif t.value == '__LINE__':
|
||||
t.type = self.t_INTEGER
|
||||
t.value = self.t_INTEGER_TYPE(t.lineno)
|
||||
|
||||
i += 1
|
||||
return tokens
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# evalexpr()
|
||||
#
|
||||
# Evaluate an expression token sequence for the purposes of evaluating
|
||||
# integral expressions.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def evalexpr(self,tokens):
|
||||
# tokens = tokenize(line)
|
||||
# Search for defined macros
|
||||
i = 0
|
||||
while i < len(tokens):
|
||||
if tokens[i].type == self.t_ID and tokens[i].value == 'defined':
|
||||
j = i + 1
|
||||
needparen = False
|
||||
result = "0L"
|
||||
while j < len(tokens):
|
||||
if tokens[j].type in self.t_WS:
|
||||
j += 1
|
||||
continue
|
||||
elif tokens[j].type == self.t_ID:
|
||||
if tokens[j].value in self.macros:
|
||||
result = "1L"
|
||||
else:
|
||||
result = "0L"
|
||||
if not needparen: break
|
||||
elif tokens[j].value == '(':
|
||||
needparen = True
|
||||
elif tokens[j].value == ')':
|
||||
break
|
||||
else:
|
||||
self.error(self.source,tokens[i].lineno,"Malformed defined()")
|
||||
j += 1
|
||||
tokens[i].type = self.t_INTEGER
|
||||
tokens[i].value = self.t_INTEGER_TYPE(result)
|
||||
del tokens[i+1:j+1]
|
||||
i += 1
|
||||
tokens = self.expand_macros(tokens)
|
||||
for i,t in enumerate(tokens):
|
||||
if t.type == self.t_ID:
|
||||
tokens[i] = copy.copy(t)
|
||||
tokens[i].type = self.t_INTEGER
|
||||
tokens[i].value = self.t_INTEGER_TYPE("0L")
|
||||
elif t.type == self.t_INTEGER:
|
||||
tokens[i] = copy.copy(t)
|
||||
# Strip off any trailing suffixes
|
||||
tokens[i].value = str(tokens[i].value)
|
||||
while tokens[i].value[-1] not in "0123456789abcdefABCDEF":
|
||||
tokens[i].value = tokens[i].value[:-1]
|
||||
|
||||
expr = "".join([str(x.value) for x in tokens])
|
||||
expr = expr.replace("&&"," and ")
|
||||
expr = expr.replace("||"," or ")
|
||||
expr = expr.replace("!"," not ")
|
||||
try:
|
||||
result = eval(expr)
|
||||
except StandardError:
|
||||
self.error(self.source,tokens[0].lineno,"Couldn't evaluate expression")
|
||||
result = 0
|
||||
return result
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# parsegen()
|
||||
#
|
||||
# Parse an input string/
|
||||
# ----------------------------------------------------------------------
|
||||
def parsegen(self,input,source=None):
|
||||
|
||||
# Replace trigraph sequences
|
||||
t = trigraph(input)
|
||||
lines = self.group_lines(t)
|
||||
|
||||
if not source:
|
||||
source = ""
|
||||
|
||||
self.define("__FILE__ \"%s\"" % source)
|
||||
|
||||
self.source = source
|
||||
chunk = []
|
||||
enable = True
|
||||
iftrigger = False
|
||||
ifstack = []
|
||||
|
||||
for x in lines:
|
||||
for i,tok in enumerate(x):
|
||||
if tok.type not in self.t_WS: break
|
||||
if tok.value == '#':
|
||||
# Preprocessor directive
|
||||
|
||||
# insert necessary whitespace instead of eaten tokens
|
||||
for tok in x:
|
||||
if tok.type in self.t_WS and '\n' in tok.value:
|
||||
chunk.append(tok)
|
||||
|
||||
dirtokens = self.tokenstrip(x[i+1:])
|
||||
if dirtokens:
|
||||
name = dirtokens[0].value
|
||||
args = self.tokenstrip(dirtokens[1:])
|
||||
else:
|
||||
name = ""
|
||||
args = []
|
||||
|
||||
if name == 'define':
|
||||
if enable:
|
||||
for tok in self.expand_macros(chunk):
|
||||
yield tok
|
||||
chunk = []
|
||||
self.define(args)
|
||||
elif name == 'include':
|
||||
if enable:
|
||||
for tok in self.expand_macros(chunk):
|
||||
yield tok
|
||||
chunk = []
|
||||
oldfile = self.macros['__FILE__']
|
||||
for tok in self.include(args):
|
||||
yield tok
|
||||
self.macros['__FILE__'] = oldfile
|
||||
self.source = source
|
||||
elif name == 'undef':
|
||||
if enable:
|
||||
for tok in self.expand_macros(chunk):
|
||||
yield tok
|
||||
chunk = []
|
||||
self.undef(args)
|
||||
elif name == 'ifdef':
|
||||
ifstack.append((enable,iftrigger))
|
||||
if enable:
|
||||
if not args[0].value in self.macros:
|
||||
enable = False
|
||||
iftrigger = False
|
||||
else:
|
||||
iftrigger = True
|
||||
elif name == 'ifndef':
|
||||
ifstack.append((enable,iftrigger))
|
||||
if enable:
|
||||
if args[0].value in self.macros:
|
||||
enable = False
|
||||
iftrigger = False
|
||||
else:
|
||||
iftrigger = True
|
||||
elif name == 'if':
|
||||
ifstack.append((enable,iftrigger))
|
||||
if enable:
|
||||
result = self.evalexpr(args)
|
||||
if not result:
|
||||
enable = False
|
||||
iftrigger = False
|
||||
else:
|
||||
iftrigger = True
|
||||
elif name == 'elif':
|
||||
if ifstack:
|
||||
if ifstack[-1][0]: # We only pay attention if outer "if" allows this
|
||||
if enable: # If already true, we flip enable False
|
||||
enable = False
|
||||
elif not iftrigger: # If False, but not triggered yet, we'll check expression
|
||||
result = self.evalexpr(args)
|
||||
if result:
|
||||
enable = True
|
||||
iftrigger = True
|
||||
else:
|
||||
self.error(self.source,dirtokens[0].lineno,"Misplaced #elif")
|
||||
|
||||
elif name == 'else':
|
||||
if ifstack:
|
||||
if ifstack[-1][0]:
|
||||
if enable:
|
||||
enable = False
|
||||
elif not iftrigger:
|
||||
enable = True
|
||||
iftrigger = True
|
||||
else:
|
||||
self.error(self.source,dirtokens[0].lineno,"Misplaced #else")
|
||||
|
||||
elif name == 'endif':
|
||||
if ifstack:
|
||||
enable,iftrigger = ifstack.pop()
|
||||
else:
|
||||
self.error(self.source,dirtokens[0].lineno,"Misplaced #endif")
|
||||
else:
|
||||
# Unknown preprocessor directive
|
||||
pass
|
||||
|
||||
else:
|
||||
# Normal text
|
||||
if enable:
|
||||
chunk.extend(x)
|
||||
|
||||
for tok in self.expand_macros(chunk):
|
||||
yield tok
|
||||
chunk = []
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# include()
|
||||
#
|
||||
# Implementation of file-inclusion
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def include(self,tokens):
|
||||
# Try to extract the filename and then process an include file
|
||||
if not tokens:
|
||||
return
|
||||
if tokens:
|
||||
if tokens[0].value != '<' and tokens[0].type != self.t_STRING:
|
||||
tokens = self.expand_macros(tokens)
|
||||
|
||||
if tokens[0].value == '<':
|
||||
# Include <...>
|
||||
i = 1
|
||||
while i < len(tokens):
|
||||
if tokens[i].value == '>':
|
||||
break
|
||||
i += 1
|
||||
else:
|
||||
print("Malformed #include <...>")
|
||||
return
|
||||
filename = "".join([x.value for x in tokens[1:i]])
|
||||
path = self.path + [""] + self.temp_path
|
||||
elif tokens[0].type == self.t_STRING:
|
||||
filename = tokens[0].value[1:-1]
|
||||
path = self.temp_path + [""] + self.path
|
||||
else:
|
||||
print("Malformed #include statement")
|
||||
return
|
||||
for p in path:
|
||||
iname = os.path.join(p,filename)
|
||||
try:
|
||||
data = open(iname,"r").read()
|
||||
dname = os.path.dirname(iname)
|
||||
if dname:
|
||||
self.temp_path.insert(0,dname)
|
||||
for tok in self.parsegen(data,filename):
|
||||
yield tok
|
||||
if dname:
|
||||
del self.temp_path[0]
|
||||
break
|
||||
except IOError:
|
||||
pass
|
||||
else:
|
||||
print("Couldn't find '%s'" % filename)
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# define()
|
||||
#
|
||||
# Define a new macro
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def define(self,tokens):
|
||||
if isinstance(tokens,(str,unicode)):
|
||||
tokens = self.tokenize(tokens)
|
||||
|
||||
linetok = tokens
|
||||
try:
|
||||
name = linetok[0]
|
||||
if len(linetok) > 1:
|
||||
mtype = linetok[1]
|
||||
else:
|
||||
mtype = None
|
||||
if not mtype:
|
||||
m = Macro(name.value,[])
|
||||
self.macros[name.value] = m
|
||||
elif mtype.type in self.t_WS:
|
||||
# A normal macro
|
||||
m = Macro(name.value,self.tokenstrip(linetok[2:]))
|
||||
self.macros[name.value] = m
|
||||
elif mtype.value == '(':
|
||||
# A macro with arguments
|
||||
tokcount, args, positions = self.collect_args(linetok[1:])
|
||||
variadic = False
|
||||
for a in args:
|
||||
if variadic:
|
||||
print("No more arguments may follow a variadic argument")
|
||||
break
|
||||
astr = "".join([str(_i.value) for _i in a])
|
||||
if astr == "...":
|
||||
variadic = True
|
||||
a[0].type = self.t_ID
|
||||
a[0].value = '__VA_ARGS__'
|
||||
variadic = True
|
||||
del a[1:]
|
||||
continue
|
||||
elif astr[-3:] == "..." and a[0].type == self.t_ID:
|
||||
variadic = True
|
||||
del a[1:]
|
||||
# If, for some reason, "." is part of the identifier, strip off the name for the purposes
|
||||
# of macro expansion
|
||||
if a[0].value[-3:] == '...':
|
||||
a[0].value = a[0].value[:-3]
|
||||
continue
|
||||
if len(a) > 1 or a[0].type != self.t_ID:
|
||||
print("Invalid macro argument")
|
||||
break
|
||||
else:
|
||||
mvalue = self.tokenstrip(linetok[1+tokcount:])
|
||||
i = 0
|
||||
while i < len(mvalue):
|
||||
if i+1 < len(mvalue):
|
||||
if mvalue[i].type in self.t_WS and mvalue[i+1].value == '##':
|
||||
del mvalue[i]
|
||||
continue
|
||||
elif mvalue[i].value == '##' and mvalue[i+1].type in self.t_WS:
|
||||
del mvalue[i+1]
|
||||
i += 1
|
||||
m = Macro(name.value,mvalue,[x[0].value for x in args],variadic)
|
||||
self.macro_prescan(m)
|
||||
self.macros[name.value] = m
|
||||
else:
|
||||
print("Bad macro definition")
|
||||
except LookupError:
|
||||
print("Bad macro definition")
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# undef()
|
||||
#
|
||||
# Undefine a macro
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
def undef(self,tokens):
|
||||
id = tokens[0].value
|
||||
try:
|
||||
del self.macros[id]
|
||||
except LookupError:
|
||||
pass
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# parse()
|
||||
#
|
||||
# Parse input text.
|
||||
# ----------------------------------------------------------------------
|
||||
def parse(self,input,source=None,ignore={}):
|
||||
self.ignore = ignore
|
||||
self.parser = self.parsegen(input,source)
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# token()
|
||||
#
|
||||
# Method to return individual tokens
|
||||
# ----------------------------------------------------------------------
|
||||
def token(self):
|
||||
try:
|
||||
while True:
|
||||
tok = next(self.parser)
|
||||
if tok.type not in self.ignore: return tok
|
||||
except StopIteration:
|
||||
self.parser = None
|
||||
return None
|
||||
|
||||
if __name__ == '__main__':
|
||||
import ply.lex as lex
|
||||
lexer = lex.lex()
|
||||
|
||||
# Run a preprocessor
|
||||
import sys
|
||||
f = open(sys.argv[1])
|
||||
input = f.read()
|
||||
|
||||
p = Preprocessor(lexer)
|
||||
p.parse(input,sys.argv[1])
|
||||
while True:
|
||||
tok = p.token()
|
||||
if not tok: break
|
||||
print(p.source, tok)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,133 +0,0 @@
|
||||
# ----------------------------------------------------------------------
|
||||
# ctokens.py
|
||||
#
|
||||
# Token specifications for symbols in ANSI C and C++. This file is
|
||||
# meant to be used as a library in other tokenizers.
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
# Reserved words
|
||||
|
||||
tokens = [
|
||||
# Literals (identifier, integer constant, float constant, string constant, char const)
|
||||
'ID', 'TYPEID', 'INTEGER', 'FLOAT', 'STRING', 'CHARACTER',
|
||||
|
||||
# Operators (+,-,*,/,%,|,&,~,^,<<,>>, ||, &&, !, <, <=, >, >=, ==, !=)
|
||||
'PLUS', 'MINUS', 'TIMES', 'DIVIDE', 'MODULO',
|
||||
'OR', 'AND', 'NOT', 'XOR', 'LSHIFT', 'RSHIFT',
|
||||
'LOR', 'LAND', 'LNOT',
|
||||
'LT', 'LE', 'GT', 'GE', 'EQ', 'NE',
|
||||
|
||||
# Assignment (=, *=, /=, %=, +=, -=, <<=, >>=, &=, ^=, |=)
|
||||
'EQUALS', 'TIMESEQUAL', 'DIVEQUAL', 'MODEQUAL', 'PLUSEQUAL', 'MINUSEQUAL',
|
||||
'LSHIFTEQUAL','RSHIFTEQUAL', 'ANDEQUAL', 'XOREQUAL', 'OREQUAL',
|
||||
|
||||
# Increment/decrement (++,--)
|
||||
'INCREMENT', 'DECREMENT',
|
||||
|
||||
# Structure dereference (->)
|
||||
'ARROW',
|
||||
|
||||
# Ternary operator (?)
|
||||
'TERNARY',
|
||||
|
||||
# Delimeters ( ) [ ] { } , . ; :
|
||||
'LPAREN', 'RPAREN',
|
||||
'LBRACKET', 'RBRACKET',
|
||||
'LBRACE', 'RBRACE',
|
||||
'COMMA', 'PERIOD', 'SEMI', 'COLON',
|
||||
|
||||
# Ellipsis (...)
|
||||
'ELLIPSIS',
|
||||
]
|
||||
|
||||
# Operators
|
||||
t_PLUS = r'\+'
|
||||
t_MINUS = r'-'
|
||||
t_TIMES = r'\*'
|
||||
t_DIVIDE = r'/'
|
||||
t_MODULO = r'%'
|
||||
t_OR = r'\|'
|
||||
t_AND = r'&'
|
||||
t_NOT = r'~'
|
||||
t_XOR = r'\^'
|
||||
t_LSHIFT = r'<<'
|
||||
t_RSHIFT = r'>>'
|
||||
t_LOR = r'\|\|'
|
||||
t_LAND = r'&&'
|
||||
t_LNOT = r'!'
|
||||
t_LT = r'<'
|
||||
t_GT = r'>'
|
||||
t_LE = r'<='
|
||||
t_GE = r'>='
|
||||
t_EQ = r'=='
|
||||
t_NE = r'!='
|
||||
|
||||
# Assignment operators
|
||||
|
||||
t_EQUALS = r'='
|
||||
t_TIMESEQUAL = r'\*='
|
||||
t_DIVEQUAL = r'/='
|
||||
t_MODEQUAL = r'%='
|
||||
t_PLUSEQUAL = r'\+='
|
||||
t_MINUSEQUAL = r'-='
|
||||
t_LSHIFTEQUAL = r'<<='
|
||||
t_RSHIFTEQUAL = r'>>='
|
||||
t_ANDEQUAL = r'&='
|
||||
t_OREQUAL = r'\|='
|
||||
t_XOREQUAL = r'\^='
|
||||
|
||||
# Increment/decrement
|
||||
t_INCREMENT = r'\+\+'
|
||||
t_DECREMENT = r'--'
|
||||
|
||||
# ->
|
||||
t_ARROW = r'->'
|
||||
|
||||
# ?
|
||||
t_TERNARY = r'\?'
|
||||
|
||||
# Delimeters
|
||||
t_LPAREN = r'\('
|
||||
t_RPAREN = r'\)'
|
||||
t_LBRACKET = r'\['
|
||||
t_RBRACKET = r'\]'
|
||||
t_LBRACE = r'\{'
|
||||
t_RBRACE = r'\}'
|
||||
t_COMMA = r','
|
||||
t_PERIOD = r'\.'
|
||||
t_SEMI = r';'
|
||||
t_COLON = r':'
|
||||
t_ELLIPSIS = r'\.\.\.'
|
||||
|
||||
# Identifiers
|
||||
t_ID = r'[A-Za-z_][A-Za-z0-9_]*'
|
||||
|
||||
# Integer literal
|
||||
t_INTEGER = r'\d+([uU]|[lL]|[uU][lL]|[lL][uU])?'
|
||||
|
||||
# Floating literal
|
||||
t_FLOAT = r'((\d+)(\.\d+)(e(\+|-)?(\d+))? | (\d+)e(\+|-)?(\d+))([lL]|[fF])?'
|
||||
|
||||
# String literal
|
||||
t_STRING = r'\"([^\\\n]|(\\.))*?\"'
|
||||
|
||||
# Character constant 'c' or L'c'
|
||||
t_CHARACTER = r'(L)?\'([^\\\n]|(\\.))*?\''
|
||||
|
||||
# Comment (C-Style)
|
||||
def t_COMMENT(t):
|
||||
r'/\*(.|\n)*?\*/'
|
||||
t.lexer.lineno += t.value.count('\n')
|
||||
return t
|
||||
|
||||
# Comment (C++-Style)
|
||||
def t_CPPCOMMENT(t):
|
||||
r'//.*\n'
|
||||
t.lexer.lineno += 1
|
||||
return t
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,74 +0,0 @@
|
||||
# ply: ygen.py
|
||||
#
|
||||
# This is a support program that auto-generates different versions of the YACC parsing
|
||||
# function with different features removed for the purposes of performance.
|
||||
#
|
||||
# Users should edit the method LParser.parsedebug() in yacc.py. The source code
|
||||
# for that method is then used to create the other methods. See the comments in
|
||||
# yacc.py for further details.
|
||||
|
||||
import os.path
|
||||
import shutil
|
||||
|
||||
def get_source_range(lines, tag):
|
||||
srclines = enumerate(lines)
|
||||
start_tag = '#--! %s-start' % tag
|
||||
end_tag = '#--! %s-end' % tag
|
||||
|
||||
for start_index, line in srclines:
|
||||
if line.strip().startswith(start_tag):
|
||||
break
|
||||
|
||||
for end_index, line in srclines:
|
||||
if line.strip().endswith(end_tag):
|
||||
break
|
||||
|
||||
return (start_index + 1, end_index)
|
||||
|
||||
def filter_section(lines, tag):
|
||||
filtered_lines = []
|
||||
include = True
|
||||
tag_text = '#--! %s' % tag
|
||||
for line in lines:
|
||||
if line.strip().startswith(tag_text):
|
||||
include = not include
|
||||
elif include:
|
||||
filtered_lines.append(line)
|
||||
return filtered_lines
|
||||
|
||||
def main():
|
||||
dirname = os.path.dirname(__file__)
|
||||
shutil.copy2(os.path.join(dirname, 'yacc.py'), os.path.join(dirname, 'yacc.py.bak'))
|
||||
with open(os.path.join(dirname, 'yacc.py'), 'r') as f:
|
||||
lines = f.readlines()
|
||||
|
||||
parse_start, parse_end = get_source_range(lines, 'parsedebug')
|
||||
parseopt_start, parseopt_end = get_source_range(lines, 'parseopt')
|
||||
parseopt_notrack_start, parseopt_notrack_end = get_source_range(lines, 'parseopt-notrack')
|
||||
|
||||
# Get the original source
|
||||
orig_lines = lines[parse_start:parse_end]
|
||||
|
||||
# Filter the DEBUG sections out
|
||||
parseopt_lines = filter_section(orig_lines, 'DEBUG')
|
||||
|
||||
# Filter the TRACKING sections out
|
||||
parseopt_notrack_lines = filter_section(parseopt_lines, 'TRACKING')
|
||||
|
||||
# Replace the parser source sections with updated versions
|
||||
lines[parseopt_notrack_start:parseopt_notrack_end] = parseopt_notrack_lines
|
||||
lines[parseopt_start:parseopt_end] = parseopt_lines
|
||||
|
||||
lines = [line.rstrip()+'\n' for line in lines]
|
||||
with open(os.path.join(dirname, 'yacc.py'), 'w') as f:
|
||||
f.writelines(lines)
|
||||
|
||||
print('Updated yacc.py')
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,11 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import sys
|
||||
|
||||
from .hook import install_import_hook, remove_import_hook
|
||||
from .parser import load, load_module, load_fp
|
||||
|
||||
__version__ = '0.3.8'
|
||||
__python__ = sys.version_info
|
||||
__all__ = ["install_import_hook", "remove_import_hook", "load", "load_module",
|
||||
"load_fp"]
|
||||
@@ -1,126 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
thriftpy._compat
|
||||
~~~~~~~~~~~~~
|
||||
|
||||
py2/py3 compatibility support.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import platform
|
||||
import sys
|
||||
import types
|
||||
|
||||
PY3 = sys.version_info[0] == 3
|
||||
PYPY = "__pypy__" in sys.modules
|
||||
|
||||
UNIX = platform.system() in ("Linux", "Darwin")
|
||||
CYTHON = False # Cython always disabled in pypy and windows
|
||||
|
||||
# only python2.7.9 and python 3.4 or above have true ssl context
|
||||
MODERN_SSL = (2, 7, 9) <= sys.version_info < (3, 0, 0) or \
|
||||
sys.version_info >= (3, 4)
|
||||
|
||||
if PY3:
|
||||
text_type = str
|
||||
string_types = (str,)
|
||||
|
||||
def u(s):
|
||||
return s
|
||||
else:
|
||||
text_type = unicode # noqa
|
||||
string_types = (str, unicode) # noqa
|
||||
|
||||
def u(s):
|
||||
if not isinstance(s, text_type):
|
||||
s = s.decode("utf-8")
|
||||
return s
|
||||
|
||||
|
||||
def with_metaclass(meta, *bases):
|
||||
"""Create a base class with a metaclass for py2 & py3
|
||||
|
||||
This code snippet is copied from six."""
|
||||
# This requires a bit of explanation: the basic idea is to make a
|
||||
# dummy metaclass for one level of class instantiation that replaces
|
||||
# itself with the actual metaclass. Because of internal type checks
|
||||
# we also need to make sure that we downgrade the custom metaclass
|
||||
# for one level to something closer to type (that's why __call__ and
|
||||
# __init__ comes back from type etc.).
|
||||
class metaclass(meta):
|
||||
__call__ = type.__call__
|
||||
__init__ = type.__init__
|
||||
|
||||
def __new__(cls, name, this_bases, d):
|
||||
if this_bases is None:
|
||||
return type.__new__(cls, name, (), d)
|
||||
return meta(name, bases, d)
|
||||
return metaclass('temporary_class', None, {})
|
||||
|
||||
|
||||
def init_func_generator(spec):
|
||||
"""Generate `__init__` function based on TPayload.default_spec
|
||||
|
||||
For example::
|
||||
|
||||
spec = [('name', 'Alice'), ('number', None)]
|
||||
|
||||
will generate::
|
||||
|
||||
def __init__(self, name='Alice', number=None):
|
||||
kwargs = locals()
|
||||
kwargs.pop('self')
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
TODO: The `locals()` part may need refine.
|
||||
"""
|
||||
if not spec:
|
||||
def __init__(self):
|
||||
pass
|
||||
return __init__
|
||||
|
||||
varnames, defaults = zip(*spec)
|
||||
varnames = ('self', ) + varnames
|
||||
|
||||
def init(self):
|
||||
self.__dict__ = locals().copy()
|
||||
del self.__dict__['self']
|
||||
|
||||
code = init.__code__
|
||||
if PY3:
|
||||
new_code = types.CodeType(len(varnames),
|
||||
0,
|
||||
len(varnames),
|
||||
code.co_stacksize,
|
||||
code.co_flags,
|
||||
code.co_code,
|
||||
code.co_consts,
|
||||
code.co_names,
|
||||
varnames,
|
||||
code.co_filename,
|
||||
"__init__",
|
||||
code.co_firstlineno,
|
||||
code.co_lnotab,
|
||||
code.co_freevars,
|
||||
code.co_cellvars)
|
||||
else:
|
||||
new_code = types.CodeType(len(varnames),
|
||||
len(varnames),
|
||||
code.co_stacksize,
|
||||
code.co_flags,
|
||||
code.co_code,
|
||||
code.co_consts,
|
||||
code.co_names,
|
||||
varnames,
|
||||
code.co_filename,
|
||||
"__init__",
|
||||
code.co_firstlineno,
|
||||
code.co_lnotab,
|
||||
code.co_freevars,
|
||||
code.co_cellvars)
|
||||
|
||||
return types.FunctionType(new_code,
|
||||
{"__builtins__": __builtins__},
|
||||
argdefs=defaults)
|
||||
@@ -1,201 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
Tracking support similar to twitter finagle-thrift.
|
||||
|
||||
Note: When using tracking, every client should have a corresponding
|
||||
server processor.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import os.path
|
||||
import time
|
||||
|
||||
from ...parser import load
|
||||
from ...thrift import TClient, TApplicationException, TMessageType, \
|
||||
TProcessor, TType
|
||||
|
||||
track_method = "__thriftpy_tracing_method_name__v2"
|
||||
track_thrift = load(os.path.join(os.path.dirname(__file__), "tracking.thrift"))
|
||||
|
||||
|
||||
__all__ = ["TTrackedClient", "TTrackedProcessor", "TrackerBase",
|
||||
"ConsoleTracker"]
|
||||
|
||||
|
||||
class RequestInfo(object):
|
||||
def __init__(self, request_id, api, seq, client, server, status, start,
|
||||
end, annotation, meta):
|
||||
"""Used to store call info.
|
||||
|
||||
:request_id: used to identity a request
|
||||
:api: api name
|
||||
:seq: sequence number
|
||||
:client: client name
|
||||
:server: server name
|
||||
:status: request status
|
||||
:start: start timestamp
|
||||
:end: end timestamp
|
||||
:annotation: application-level key-value datas
|
||||
"""
|
||||
self.request_id = request_id
|
||||
self.api = api
|
||||
self.seq = seq
|
||||
self.client = client
|
||||
self.server = server
|
||||
self.status = status
|
||||
self.start = start
|
||||
self.end = end
|
||||
self.annotation = annotation
|
||||
self.meta = meta
|
||||
|
||||
|
||||
class TTrackedClient(TClient):
|
||||
def __init__(self, tracker_handler, *args, **kwargs):
|
||||
super(TTrackedClient, self).__init__(*args, **kwargs)
|
||||
|
||||
self.tracker = tracker_handler
|
||||
self._upgraded = False
|
||||
|
||||
try:
|
||||
self._negotiation()
|
||||
self._upgraded = True
|
||||
except TApplicationException as e:
|
||||
if e.type != TApplicationException.UNKNOWN_METHOD:
|
||||
raise
|
||||
|
||||
def _negotiation(self):
|
||||
self._oprot.write_message_begin(track_method, TMessageType.CALL,
|
||||
self._seqid)
|
||||
args = track_thrift.UpgradeArgs()
|
||||
self.tracker.init_handshake_info(args)
|
||||
args.write(self._oprot)
|
||||
self._oprot.write_message_end()
|
||||
self._oprot.trans.flush()
|
||||
|
||||
api, msg_type, seqid = self._iprot.read_message_begin()
|
||||
if msg_type == TMessageType.EXCEPTION:
|
||||
x = TApplicationException()
|
||||
x.read(self._iprot)
|
||||
self._iprot.read_message_end()
|
||||
raise x
|
||||
else:
|
||||
result = track_thrift.UpgradeReply()
|
||||
result.read(self._iprot)
|
||||
self._iprot.read_message_end()
|
||||
|
||||
def _send(self, _api, **kwargs):
|
||||
if self._upgraded:
|
||||
self._header = track_thrift.RequestHeader()
|
||||
self.tracker.gen_header(self._header)
|
||||
self._header.write(self._oprot)
|
||||
|
||||
self.send_start = int(time.time() * 1000)
|
||||
super(TTrackedClient, self)._send(_api, **kwargs)
|
||||
|
||||
def _req(self, _api, *args, **kwargs):
|
||||
if not self._upgraded:
|
||||
return super(TTrackedClient, self)._req(_api, *args, **kwargs)
|
||||
|
||||
exception = None
|
||||
status = False
|
||||
|
||||
try:
|
||||
res = super(TTrackedClient, self)._req(_api, *args, **kwargs)
|
||||
status = True
|
||||
return res
|
||||
except BaseException as e:
|
||||
exception = e
|
||||
raise
|
||||
finally:
|
||||
header_info = RequestInfo(
|
||||
request_id=self._header.request_id,
|
||||
seq=self._header.seq,
|
||||
client=self.tracker.client,
|
||||
server=self.tracker.server,
|
||||
api=_api,
|
||||
status=status,
|
||||
start=self.send_start,
|
||||
end=int(time.time() * 1000),
|
||||
annotation=self.tracker.annotation,
|
||||
meta=self._header.meta,
|
||||
)
|
||||
self.tracker.record(header_info, exception)
|
||||
|
||||
|
||||
class TTrackedProcessor(TProcessor):
|
||||
def __init__(self, tracker_handler, *args, **kwargs):
|
||||
super(TTrackedProcessor, self).__init__(*args, **kwargs)
|
||||
|
||||
self.tracker = tracker_handler
|
||||
self._upgraded = False
|
||||
|
||||
def process(self, iprot, oprot):
|
||||
if not self._upgraded:
|
||||
res = self._try_upgrade(iprot)
|
||||
else:
|
||||
request_header = track_thrift.RequestHeader()
|
||||
request_header.read(iprot)
|
||||
self.tracker.handle(request_header)
|
||||
res = super(TTrackedProcessor, self).process_in(iprot)
|
||||
|
||||
self._do_process(iprot, oprot, *res)
|
||||
|
||||
def _try_upgrade(self, iprot):
|
||||
api, msg_type, seqid = iprot.read_message_begin()
|
||||
if msg_type == TMessageType.CALL and api == track_method:
|
||||
self._upgraded = True
|
||||
|
||||
args = track_thrift.UpgradeArgs()
|
||||
args.read(iprot)
|
||||
self.tracker.handle_handshake_info(args)
|
||||
result = track_thrift.UpgradeReply()
|
||||
result.oneway = False
|
||||
|
||||
def call():
|
||||
pass
|
||||
iprot.read_message_end()
|
||||
else:
|
||||
result, call = self._process_in(api, iprot)
|
||||
|
||||
return api, seqid, result, call
|
||||
|
||||
def _process_in(self, api, iprot):
|
||||
if api not in self._service.thrift_services:
|
||||
iprot.skip(TType.STRUCT)
|
||||
iprot.read_message_end()
|
||||
return TApplicationException(
|
||||
TApplicationException.UNKNOWN_METHOD), None
|
||||
|
||||
args = getattr(self._service, api + "_args")()
|
||||
args.read(iprot)
|
||||
iprot.read_message_end()
|
||||
result = getattr(self._service, api + "_result")()
|
||||
|
||||
# convert kwargs to args
|
||||
api_args = [args.thrift_spec[k][1]
|
||||
for k in sorted(args.thrift_spec)]
|
||||
|
||||
def call():
|
||||
return getattr(self._handler, api)(
|
||||
*(args.__dict__[k] for k in api_args)
|
||||
)
|
||||
|
||||
return result, call
|
||||
|
||||
def _do_process(self, iprot, oprot, api, seqid, result, call):
|
||||
if isinstance(result, TApplicationException):
|
||||
return self.send_exception(oprot, api, result, seqid)
|
||||
|
||||
try:
|
||||
result.success = call()
|
||||
except Exception as e:
|
||||
# raise if api don't have throws
|
||||
self.handle_exception(e, result)
|
||||
|
||||
if not result.oneway:
|
||||
self.send_result(oprot, api, result, seqid)
|
||||
|
||||
|
||||
from .tracker import TrackerBase, ConsoleTracker # noqa
|
||||
@@ -1,113 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import contextlib
|
||||
import copy
|
||||
import threading
|
||||
import uuid
|
||||
|
||||
ctx = threading.local()
|
||||
|
||||
|
||||
class TrackerBase(object):
|
||||
def __init__(self, client=None, server=None):
|
||||
self.client = client
|
||||
self.server = server
|
||||
|
||||
def handle(self, header):
|
||||
ctx.header = header
|
||||
ctx.counter = 0
|
||||
|
||||
def gen_header(self, header):
|
||||
header.request_id = self.get_request_id()
|
||||
|
||||
if not hasattr(ctx, "counter"):
|
||||
ctx.counter = 0
|
||||
|
||||
ctx.counter += 1
|
||||
|
||||
if hasattr(ctx, "header"):
|
||||
header.seq = "{prev_seq}.{cur_counter}".format(
|
||||
prev_seq=ctx.header.seq, cur_counter=ctx.counter)
|
||||
header.meta = ctx.header.meta
|
||||
else:
|
||||
header.meta = {}
|
||||
header.seq = str(ctx.counter)
|
||||
|
||||
if hasattr(ctx, "meta"):
|
||||
header.meta.update(ctx.meta)
|
||||
|
||||
def record(self, header, exception):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
@contextlib.contextmanager
|
||||
def counter(cls, init=0):
|
||||
"""Context for manually setting counter of seq number.
|
||||
|
||||
:init: init value
|
||||
"""
|
||||
if not hasattr(ctx, "counter"):
|
||||
ctx.counter = 0
|
||||
|
||||
old = ctx.counter
|
||||
ctx.counter = init
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
ctx.counter = old
|
||||
|
||||
@classmethod
|
||||
@contextlib.contextmanager
|
||||
def annotate(cls, **kwargs):
|
||||
ctx.annotation = kwargs
|
||||
try:
|
||||
yield ctx.annotation
|
||||
finally:
|
||||
del ctx.annotation
|
||||
|
||||
@classmethod
|
||||
@contextlib.contextmanager
|
||||
def add_meta(cls, **kwds):
|
||||
if hasattr(ctx, 'meta'):
|
||||
old_dict = copy.copy(ctx.meta)
|
||||
ctx.meta.update(kwds)
|
||||
try:
|
||||
yield ctx.meta
|
||||
finally:
|
||||
ctx.meta = old_dict
|
||||
else:
|
||||
ctx.meta = kwds
|
||||
try:
|
||||
yield ctx.meta
|
||||
finally:
|
||||
del ctx.meta
|
||||
|
||||
@property
|
||||
def meta(self):
|
||||
meta = ctx.header.meta if hasattr(ctx, "header") else {}
|
||||
if hasattr(ctx, "meta"):
|
||||
meta.update(ctx.meta)
|
||||
return meta
|
||||
|
||||
@property
|
||||
def annotation(self):
|
||||
return ctx.annotation if hasattr(ctx, "annotation") else {}
|
||||
|
||||
def get_request_id(self):
|
||||
if hasattr(ctx, "header"):
|
||||
return ctx.header.request_id
|
||||
return str(uuid.uuid4())
|
||||
|
||||
def init_handshake_info(self, handshake_obj):
|
||||
pass
|
||||
|
||||
def handle_handshake_info(self, handshake_obj):
|
||||
pass
|
||||
|
||||
|
||||
class ConsoleTracker(TrackerBase):
|
||||
def record(self, header, exception):
|
||||
print(header)
|
||||
@@ -1,16 +0,0 @@
|
||||
/*
|
||||
* This is the structure used to send call info to server.
|
||||
*/
|
||||
struct RequestHeader {
|
||||
1: string request_id
|
||||
2: string seq
|
||||
3: map<string, string> meta
|
||||
}
|
||||
|
||||
/**
|
||||
* This is the struct that a successful upgrade will reply with.
|
||||
*/
|
||||
struct UpgradeReply {}
|
||||
struct UpgradeArgs {
|
||||
1: string app_id
|
||||
}
|
||||
@@ -1,35 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import sys
|
||||
|
||||
from .parser import load_module
|
||||
|
||||
|
||||
class ThriftImporter(object):
|
||||
def __init__(self, extension="_thrift"):
|
||||
self.extension = extension
|
||||
|
||||
def __eq__(self, other):
|
||||
return self.__class__.__module__ == other.__class__.__module__ and \
|
||||
self.__class__.__name__ == other.__class__.__name__ and \
|
||||
self.extension == other.extension
|
||||
|
||||
def find_module(self, fullname, path=None):
|
||||
if fullname.endswith(self.extension):
|
||||
return self
|
||||
|
||||
def load_module(self, fullname):
|
||||
return load_module(fullname)
|
||||
_imp = ThriftImporter()
|
||||
|
||||
|
||||
def install_import_hook():
|
||||
global _imp
|
||||
sys.meta_path[:] = [x for x in sys.meta_path if _imp != x] + [_imp]
|
||||
|
||||
|
||||
def remove_import_hook():
|
||||
global _imp
|
||||
sys.meta_path[:] = [x for x in sys.meta_path if _imp != x]
|
||||
@@ -1,77 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
thriftpy.parser
|
||||
~~~~~~~~~~~~~~~
|
||||
|
||||
Thrift parser using ply
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
from .parser import parse, parse_fp
|
||||
|
||||
|
||||
def load(path, module_name=None, include_dirs=None, include_dir=None):
|
||||
"""Load thrift file as a module.
|
||||
|
||||
The module loaded and objects inside may only be pickled if module_name
|
||||
was provided.
|
||||
|
||||
Note: `include_dir` will be depreacated in the future, use `include_dirs`
|
||||
instead. If `include_dir` was provided (not None), it will be appended to
|
||||
`include_dirs`.
|
||||
"""
|
||||
real_module = bool(module_name)
|
||||
thrift = parse(path, module_name, include_dirs=include_dirs,
|
||||
include_dir=include_dir)
|
||||
|
||||
if real_module:
|
||||
sys.modules[module_name] = thrift
|
||||
return thrift
|
||||
|
||||
|
||||
def load_fp(source, module_name):
|
||||
"""Load thrift file like object as a module.
|
||||
"""
|
||||
thrift = parse_fp(source, module_name)
|
||||
sys.modules[module_name] = thrift
|
||||
return thrift
|
||||
|
||||
|
||||
def _import_module(import_name):
|
||||
if '.' in import_name:
|
||||
module, obj = import_name.rsplit('.', 1)
|
||||
return getattr(__import__(module, None, None, [obj]), obj)
|
||||
else:
|
||||
return __import__(import_name)
|
||||
|
||||
|
||||
def load_module(fullname):
|
||||
"""Load thrift_file by fullname, fullname should have '_thrift' as
|
||||
suffix.
|
||||
The loader will replace the '_thrift' with '.thrift' and use it as
|
||||
filename to locate the real thrift file.
|
||||
"""
|
||||
if not fullname.endswith("_thrift"):
|
||||
raise ImportError(
|
||||
"ThriftPy can only load module with '_thrift' suffix")
|
||||
|
||||
if fullname in sys.modules:
|
||||
return sys.modules[fullname]
|
||||
|
||||
if '.' in fullname:
|
||||
module_name, thrift_module_name = fullname.rsplit('.', 1)
|
||||
module = _import_module(module_name)
|
||||
path_prefix = os.path.dirname(os.path.abspath(module.__file__))
|
||||
path = os.path.join(path_prefix, thrift_module_name)
|
||||
else:
|
||||
path = fullname
|
||||
thrift_file = "{0}.thrift".format(path[:-7])
|
||||
|
||||
module = load(thrift_file, module_name=fullname)
|
||||
sys.modules[fullname] = module
|
||||
return sys.modules[fullname]
|
||||
@@ -1,15 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
|
||||
class ThriftParserError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class ThriftLexerError(ThriftParserError):
|
||||
pass
|
||||
|
||||
|
||||
class ThriftGrammerError(ThriftParserError):
|
||||
pass
|
||||
@@ -1,257 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
from .exc import ThriftLexerError
|
||||
|
||||
literals = ':;,=*{}()<>[]'
|
||||
|
||||
|
||||
thrift_reserved_keywords = (
|
||||
'BEGIN',
|
||||
'END',
|
||||
'__CLASS__',
|
||||
'__DIR__',
|
||||
'__FILE__',
|
||||
'__FUNCTION__',
|
||||
'__LINE__',
|
||||
'__METHOD__',
|
||||
'__NAMESPACE__',
|
||||
'abstract',
|
||||
'alias',
|
||||
'and',
|
||||
'args',
|
||||
'as',
|
||||
'assert',
|
||||
'begin',
|
||||
'break',
|
||||
'case',
|
||||
'catch',
|
||||
'class',
|
||||
'clone',
|
||||
'continue',
|
||||
'declare',
|
||||
'def',
|
||||
'default',
|
||||
'del',
|
||||
'delete',
|
||||
'do',
|
||||
'dynamic',
|
||||
'elif',
|
||||
'else',
|
||||
'elseif',
|
||||
'elsif',
|
||||
'end',
|
||||
'enddeclare',
|
||||
'endfor',
|
||||
'endforeach',
|
||||
'endif',
|
||||
'endswitch',
|
||||
'endwhile',
|
||||
'ensure',
|
||||
'except',
|
||||
'exec',
|
||||
'finally',
|
||||
'float',
|
||||
'for',
|
||||
'foreach',
|
||||
'function',
|
||||
'global',
|
||||
'goto',
|
||||
'if',
|
||||
'implements',
|
||||
'import',
|
||||
'in',
|
||||
'inline',
|
||||
'instanceof',
|
||||
'interface',
|
||||
'is',
|
||||
'lambda',
|
||||
'module',
|
||||
'native',
|
||||
'new',
|
||||
'next',
|
||||
'nil',
|
||||
'not',
|
||||
'or',
|
||||
'pass',
|
||||
'public',
|
||||
'print',
|
||||
'private',
|
||||
'protected',
|
||||
'public',
|
||||
'raise',
|
||||
'redo',
|
||||
'rescue',
|
||||
'retry',
|
||||
'register',
|
||||
'return',
|
||||
'self',
|
||||
'sizeof',
|
||||
'static',
|
||||
'super',
|
||||
'switch',
|
||||
'synchronized',
|
||||
'then',
|
||||
'this',
|
||||
'throw',
|
||||
'transient',
|
||||
'try',
|
||||
'undef',
|
||||
'union',
|
||||
'unless',
|
||||
'unsigned',
|
||||
'until',
|
||||
'use',
|
||||
'var',
|
||||
'virtual',
|
||||
'volatile',
|
||||
'when',
|
||||
'while',
|
||||
'with',
|
||||
'xor',
|
||||
'yield'
|
||||
)
|
||||
|
||||
|
||||
keywords = (
|
||||
'namespace',
|
||||
'include',
|
||||
'void',
|
||||
'bool',
|
||||
'byte',
|
||||
'i16',
|
||||
'i32',
|
||||
'i64',
|
||||
'double',
|
||||
'string',
|
||||
'binary',
|
||||
'map',
|
||||
'list',
|
||||
'set',
|
||||
'oneway',
|
||||
'typedef',
|
||||
'struct',
|
||||
'union',
|
||||
'exception',
|
||||
'extends',
|
||||
'throws',
|
||||
'service',
|
||||
'enum',
|
||||
'const',
|
||||
'required',
|
||||
'optional',
|
||||
)
|
||||
|
||||
|
||||
tokens = (
|
||||
'BOOLCONSTANT',
|
||||
'INTCONSTANT',
|
||||
'DUBCONSTANT',
|
||||
'LITERAL',
|
||||
'IDENTIFIER',
|
||||
) + tuple(map(lambda kw: kw.upper(), keywords))
|
||||
|
||||
|
||||
t_ignore = ' \t\r' # whitespace
|
||||
|
||||
|
||||
def t_error(t):
|
||||
raise ThriftLexerError('Illegal characher %r at line %d' %
|
||||
(t.value[0], t.lineno))
|
||||
|
||||
|
||||
def t_newline(t):
|
||||
r'\n+'
|
||||
t.lexer.lineno += len(t.value)
|
||||
|
||||
|
||||
def t_ignore_SILLYCOMM(t):
|
||||
r'\/\*\**\*\/'
|
||||
t.lexer.lineno += t.value.count('\n')
|
||||
|
||||
|
||||
def t_ignore_MULTICOMM(t):
|
||||
r'\/\*[^*]\/*([^*/]|[^*]\/|\*[^/])*\**\*\/'
|
||||
t.lexer.lineno += t.value.count('\n')
|
||||
|
||||
|
||||
def t_ignore_DOCTEXT(t):
|
||||
r'\/\*\*([^*/]|[^*]\/|\*[^/])*\**\*\/'
|
||||
t.lexer.lineno += t.value.count('\n')
|
||||
|
||||
|
||||
def t_ignore_UNIXCOMMENT(t):
|
||||
r'\#[^\n]*'
|
||||
|
||||
|
||||
def t_ignore_COMMENT(t):
|
||||
r'\/\/[^\n]*'
|
||||
|
||||
|
||||
def t_BOOLCONSTANT(t):
|
||||
r'true|false'
|
||||
t.value = t.value == 'true'
|
||||
return t
|
||||
|
||||
|
||||
def t_DUBCONSTANT(t):
|
||||
r'-?\d+\.\d*(e-?\d+)?'
|
||||
t.value = float(t.value)
|
||||
return t
|
||||
|
||||
|
||||
def t_HEXCONSTANT(t):
|
||||
r'0x[0-9A-Fa-f]+'
|
||||
t.value = int(t.value, 16)
|
||||
t.type = 'INTCONSTANT'
|
||||
return t
|
||||
|
||||
|
||||
def t_INTCONSTANT(t):
|
||||
r'[+-]?[0-9]+'
|
||||
t.value = int(t.value)
|
||||
return t
|
||||
|
||||
|
||||
def t_LITERAL(t):
|
||||
r'(\"([^\\\n]|(\\.))*?\")|\'([^\\\n]|(\\.))*?\''
|
||||
s = t.value[1:-1]
|
||||
maps = {
|
||||
't': '\t',
|
||||
'r': '\r',
|
||||
'n': '\n',
|
||||
'\\': '\\',
|
||||
'\'': '\'',
|
||||
'"': '\"'
|
||||
}
|
||||
i = 0
|
||||
length = len(s)
|
||||
val = ''
|
||||
while i < length:
|
||||
if s[i] == '\\':
|
||||
i += 1
|
||||
if s[i] in maps:
|
||||
val += maps[s[i]]
|
||||
else:
|
||||
msg = 'Unexcepted escaping characher: %s' % s[i]
|
||||
raise ThriftLexerError(msg)
|
||||
else:
|
||||
val += s[i]
|
||||
|
||||
i += 1
|
||||
|
||||
t.value = val
|
||||
return t
|
||||
|
||||
|
||||
def t_IDENTIFIER(t):
|
||||
r'[a-zA-Z_](\.[a-zA-Z_0-9]|[a-zA-Z_0-9])*'
|
||||
|
||||
if t.value in keywords:
|
||||
t.type = t.value.upper()
|
||||
return t
|
||||
if t.value in thrift_reserved_keywords:
|
||||
raise ThriftLexerError('Cannot use reserved language keyword: %r'
|
||||
' at line %d' % (t.value, t.lineno))
|
||||
return t
|
||||
@@ -1,820 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
IDL Ref:
|
||||
https://thrift.apache.org/docs/idl
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import collections
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
from ply import lex, yacc
|
||||
from .lexer import * # noqa
|
||||
from .exc import ThriftParserError, ThriftGrammerError
|
||||
from ..thrift import gen_init, TType, TPayload, TException
|
||||
|
||||
|
||||
def p_error(p):
|
||||
if p is None:
|
||||
raise ThriftGrammerError('Grammer error at EOF')
|
||||
raise ThriftGrammerError('Grammer error %r at line %d' %
|
||||
(p.value, p.lineno))
|
||||
|
||||
|
||||
def p_start(p):
|
||||
'''start : header definition'''
|
||||
|
||||
|
||||
def p_header(p):
|
||||
'''header : header_unit_ header
|
||||
|'''
|
||||
|
||||
|
||||
def p_header_unit_(p):
|
||||
'''header_unit_ : header_unit ';'
|
||||
| header_unit'''
|
||||
|
||||
|
||||
def p_header_unit(p):
|
||||
'''header_unit : include
|
||||
| namespace'''
|
||||
|
||||
|
||||
def p_include(p):
|
||||
'''include : INCLUDE LITERAL'''
|
||||
thrift = thrift_stack[-1]
|
||||
|
||||
if thrift.__thrift_file__ is None:
|
||||
raise ThriftParserError('Unexcepted include statement while loading'
|
||||
'from file like object.')
|
||||
|
||||
for include_dir in include_dirs_:
|
||||
path = os.path.join(include_dir, p[2])
|
||||
if os.path.exists(path):
|
||||
child = parse(path)
|
||||
setattr(thrift, child.__name__, child)
|
||||
_add_thrift_meta('includes', child)
|
||||
return
|
||||
raise ThriftParserError(('Couldn\'t include thrift %s in any '
|
||||
'directories provided') % p[2])
|
||||
|
||||
|
||||
def p_namespace(p):
|
||||
'''namespace : NAMESPACE namespace_scope IDENTIFIER'''
|
||||
# namespace is useless in thriftpy
|
||||
# if p[2] == 'py' or p[2] == '*':
|
||||
# setattr(thrift_stack[-1], '__name__', p[3])
|
||||
|
||||
|
||||
def p_namespace_scope(p):
|
||||
'''namespace_scope : '*'
|
||||
| IDENTIFIER'''
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
def p_sep(p):
|
||||
'''sep : ','
|
||||
| ';'
|
||||
'''
|
||||
|
||||
|
||||
def p_definition(p):
|
||||
'''definition : definition definition_unit_
|
||||
|'''
|
||||
|
||||
|
||||
def p_definition_unit_(p):
|
||||
'''definition_unit_ : definition_unit ';'
|
||||
| definition_unit'''
|
||||
|
||||
|
||||
def p_definition_unit(p):
|
||||
'''definition_unit : const
|
||||
| ttype
|
||||
'''
|
||||
|
||||
|
||||
def p_const(p):
|
||||
'''const : CONST field_type IDENTIFIER '=' const_value
|
||||
| CONST field_type IDENTIFIER '=' const_value sep'''
|
||||
|
||||
try:
|
||||
val = _cast(p[2])(p[5])
|
||||
except AssertionError:
|
||||
raise ThriftParserError('Type error for constant %s at line %d' %
|
||||
(p[3], p.lineno(3)))
|
||||
setattr(thrift_stack[-1], p[3], val)
|
||||
_add_thrift_meta('consts', val)
|
||||
|
||||
|
||||
def p_const_value(p):
|
||||
'''const_value : INTCONSTANT
|
||||
| DUBCONSTANT
|
||||
| LITERAL
|
||||
| BOOLCONSTANT
|
||||
| const_list
|
||||
| const_map
|
||||
| const_ref'''
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
def p_const_list(p):
|
||||
'''const_list : '[' const_list_seq ']' '''
|
||||
p[0] = p[2]
|
||||
|
||||
|
||||
def p_const_list_seq(p):
|
||||
'''const_list_seq : const_value sep const_list_seq
|
||||
| const_value const_list_seq
|
||||
|'''
|
||||
_parse_seq(p)
|
||||
|
||||
|
||||
def p_const_map(p):
|
||||
'''const_map : '{' const_map_seq '}' '''
|
||||
p[0] = dict(p[2])
|
||||
|
||||
|
||||
def p_const_map_seq(p):
|
||||
'''const_map_seq : const_map_item sep const_map_seq
|
||||
| const_map_item const_map_seq
|
||||
|'''
|
||||
_parse_seq(p)
|
||||
|
||||
|
||||
def p_const_map_item(p):
|
||||
'''const_map_item : const_value ':' const_value '''
|
||||
p[0] = [p[1], p[3]]
|
||||
|
||||
|
||||
def p_const_ref(p):
|
||||
'''const_ref : IDENTIFIER'''
|
||||
child = thrift_stack[-1]
|
||||
|
||||
for name in p[1].split('.'):
|
||||
father = child
|
||||
child = getattr(child, name, None)
|
||||
if child is None:
|
||||
raise ThriftParserError('Cann\'t find name %r at line %d'
|
||||
% (p[1], p.lineno(1)))
|
||||
|
||||
if _get_ttype(child) is None or _get_ttype(father) == TType.I32:
|
||||
# child is a constant or enum value
|
||||
p[0] = child
|
||||
else:
|
||||
raise ThriftParserError('No enum value or constant found '
|
||||
'named %r' % p[1])
|
||||
|
||||
|
||||
def p_ttype(p):
|
||||
'''ttype : typedef
|
||||
| enum
|
||||
| struct
|
||||
| union
|
||||
| exception
|
||||
| service'''
|
||||
|
||||
|
||||
def p_typedef(p):
|
||||
'''typedef : TYPEDEF field_type IDENTIFIER'''
|
||||
setattr(thrift_stack[-1], p[3], p[2])
|
||||
|
||||
|
||||
def p_enum(p): # noqa
|
||||
'''enum : ENUM IDENTIFIER '{' enum_seq '}' '''
|
||||
val = _make_enum(p[2], p[4])
|
||||
setattr(thrift_stack[-1], p[2], val)
|
||||
_add_thrift_meta('enums', val)
|
||||
|
||||
|
||||
def p_enum_seq(p):
|
||||
'''enum_seq : enum_item sep enum_seq
|
||||
| enum_item enum_seq
|
||||
|'''
|
||||
_parse_seq(p)
|
||||
|
||||
|
||||
def p_enum_item(p):
|
||||
'''enum_item : IDENTIFIER '=' INTCONSTANT
|
||||
| IDENTIFIER
|
||||
|'''
|
||||
if len(p) == 4:
|
||||
p[0] = [p[1], p[3]]
|
||||
elif len(p) == 2:
|
||||
p[0] = [p[1], None]
|
||||
|
||||
|
||||
def p_struct(p):
|
||||
'''struct : seen_struct '{' field_seq '}' '''
|
||||
val = _fill_in_struct(p[1], p[3])
|
||||
_add_thrift_meta('structs', val)
|
||||
|
||||
|
||||
def p_seen_struct(p):
|
||||
'''seen_struct : STRUCT IDENTIFIER '''
|
||||
val = _make_empty_struct(p[2])
|
||||
setattr(thrift_stack[-1], p[2], val)
|
||||
p[0] = val
|
||||
|
||||
|
||||
def p_union(p):
|
||||
'''union : seen_union '{' field_seq '}' '''
|
||||
val = _fill_in_struct(p[1], p[3])
|
||||
_add_thrift_meta('unions', val)
|
||||
|
||||
|
||||
def p_seen_union(p):
|
||||
'''seen_union : UNION IDENTIFIER '''
|
||||
val = _make_empty_struct(p[2])
|
||||
setattr(thrift_stack[-1], p[2], val)
|
||||
p[0] = val
|
||||
|
||||
|
||||
def p_exception(p):
|
||||
'''exception : EXCEPTION IDENTIFIER '{' field_seq '}' '''
|
||||
val = _make_struct(p[2], p[4], base_cls=TException)
|
||||
setattr(thrift_stack[-1], p[2], val)
|
||||
_add_thrift_meta('exceptions', val)
|
||||
|
||||
|
||||
def p_service(p):
|
||||
'''service : SERVICE IDENTIFIER '{' function_seq '}'
|
||||
| SERVICE IDENTIFIER EXTENDS IDENTIFIER '{' function_seq '}'
|
||||
'''
|
||||
thrift = thrift_stack[-1]
|
||||
|
||||
if len(p) == 8:
|
||||
extends = thrift
|
||||
for name in p[4].split('.'):
|
||||
extends = getattr(extends, name, None)
|
||||
if extends is None:
|
||||
raise ThriftParserError('Can\'t find service %r for '
|
||||
'service %r to extend' %
|
||||
(p[4], p[2]))
|
||||
|
||||
if not hasattr(extends, 'thrift_services'):
|
||||
raise ThriftParserError('Can\'t extends %r, not a service'
|
||||
% p[4])
|
||||
|
||||
else:
|
||||
extends = None
|
||||
|
||||
val = _make_service(p[2], p[len(p) - 2], extends)
|
||||
setattr(thrift, p[2], val)
|
||||
_add_thrift_meta('services', val)
|
||||
|
||||
|
||||
def p_function(p):
|
||||
'''function : ONEWAY function_type IDENTIFIER '(' field_seq ')' throws
|
||||
| ONEWAY function_type IDENTIFIER '(' field_seq ')'
|
||||
| function_type IDENTIFIER '(' field_seq ')' throws
|
||||
| function_type IDENTIFIER '(' field_seq ')' '''
|
||||
|
||||
if p[1] == 'oneway':
|
||||
oneway = True
|
||||
base = 1
|
||||
else:
|
||||
oneway = False
|
||||
base = 0
|
||||
|
||||
if p[len(p) - 1] == ')':
|
||||
throws = []
|
||||
else:
|
||||
throws = p[len(p) - 1]
|
||||
|
||||
p[0] = [oneway, p[base + 1], p[base + 2], p[base + 4], throws]
|
||||
|
||||
|
||||
def p_function_seq(p):
|
||||
'''function_seq : function sep function_seq
|
||||
| function function_seq
|
||||
|'''
|
||||
_parse_seq(p)
|
||||
|
||||
|
||||
def p_throws(p):
|
||||
'''throws : THROWS '(' field_seq ')' '''
|
||||
p[0] = p[3]
|
||||
|
||||
|
||||
def p_function_type(p):
|
||||
'''function_type : field_type
|
||||
| VOID'''
|
||||
if p[1] == 'void':
|
||||
p[0] = TType.VOID
|
||||
else:
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
def p_field_seq(p):
|
||||
'''field_seq : field sep field_seq
|
||||
| field field_seq
|
||||
|'''
|
||||
_parse_seq(p)
|
||||
|
||||
|
||||
def p_field(p):
|
||||
'''field : field_id field_req field_type IDENTIFIER
|
||||
| field_id field_req field_type IDENTIFIER '=' const_value'''
|
||||
|
||||
if len(p) == 7:
|
||||
try:
|
||||
val = _cast(p[3])(p[6])
|
||||
except AssertionError:
|
||||
raise ThriftParserError(
|
||||
'Type error for field %s '
|
||||
'at line %d' % (p[4], p.lineno(4)))
|
||||
else:
|
||||
val = None
|
||||
|
||||
p[0] = [p[1], p[2], p[3], p[4], val]
|
||||
|
||||
|
||||
def p_field_id(p):
|
||||
'''field_id : INTCONSTANT ':' '''
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
def p_field_req(p):
|
||||
'''field_req : REQUIRED
|
||||
| OPTIONAL
|
||||
|'''
|
||||
if len(p) == 2:
|
||||
p[0] = p[1] == 'required'
|
||||
elif len(p) == 1:
|
||||
p[0] = False # default: required=False
|
||||
|
||||
|
||||
def p_field_type(p):
|
||||
'''field_type : ref_type
|
||||
| definition_type'''
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
def p_ref_type(p):
|
||||
'''ref_type : IDENTIFIER'''
|
||||
ref_type = thrift_stack[-1]
|
||||
|
||||
for name in p[1].split('.'):
|
||||
ref_type = getattr(ref_type, name, None)
|
||||
if ref_type is None:
|
||||
raise ThriftParserError('No type found: %r, at line %d' %
|
||||
(p[1], p.lineno(1)))
|
||||
|
||||
if hasattr(ref_type, '_ttype'):
|
||||
p[0] = getattr(ref_type, '_ttype'), ref_type
|
||||
else:
|
||||
p[0] = ref_type
|
||||
|
||||
|
||||
def p_base_type(p): # noqa
|
||||
'''base_type : BOOL
|
||||
| BYTE
|
||||
| I16
|
||||
| I32
|
||||
| I64
|
||||
| DOUBLE
|
||||
| STRING
|
||||
| BINARY'''
|
||||
if p[1] == 'bool':
|
||||
p[0] = TType.BOOL
|
||||
if p[1] == 'byte':
|
||||
p[0] = TType.BYTE
|
||||
if p[1] == 'i16':
|
||||
p[0] = TType.I16
|
||||
if p[1] == 'i32':
|
||||
p[0] = TType.I32
|
||||
if p[1] == 'i64':
|
||||
p[0] = TType.I64
|
||||
if p[1] == 'double':
|
||||
p[0] = TType.DOUBLE
|
||||
if p[1] == 'string':
|
||||
p[0] = TType.STRING
|
||||
if p[1] == 'binary':
|
||||
p[0] = TType.BINARY
|
||||
|
||||
|
||||
def p_container_type(p):
|
||||
'''container_type : map_type
|
||||
| list_type
|
||||
| set_type'''
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
def p_map_type(p):
|
||||
'''map_type : MAP '<' field_type ',' field_type '>' '''
|
||||
p[0] = TType.MAP, (p[3], p[5])
|
||||
|
||||
|
||||
def p_list_type(p):
|
||||
'''list_type : LIST '<' field_type '>' '''
|
||||
p[0] = TType.LIST, p[3]
|
||||
|
||||
|
||||
def p_set_type(p):
|
||||
'''set_type : SET '<' field_type '>' '''
|
||||
p[0] = TType.SET, p[3]
|
||||
|
||||
|
||||
def p_definition_type(p):
|
||||
'''definition_type : base_type
|
||||
| container_type'''
|
||||
p[0] = p[1]
|
||||
|
||||
|
||||
thrift_stack = []
|
||||
include_dirs_ = ['.']
|
||||
thrift_cache = {}
|
||||
|
||||
|
||||
def parse(path, module_name=None, include_dirs=None, include_dir=None,
|
||||
lexer=None, parser=None, enable_cache=True):
|
||||
"""Parse a single thrift file to module object, e.g.::
|
||||
|
||||
>>> from thriftpy.parser.parser import parse
|
||||
>>> note_thrift = parse("path/to/note.thrift")
|
||||
<module 'note_thrift' (built-in)>
|
||||
|
||||
:param path: file path to parse, should be a string ending with '.thrift'.
|
||||
:param module_name: the name for parsed module, the default is the basename
|
||||
without extension of `path`.
|
||||
:param include_dirs: directories to find thrift files while processing
|
||||
the `include` directive, by default: ['.'].
|
||||
:param include_dir: directory to find child thrift files. Note this keyword
|
||||
parameter will be deprecated in the future, it exists
|
||||
for compatiable reason. If it's provided (not `None`),
|
||||
it will be appended to `include_dirs`.
|
||||
:param lexer: ply lexer to use, if not provided, `parse` will new one.
|
||||
:param parser: ply parser to use, if not provided, `parse` will new one.
|
||||
:param enable_cache: if this is set to be `True`, parsed module will be
|
||||
cached, this is enabled by default. If `module_name`
|
||||
is provided, use it as cache key, else use the `path`.
|
||||
"""
|
||||
if os.name == 'nt' and sys.version_info < (3, 2):
|
||||
os.path.samefile = lambda f1, f2: os.stat(f1) == os.stat(f2)
|
||||
|
||||
# dead include checking on current stack
|
||||
for thrift in thrift_stack:
|
||||
if thrift.__thrift_file__ is not None and \
|
||||
os.path.samefile(path, thrift.__thrift_file__):
|
||||
raise ThriftParserError('Dead including on %s' % path)
|
||||
|
||||
global thrift_cache
|
||||
|
||||
cache_key = module_name or os.path.normpath(path)
|
||||
|
||||
if enable_cache and cache_key in thrift_cache:
|
||||
return thrift_cache[cache_key]
|
||||
|
||||
if lexer is None:
|
||||
lexer = lex.lex()
|
||||
if parser is None:
|
||||
parser = yacc.yacc(debug=False, write_tables=0)
|
||||
|
||||
global include_dirs_
|
||||
|
||||
if include_dirs is not None:
|
||||
include_dirs_ = include_dirs
|
||||
if include_dir is not None:
|
||||
include_dirs_.append(include_dir)
|
||||
|
||||
if not path.endswith('.thrift'):
|
||||
raise ThriftParserError('Path should end with .thrift')
|
||||
|
||||
with open(path) as fh:
|
||||
data = fh.read()
|
||||
|
||||
if module_name is not None and not module_name.endswith('_thrift'):
|
||||
raise ThriftParserError('ThriftPy can only generate module with '
|
||||
'\'_thrift\' suffix')
|
||||
|
||||
if module_name is None:
|
||||
basename = os.path.basename(path)
|
||||
module_name = os.path.splitext(basename)[0]
|
||||
|
||||
thrift = types.ModuleType(module_name)
|
||||
setattr(thrift, '__thrift_file__', path)
|
||||
thrift_stack.append(thrift)
|
||||
lexer.lineno = 1
|
||||
parser.parse(data)
|
||||
thrift_stack.pop()
|
||||
|
||||
if enable_cache:
|
||||
thrift_cache[cache_key] = thrift
|
||||
return thrift
|
||||
|
||||
|
||||
def parse_fp(source, module_name, lexer=None, parser=None, enable_cache=True):
|
||||
"""Parse a file-like object to thrift module object, e.g.::
|
||||
|
||||
>>> from thriftpy.parser.parser import parse_fp
|
||||
>>> with open("path/to/note.thrift") as fp:
|
||||
parse_fp(fp, "note_thrift")
|
||||
<module 'note_thrift' (built-in)>
|
||||
|
||||
:param source: file-like object, expected to have a method named `read`.
|
||||
:param module_name: the name for parsed module, shoule be endswith
|
||||
'_thrift'.
|
||||
:param lexer: ply lexer to use, if not provided, `parse` will new one.
|
||||
:param parser: ply parser to use, if not provided, `parse` will new one.
|
||||
:param enable_cache: if this is set to be `True`, parsed module will be
|
||||
cached by `module_name`, this is enabled by default.
|
||||
"""
|
||||
if not module_name.endswith('_thrift'):
|
||||
raise ThriftParserError('ThriftPy can only generate module with '
|
||||
'\'_thrift\' suffix')
|
||||
|
||||
if enable_cache and module_name in thrift_cache:
|
||||
return thrift_cache[module_name]
|
||||
|
||||
if not hasattr(source, 'read'):
|
||||
raise ThriftParserError('Except `source` to be a file-like object with'
|
||||
'a method named \'read\'')
|
||||
|
||||
if lexer is None:
|
||||
lexer = lex.lex()
|
||||
if parser is None:
|
||||
parser = yacc.yacc(debug=False, write_tables=0)
|
||||
|
||||
data = source.read()
|
||||
|
||||
thrift = types.ModuleType(module_name)
|
||||
setattr(thrift, '__thrift_file__', None)
|
||||
thrift_stack.append(thrift)
|
||||
lexer.lineno = 1
|
||||
parser.parse(data)
|
||||
thrift_stack.pop()
|
||||
|
||||
if enable_cache:
|
||||
thrift_cache[module_name] = thrift
|
||||
return thrift
|
||||
|
||||
|
||||
def _add_thrift_meta(key, val):
|
||||
thrift = thrift_stack[-1]
|
||||
|
||||
if not hasattr(thrift, '__thrift_meta__'):
|
||||
meta = collections.defaultdict(list)
|
||||
setattr(thrift, '__thrift_meta__', meta)
|
||||
else:
|
||||
meta = getattr(thrift, '__thrift_meta__')
|
||||
|
||||
meta[key].append(val)
|
||||
|
||||
|
||||
def _parse_seq(p):
|
||||
if len(p) == 4:
|
||||
p[0] = [p[1]] + p[3]
|
||||
elif len(p) == 3:
|
||||
p[0] = [p[1]] + p[2]
|
||||
elif len(p) == 1:
|
||||
p[0] = []
|
||||
|
||||
|
||||
def _cast(t): # noqa
|
||||
if t == TType.BOOL:
|
||||
return _cast_bool
|
||||
if t == TType.BYTE:
|
||||
return _cast_byte
|
||||
if t == TType.I16:
|
||||
return _cast_i16
|
||||
if t == TType.I32:
|
||||
return _cast_i32
|
||||
if t == TType.I64:
|
||||
return _cast_i64
|
||||
if t == TType.DOUBLE:
|
||||
return _cast_double
|
||||
if t == TType.STRING:
|
||||
return _cast_string
|
||||
if t == TType.BINARY:
|
||||
return _cast_binary
|
||||
if t[0] == TType.LIST:
|
||||
return _cast_list(t)
|
||||
if t[0] == TType.SET:
|
||||
return _cast_set(t)
|
||||
if t[0] == TType.MAP:
|
||||
return _cast_map(t)
|
||||
if t[0] == TType.I32:
|
||||
return _cast_enum(t)
|
||||
if t[0] == TType.STRUCT:
|
||||
return _cast_struct(t)
|
||||
|
||||
|
||||
def _cast_bool(v):
|
||||
assert isinstance(v, (bool, int))
|
||||
return bool(v)
|
||||
|
||||
|
||||
def _cast_byte(v):
|
||||
assert isinstance(v, str)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_i16(v):
|
||||
assert isinstance(v, int)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_i32(v):
|
||||
assert isinstance(v, int)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_i64(v):
|
||||
assert isinstance(v, int)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_double(v):
|
||||
assert isinstance(v, float)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_string(v):
|
||||
assert isinstance(v, str)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_binary(v):
|
||||
assert isinstance(v, str)
|
||||
return v
|
||||
|
||||
|
||||
def _cast_list(t):
|
||||
assert t[0] == TType.LIST
|
||||
|
||||
def __cast_list(v):
|
||||
assert isinstance(v, list)
|
||||
map(_cast(t[1]), v)
|
||||
return v
|
||||
return __cast_list
|
||||
|
||||
|
||||
def _cast_set(t):
|
||||
assert t[0] == TType.SET
|
||||
|
||||
def __cast_set(v):
|
||||
assert isinstance(v, (list, set))
|
||||
map(_cast(t[1]), v)
|
||||
if not isinstance(v, set):
|
||||
return set(v)
|
||||
return v
|
||||
return __cast_set
|
||||
|
||||
|
||||
def _cast_map(t):
|
||||
assert t[0] == TType.MAP
|
||||
|
||||
def __cast_map(v):
|
||||
assert isinstance(v, dict)
|
||||
for key in v:
|
||||
v[_cast(t[1][0])(key)] = \
|
||||
_cast(t[1][1])(v[key])
|
||||
return v
|
||||
return __cast_map
|
||||
|
||||
|
||||
def _cast_enum(t):
|
||||
assert t[0] == TType.I32
|
||||
|
||||
def __cast_enum(v):
|
||||
assert isinstance(v, int)
|
||||
if v in t[1]._VALUES_TO_NAMES:
|
||||
return v
|
||||
raise ThriftParserError('Couldn\'t find a named value in enum '
|
||||
'%s for value %d' % (t[1].__name__, v))
|
||||
return __cast_enum
|
||||
|
||||
|
||||
def _cast_struct(t): # struct/exception/union
|
||||
assert t[0] == TType.STRUCT
|
||||
|
||||
def __cast_struct(v):
|
||||
if isinstance(v, t[1]):
|
||||
return v # already cast
|
||||
|
||||
assert isinstance(v, dict)
|
||||
tspec = getattr(t[1], '_tspec')
|
||||
|
||||
for key in tspec: # requirement check
|
||||
if tspec[key][0] and key not in v:
|
||||
raise ThriftParserError('Field %r was required to create '
|
||||
'constant for type %r' %
|
||||
(key, t[1].__name__))
|
||||
|
||||
for key in v: # cast values
|
||||
if key not in tspec:
|
||||
raise ThriftParserError('No field named %r was '
|
||||
'found in struct of type %r' %
|
||||
(key, t[1].__name__))
|
||||
v[key] = _cast(tspec[key][1])(v[key])
|
||||
return t[1](**v)
|
||||
return __cast_struct
|
||||
|
||||
|
||||
def _make_enum(name, kvs):
|
||||
attrs = {'__module__': thrift_stack[-1].__name__, '_ttype': TType.I32}
|
||||
cls = type(name, (object, ), attrs)
|
||||
|
||||
_values_to_names = {}
|
||||
_names_to_values = {}
|
||||
|
||||
if kvs:
|
||||
val = kvs[0][1]
|
||||
if val is None:
|
||||
val = -1
|
||||
for item in kvs:
|
||||
if item[1] is None:
|
||||
item[1] = val + 1
|
||||
val = item[1]
|
||||
for key, val in kvs:
|
||||
setattr(cls, key, val)
|
||||
_values_to_names[val] = key
|
||||
_names_to_values[key] = val
|
||||
setattr(cls, '_VALUES_TO_NAMES', _values_to_names)
|
||||
setattr(cls, '_NAMES_TO_VALUES', _names_to_values)
|
||||
return cls
|
||||
|
||||
|
||||
def _make_empty_struct(name, ttype=TType.STRUCT, base_cls=TPayload):
|
||||
attrs = {'__module__': thrift_stack[-1].__name__, '_ttype': ttype}
|
||||
return type(name, (base_cls, ), attrs)
|
||||
|
||||
|
||||
def _fill_in_struct(cls, fields, _gen_init=True):
|
||||
thrift_spec = {}
|
||||
default_spec = []
|
||||
_tspec = {}
|
||||
|
||||
for field in fields:
|
||||
if field[0] in thrift_spec or field[3] in _tspec:
|
||||
raise ThriftGrammerError(('\'%d:%s\' field identifier/name has '
|
||||
'already been used') % (
|
||||
field[0], field[3]))
|
||||
ttype = field[2]
|
||||
thrift_spec[field[0]] = _ttype_spec(ttype, field[3], field[1])
|
||||
default_spec.append((field[3], field[4]))
|
||||
_tspec[field[3]] = field[1], ttype
|
||||
setattr(cls, 'thrift_spec', thrift_spec)
|
||||
setattr(cls, 'default_spec', default_spec)
|
||||
setattr(cls, '_tspec', _tspec)
|
||||
if _gen_init:
|
||||
gen_init(cls, thrift_spec, default_spec)
|
||||
return cls
|
||||
|
||||
|
||||
def _make_struct(name, fields, ttype=TType.STRUCT, base_cls=TPayload,
|
||||
_gen_init=True):
|
||||
cls = _make_empty_struct(name, ttype=ttype, base_cls=base_cls)
|
||||
return _fill_in_struct(cls, fields, _gen_init=_gen_init)
|
||||
|
||||
|
||||
def _make_service(name, funcs, extends):
|
||||
if extends is None:
|
||||
extends = object
|
||||
|
||||
attrs = {'__module__': thrift_stack[-1].__name__}
|
||||
cls = type(name, (extends, ), attrs)
|
||||
thrift_services = []
|
||||
|
||||
for func in funcs:
|
||||
func_name = func[2]
|
||||
# args payload cls
|
||||
args_name = '%s_args' % func_name
|
||||
args_fields = func[3]
|
||||
args_cls = _make_struct(args_name, args_fields)
|
||||
setattr(cls, args_name, args_cls)
|
||||
# result payload cls
|
||||
result_name = '%s_result' % func_name
|
||||
result_type = func[1]
|
||||
result_throws = func[4]
|
||||
result_oneway = func[0]
|
||||
result_cls = _make_struct(result_name, result_throws,
|
||||
_gen_init=False)
|
||||
setattr(result_cls, 'oneway', result_oneway)
|
||||
if result_type != TType.VOID:
|
||||
result_cls.thrift_spec[0] = _ttype_spec(result_type, 'success')
|
||||
result_cls.default_spec.insert(0, ('success', None))
|
||||
gen_init(result_cls, result_cls.thrift_spec, result_cls.default_spec)
|
||||
setattr(cls, result_name, result_cls)
|
||||
thrift_services.append(func_name)
|
||||
if extends is not None and hasattr(extends, 'thrift_services'):
|
||||
thrift_services.extend(extends.thrift_services)
|
||||
setattr(cls, 'thrift_services', thrift_services)
|
||||
return cls
|
||||
|
||||
|
||||
def _ttype_spec(ttype, name, required=False):
|
||||
if isinstance(ttype, int):
|
||||
return ttype, name, required
|
||||
else:
|
||||
return ttype[0], name, ttype[1], required
|
||||
|
||||
|
||||
def _get_ttype(inst, default_ttype=None):
|
||||
if hasattr(inst, '__dict__') and '_ttype' in inst.__dict__:
|
||||
return inst.__dict__['_ttype']
|
||||
return default_ttype
|
||||
@@ -1,26 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
from thriftpy._compat import PYPY, CYTHON
|
||||
from .binary import TBinaryProtocol, TBinaryProtocolFactory
|
||||
from .compact import TCompactProtocol, TCompactProtocolFactory
|
||||
from .json import TJSONProtocol, TJSONProtocolFactory
|
||||
from .multiplex import TMultiplexedProtocol, TMultiplexedProtocolFactory
|
||||
|
||||
if not PYPY:
|
||||
# enable cython binary by default for CPython.
|
||||
if CYTHON:
|
||||
from .cybin import TCyBinaryProtocol, TCyBinaryProtocolFactory
|
||||
TBinaryProtocol = TCyBinaryProtocol # noqa
|
||||
TBinaryProtocolFactory = TCyBinaryProtocolFactory # noqa
|
||||
else:
|
||||
# disable cython binary protocol for PYPY since it's slower.
|
||||
TCyBinaryProtocol = TBinaryProtocol
|
||||
TCyBinaryProtocolFactory = TBinaryProtocolFactory
|
||||
|
||||
__all__ = ['TBinaryProtocol', 'TBinaryProtocolFactory',
|
||||
'TCyBinaryProtocol', 'TCyBinaryProtocolFactory',
|
||||
'TJSONProtocol', 'TJSONProtocolFactory',
|
||||
'TMultiplexedProtocol', 'TMultiplexedProtocolFactory',
|
||||
'TCompactProtocol', 'TCompactProtocolFactory']
|
||||
@@ -1,401 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import struct
|
||||
|
||||
from .exc import TProtocolException
|
||||
from ..thrift import TType
|
||||
|
||||
# VERSION_MASK = 0xffff0000
|
||||
VERSION_MASK = -65536
|
||||
# VERSION_1 = 0x80010000
|
||||
VERSION_1 = -2147418112
|
||||
TYPE_MASK = 0x000000ff
|
||||
|
||||
|
||||
def pack_i8(byte):
|
||||
return struct.pack("!b", byte)
|
||||
|
||||
|
||||
def pack_i16(i16):
|
||||
return struct.pack("!h", i16)
|
||||
|
||||
|
||||
def pack_i32(i32):
|
||||
return struct.pack("!i", i32)
|
||||
|
||||
|
||||
def pack_i64(i64):
|
||||
return struct.pack("!q", i64)
|
||||
|
||||
|
||||
def pack_double(dub):
|
||||
return struct.pack("!d", dub)
|
||||
|
||||
|
||||
def pack_string(string):
|
||||
return struct.pack("!i%ds" % len(string), len(string), string)
|
||||
|
||||
|
||||
def unpack_i8(buf):
|
||||
return struct.unpack("!b", buf)[0]
|
||||
|
||||
|
||||
def unpack_i16(buf):
|
||||
return struct.unpack("!h", buf)[0]
|
||||
|
||||
|
||||
def unpack_i32(buf):
|
||||
return struct.unpack("!i", buf)[0]
|
||||
|
||||
|
||||
def unpack_i64(buf):
|
||||
return struct.unpack("!q", buf)[0]
|
||||
|
||||
|
||||
def unpack_double(buf):
|
||||
return struct.unpack("!d", buf)[0]
|
||||
|
||||
|
||||
def write_message_begin(outbuf, name, ttype, seqid, strict=True):
|
||||
if strict:
|
||||
outbuf.write(pack_i32(VERSION_1 | ttype))
|
||||
outbuf.write(pack_string(name.encode('utf-8')))
|
||||
else:
|
||||
outbuf.write(pack_string(name.encode('utf-8')))
|
||||
outbuf.write(pack_i8(ttype))
|
||||
|
||||
outbuf.write(pack_i32(seqid))
|
||||
|
||||
|
||||
def write_field_begin(outbuf, ttype, fid):
|
||||
outbuf.write(pack_i8(ttype) + pack_i16(fid))
|
||||
|
||||
|
||||
def write_field_stop(outbuf):
|
||||
outbuf.write(pack_i8(TType.STOP))
|
||||
|
||||
|
||||
def write_list_begin(outbuf, etype, size):
|
||||
outbuf.write(pack_i8(etype) + pack_i32(size))
|
||||
|
||||
|
||||
def write_map_begin(outbuf, ktype, vtype, size):
|
||||
outbuf.write(pack_i8(ktype) + pack_i8(vtype) + pack_i32(size))
|
||||
|
||||
|
||||
def write_val(outbuf, ttype, val, spec=None):
|
||||
if ttype == TType.BOOL:
|
||||
if val:
|
||||
outbuf.write(pack_i8(1))
|
||||
else:
|
||||
outbuf.write(pack_i8(0))
|
||||
|
||||
elif ttype == TType.BYTE:
|
||||
outbuf.write(pack_i8(val))
|
||||
|
||||
elif ttype == TType.I16:
|
||||
outbuf.write(pack_i16(val))
|
||||
|
||||
elif ttype == TType.I32:
|
||||
outbuf.write(pack_i32(val))
|
||||
|
||||
elif ttype == TType.I64:
|
||||
outbuf.write(pack_i64(val))
|
||||
|
||||
elif ttype == TType.DOUBLE:
|
||||
outbuf.write(pack_double(val))
|
||||
|
||||
elif ttype == TType.STRING:
|
||||
if not isinstance(val, bytes):
|
||||
val = val.encode('utf-8')
|
||||
outbuf.write(pack_string(val))
|
||||
|
||||
elif ttype == TType.SET or ttype == TType.LIST:
|
||||
if isinstance(spec, tuple):
|
||||
e_type, t_spec = spec[0], spec[1]
|
||||
else:
|
||||
e_type, t_spec = spec, None
|
||||
|
||||
val_len = len(val)
|
||||
write_list_begin(outbuf, e_type, val_len)
|
||||
for e_val in val:
|
||||
write_val(outbuf, e_type, e_val, t_spec)
|
||||
|
||||
elif ttype == TType.MAP:
|
||||
if isinstance(spec[0], int):
|
||||
k_type = spec[0]
|
||||
k_spec = None
|
||||
else:
|
||||
k_type, k_spec = spec[0]
|
||||
|
||||
if isinstance(spec[1], int):
|
||||
v_type = spec[1]
|
||||
v_spec = None
|
||||
else:
|
||||
v_type, v_spec = spec[1]
|
||||
|
||||
write_map_begin(outbuf, k_type, v_type, len(val))
|
||||
for k in iter(val):
|
||||
write_val(outbuf, k_type, k, k_spec)
|
||||
write_val(outbuf, v_type, val[k], v_spec)
|
||||
|
||||
elif ttype == TType.STRUCT:
|
||||
for fid in iter(val.thrift_spec):
|
||||
f_spec = val.thrift_spec[fid]
|
||||
if len(f_spec) == 3:
|
||||
f_type, f_name, f_req = f_spec
|
||||
f_container_spec = None
|
||||
else:
|
||||
f_type, f_name, f_container_spec, f_req = f_spec
|
||||
|
||||
v = getattr(val, f_name)
|
||||
if v is None:
|
||||
continue
|
||||
|
||||
write_field_begin(outbuf, f_type, fid)
|
||||
write_val(outbuf, f_type, v, f_container_spec)
|
||||
write_field_stop(outbuf)
|
||||
|
||||
|
||||
def read_message_begin(inbuf, strict=True):
|
||||
sz = unpack_i32(inbuf.read(4))
|
||||
if sz < 0:
|
||||
version = sz & VERSION_MASK
|
||||
if version != VERSION_1:
|
||||
raise TProtocolException(
|
||||
type=TProtocolException.BAD_VERSION,
|
||||
message='Bad version in read_message_begin: %d' % (sz))
|
||||
name_sz = unpack_i32(inbuf.read(4))
|
||||
name = inbuf.read(name_sz).decode('utf-8')
|
||||
|
||||
type_ = sz & TYPE_MASK
|
||||
else:
|
||||
if strict:
|
||||
raise TProtocolException(type=TProtocolException.BAD_VERSION,
|
||||
message='No protocol version header')
|
||||
|
||||
name = inbuf.read(sz).decode('utf-8')
|
||||
type_ = unpack_i8(inbuf.read(1))
|
||||
|
||||
seqid = unpack_i32(inbuf.read(4))
|
||||
|
||||
return name, type_, seqid
|
||||
|
||||
|
||||
def read_field_begin(inbuf):
|
||||
f_type = unpack_i8(inbuf.read(1))
|
||||
if f_type == TType.STOP:
|
||||
return f_type, 0
|
||||
|
||||
return f_type, unpack_i16(inbuf.read(2))
|
||||
|
||||
|
||||
def read_list_begin(inbuf):
|
||||
e_type = unpack_i8(inbuf.read(1))
|
||||
sz = unpack_i32(inbuf.read(4))
|
||||
return e_type, sz
|
||||
|
||||
|
||||
def read_map_begin(inbuf):
|
||||
k_type, v_type = unpack_i8(inbuf.read(1)), unpack_i8(inbuf.read(1))
|
||||
sz = unpack_i32(inbuf.read(4))
|
||||
return k_type, v_type, sz
|
||||
|
||||
|
||||
def read_val(inbuf, ttype, spec=None, decode_response=True):
|
||||
if ttype == TType.BOOL:
|
||||
return bool(unpack_i8(inbuf.read(1)))
|
||||
|
||||
elif ttype == TType.BYTE:
|
||||
return unpack_i8(inbuf.read(1))
|
||||
|
||||
elif ttype == TType.I16:
|
||||
return unpack_i16(inbuf.read(2))
|
||||
|
||||
elif ttype == TType.I32:
|
||||
return unpack_i32(inbuf.read(4))
|
||||
|
||||
elif ttype == TType.I64:
|
||||
return unpack_i64(inbuf.read(8))
|
||||
|
||||
elif ttype == TType.DOUBLE:
|
||||
return unpack_double(inbuf.read(8))
|
||||
|
||||
elif ttype == TType.STRING:
|
||||
sz = unpack_i32(inbuf.read(4))
|
||||
byte_payload = inbuf.read(sz)
|
||||
|
||||
# Since we cannot tell if we're getting STRING or BINARY
|
||||
# if not asked not to decode, try both
|
||||
if decode_response:
|
||||
try:
|
||||
return byte_payload.decode('utf-8')
|
||||
except UnicodeDecodeError:
|
||||
pass
|
||||
return byte_payload
|
||||
|
||||
elif ttype == TType.SET or ttype == TType.LIST:
|
||||
if isinstance(spec, tuple):
|
||||
v_type, v_spec = spec[0], spec[1]
|
||||
else:
|
||||
v_type, v_spec = spec, None
|
||||
|
||||
result = []
|
||||
r_type, sz = read_list_begin(inbuf)
|
||||
# the v_type is useless here since we already get it from spec
|
||||
if r_type != v_type:
|
||||
for _ in range(sz):
|
||||
skip(inbuf, r_type)
|
||||
return []
|
||||
|
||||
for i in range(sz):
|
||||
result.append(read_val(inbuf, v_type, v_spec, decode_response))
|
||||
return result
|
||||
|
||||
elif ttype == TType.MAP:
|
||||
if isinstance(spec[0], int):
|
||||
k_type = spec[0]
|
||||
k_spec = None
|
||||
else:
|
||||
k_type, k_spec = spec[0]
|
||||
|
||||
if isinstance(spec[1], int):
|
||||
v_type = spec[1]
|
||||
v_spec = None
|
||||
else:
|
||||
v_type, v_spec = spec[1]
|
||||
|
||||
result = {}
|
||||
sk_type, sv_type, sz = read_map_begin(inbuf)
|
||||
if sk_type != k_type or sv_type != v_type:
|
||||
for _ in range(sz):
|
||||
skip(inbuf, sk_type)
|
||||
skip(inbuf, sv_type)
|
||||
return {}
|
||||
|
||||
for i in range(sz):
|
||||
k_val = read_val(inbuf, k_type, k_spec, decode_response)
|
||||
v_val = read_val(inbuf, v_type, v_spec, decode_response)
|
||||
result[k_val] = v_val
|
||||
|
||||
return result
|
||||
|
||||
elif ttype == TType.STRUCT:
|
||||
obj = spec()
|
||||
read_struct(inbuf, obj, decode_response)
|
||||
return obj
|
||||
|
||||
|
||||
def read_struct(inbuf, obj, decode_response=True):
|
||||
while True:
|
||||
f_type, fid = read_field_begin(inbuf)
|
||||
if f_type == TType.STOP:
|
||||
break
|
||||
|
||||
if fid not in obj.thrift_spec:
|
||||
skip(inbuf, f_type)
|
||||
continue
|
||||
|
||||
if len(obj.thrift_spec[fid]) == 3:
|
||||
sf_type, f_name, f_req = obj.thrift_spec[fid]
|
||||
f_container_spec = None
|
||||
else:
|
||||
sf_type, f_name, f_container_spec, f_req = obj.thrift_spec[fid]
|
||||
|
||||
# it really should equal here. but since we already wasted
|
||||
# space storing the duplicate info, let's check it.
|
||||
if f_type != sf_type:
|
||||
skip(inbuf, f_type)
|
||||
continue
|
||||
|
||||
setattr(obj, f_name,
|
||||
read_val(inbuf, f_type, f_container_spec, decode_response))
|
||||
|
||||
|
||||
def skip(inbuf, ftype):
|
||||
if ftype == TType.BOOL or ftype == TType.BYTE:
|
||||
inbuf.read(1)
|
||||
|
||||
elif ftype == TType.I16:
|
||||
inbuf.read(2)
|
||||
|
||||
elif ftype == TType.I32:
|
||||
inbuf.read(4)
|
||||
|
||||
elif ftype == TType.I64:
|
||||
inbuf.read(8)
|
||||
|
||||
elif ftype == TType.DOUBLE:
|
||||
inbuf.read(8)
|
||||
|
||||
elif ftype == TType.STRING:
|
||||
inbuf.read(unpack_i32(inbuf.read(4)))
|
||||
|
||||
elif ftype == TType.SET or ftype == TType.LIST:
|
||||
v_type, sz = read_list_begin(inbuf)
|
||||
for i in range(sz):
|
||||
skip(inbuf, v_type)
|
||||
|
||||
elif ftype == TType.MAP:
|
||||
k_type, v_type, sz = read_map_begin(inbuf)
|
||||
for i in range(sz):
|
||||
skip(inbuf, k_type)
|
||||
skip(inbuf, v_type)
|
||||
|
||||
elif ftype == TType.STRUCT:
|
||||
while True:
|
||||
f_type, fid = read_field_begin(inbuf)
|
||||
if f_type == TType.STOP:
|
||||
break
|
||||
skip(inbuf, f_type)
|
||||
|
||||
|
||||
class TBinaryProtocol(object):
|
||||
"""Binary implementation of the Thrift protocol driver."""
|
||||
|
||||
def __init__(self, trans,
|
||||
strict_read=True, strict_write=True,
|
||||
decode_response=True):
|
||||
self.trans = trans
|
||||
self.strict_read = strict_read
|
||||
self.strict_write = strict_write
|
||||
self.decode_response = decode_response
|
||||
|
||||
def skip(self, ttype):
|
||||
skip(self.trans, ttype)
|
||||
|
||||
def read_message_begin(self):
|
||||
api, ttype, seqid = read_message_begin(
|
||||
self.trans, strict=self.strict_read)
|
||||
return api, ttype, seqid
|
||||
|
||||
def read_message_end(self):
|
||||
pass
|
||||
|
||||
def write_message_begin(self, name, ttype, seqid):
|
||||
write_message_begin(self.trans, name, ttype, seqid,
|
||||
strict=self.strict_write)
|
||||
|
||||
def write_message_end(self):
|
||||
pass
|
||||
|
||||
def read_struct(self, obj):
|
||||
return read_struct(self.trans, obj, self.decode_response)
|
||||
|
||||
def write_struct(self, obj):
|
||||
write_val(self.trans, TType.STRUCT, obj)
|
||||
|
||||
|
||||
class TBinaryProtocolFactory(object):
|
||||
def __init__(self, strict_read=True, strict_write=True,
|
||||
decode_response=True):
|
||||
self.strict_read = strict_read
|
||||
self.strict_write = strict_write
|
||||
self.decode_response = decode_response
|
||||
|
||||
def get_protocol(self, trans):
|
||||
return TBinaryProtocol(trans,
|
||||
self.strict_read, self.strict_write,
|
||||
self.decode_response)
|
||||
@@ -1,565 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import array
|
||||
from struct import pack, unpack
|
||||
|
||||
from thriftpy._compat import PY3
|
||||
from .exc import TProtocolException
|
||||
from ..thrift import TException
|
||||
from ..thrift import TType
|
||||
|
||||
CLEAR = 0
|
||||
FIELD_WRITE = 1
|
||||
VALUE_WRITE = 2
|
||||
CONTAINER_WRITE = 3
|
||||
BOOL_WRITE = 4
|
||||
FIELD_READ = 5
|
||||
CONTAINER_READ = 6
|
||||
VALUE_READ = 7
|
||||
BOOL_READ = 8
|
||||
|
||||
|
||||
def check_integer_limits(i, bits):
|
||||
if bits == 8 and (i < -128 or i > 127):
|
||||
raise TProtocolException(TProtocolException.INVALID_DATA,
|
||||
"i8 requires -128 <= number <= 127")
|
||||
elif bits == 16 and (i < -32768 or i > 32767):
|
||||
raise TProtocolException(TProtocolException.INVALID_DATA,
|
||||
"i16 requires -32768 <= number <= 32767")
|
||||
elif bits == 32 and (i < -2147483648 or i > 2147483647):
|
||||
raise TProtocolException(
|
||||
TProtocolException.INVALID_DATA,
|
||||
"i32 requires -2147483648 <= number <= 2147483647")
|
||||
elif bits == 64 and (i < -9223372036854775808 or i > 9223372036854775807):
|
||||
raise TProtocolException(
|
||||
TProtocolException.INVALID_DATA,
|
||||
"i64 requires -9223372036854775808 <= number <= \
|
||||
9223372036854775807")
|
||||
|
||||
|
||||
def make_zig_zag(n, bits):
|
||||
check_integer_limits(n, bits)
|
||||
return (n << 1) ^ (n >> (bits - 1))
|
||||
|
||||
|
||||
def from_zig_zag(n):
|
||||
return (n >> 1) ^ -(n & 1)
|
||||
|
||||
|
||||
def write_varint(trans, n):
|
||||
out = []
|
||||
while True:
|
||||
if n & ~0x7f == 0:
|
||||
out.append(n)
|
||||
break
|
||||
else:
|
||||
out.append((n & 0xff) | 0x80)
|
||||
n = n >> 7
|
||||
data = array.array('B', out).tostring()
|
||||
|
||||
if PY3:
|
||||
trans.write(data)
|
||||
else:
|
||||
trans.write(bytes(data))
|
||||
|
||||
|
||||
def read_varint(trans):
|
||||
result = 0
|
||||
shift = 0
|
||||
|
||||
while True:
|
||||
x = trans.read(1)
|
||||
byte = ord(x)
|
||||
result |= (byte & 0x7f) << shift
|
||||
if byte >> 7 == 0:
|
||||
return result
|
||||
shift += 7
|
||||
|
||||
|
||||
class CompactType(object):
|
||||
STOP = 0x00
|
||||
TRUE = 0x01
|
||||
FALSE = 0x02
|
||||
BYTE = 0x03
|
||||
I16 = 0x04
|
||||
I32 = 0x05
|
||||
I64 = 0x06
|
||||
DOUBLE = 0x07
|
||||
BINARY = 0x08
|
||||
LIST = 0x09
|
||||
SET = 0x0A
|
||||
MAP = 0x0B
|
||||
STRUCT = 0x0C
|
||||
|
||||
|
||||
CTYPES = {
|
||||
TType.STOP: CompactType.STOP,
|
||||
TType.BOOL: CompactType.TRUE,
|
||||
TType.BYTE: CompactType.BYTE,
|
||||
TType.I16: CompactType.I16,
|
||||
TType.I32: CompactType.I32,
|
||||
TType.I64: CompactType.I64,
|
||||
TType.DOUBLE: CompactType.DOUBLE,
|
||||
TType.STRING: CompactType.BINARY,
|
||||
TType.STRUCT: CompactType.STRUCT,
|
||||
TType.LIST: CompactType.LIST,
|
||||
TType.SET: CompactType.SET,
|
||||
TType.MAP: CompactType.MAP
|
||||
}
|
||||
TTYPES = dict((v, k) for k, v in CTYPES.items())
|
||||
TTYPES[CompactType.FALSE] = TType.BOOL
|
||||
|
||||
|
||||
class TCompactProtocol(object):
|
||||
"""Compact implementation of the Thrift protocol driver."""
|
||||
PROTOCOL_ID = 0x82
|
||||
VERSION = 1
|
||||
VERSION_MASK = 0x1f
|
||||
TYPE_MASK = 0xe0
|
||||
TYPE_BITS = 0x07
|
||||
TYPE_SHIFT_AMOUNT = 5
|
||||
|
||||
def __init__(self, trans, decode_response=True):
|
||||
self.trans = trans
|
||||
self._last_fid = 0
|
||||
self._bool_fid = None
|
||||
self._bool_value = None
|
||||
self._structs = []
|
||||
self.decode_response = decode_response
|
||||
|
||||
def _get_ttype(self, byte):
|
||||
return TTYPES[byte & 0x0f]
|
||||
|
||||
def _read_size(self):
|
||||
result = read_varint(self.trans)
|
||||
if result < 0:
|
||||
raise TException("Length < 0")
|
||||
return result
|
||||
|
||||
def read_message_begin(self):
|
||||
proto_id = self.read_ubyte()
|
||||
if proto_id != self.PROTOCOL_ID:
|
||||
raise TProtocolException(TProtocolException.BAD_VERSION,
|
||||
'Bad protocol id in the message: %d'
|
||||
% proto_id)
|
||||
|
||||
ver_type = self.read_ubyte()
|
||||
type = (ver_type >> self.TYPE_SHIFT_AMOUNT) & self.TYPE_BITS
|
||||
version = ver_type & self.VERSION_MASK
|
||||
if version != self.VERSION:
|
||||
raise TProtocolException(TProtocolException.BAD_VERSION,
|
||||
'Bad version: %d (expect %d)'
|
||||
% (version, self.VERSION))
|
||||
seqid = read_varint(self.trans)
|
||||
name = self.read_string()
|
||||
return name, type, seqid
|
||||
|
||||
def read_message_end(self):
|
||||
assert len(self._structs) == 0
|
||||
|
||||
def read_field_begin(self):
|
||||
type = self.read_ubyte()
|
||||
if type & 0x0f == TType.STOP:
|
||||
return None, 0, 0
|
||||
|
||||
delta = type >> 4
|
||||
if delta == 0:
|
||||
fid = from_zig_zag(read_varint(self.trans))
|
||||
else:
|
||||
fid = self._last_fid + delta
|
||||
self._last_fid = fid
|
||||
|
||||
type = type & 0x0f
|
||||
if type == CompactType.TRUE:
|
||||
self._bool_value = True
|
||||
elif type == CompactType.FALSE:
|
||||
self._bool_value = False
|
||||
|
||||
return None, self._get_ttype(type), fid
|
||||
|
||||
def read_field_end(self):
|
||||
pass
|
||||
|
||||
def read_struct_begin(self):
|
||||
self._structs.append(self._last_fid)
|
||||
self._last_fid = 0
|
||||
|
||||
def read_struct_end(self):
|
||||
self._last_fid = self._structs.pop()
|
||||
|
||||
def read_map_begin(self):
|
||||
size = self._read_size()
|
||||
types = 0
|
||||
if size > 0:
|
||||
types = self.read_ubyte()
|
||||
vtype = self._get_ttype(types)
|
||||
ktype = self._get_ttype(types >> 4)
|
||||
return (ktype, vtype, size)
|
||||
|
||||
def read_collection_begin(self):
|
||||
size_type = self.read_ubyte()
|
||||
size = size_type >> 4
|
||||
type = self._get_ttype(size_type)
|
||||
if size == 15:
|
||||
size = self._read_size()
|
||||
return type, size
|
||||
|
||||
def read_collection_end(self):
|
||||
pass
|
||||
|
||||
def read_byte(self):
|
||||
result, = unpack('!b', self.trans.read(1))
|
||||
return result
|
||||
|
||||
def read_ubyte(self):
|
||||
result, = unpack('!B', self.trans.read(1))
|
||||
return result
|
||||
|
||||
def read_int(self):
|
||||
return from_zig_zag(read_varint(self.trans))
|
||||
|
||||
def read_double(self):
|
||||
buff = self.trans.read(8)
|
||||
val, = unpack('<d', buff)
|
||||
return val
|
||||
|
||||
def read_string(self):
|
||||
len = self._read_size()
|
||||
byte_payload = self.trans.read(len)
|
||||
|
||||
if self.decode_response:
|
||||
try:
|
||||
byte_payload = byte_payload.decode('utf-8')
|
||||
except UnicodeDecodeError:
|
||||
pass
|
||||
return byte_payload
|
||||
|
||||
def read_bool(self):
|
||||
if self._bool_value is not None:
|
||||
result = self._bool_value
|
||||
self._bool_value = None
|
||||
return result
|
||||
return self.read_byte() == CompactType.TRUE
|
||||
|
||||
def read_struct(self, obj):
|
||||
self.read_struct_begin()
|
||||
while True:
|
||||
fname, ftype, fid = self.read_field_begin()
|
||||
if ftype == TType.STOP:
|
||||
break
|
||||
|
||||
if fid not in obj.thrift_spec:
|
||||
self.skip(ftype)
|
||||
continue
|
||||
|
||||
try:
|
||||
field = obj.thrift_spec[fid]
|
||||
except IndexError:
|
||||
self.skip(ftype)
|
||||
raise
|
||||
else:
|
||||
if field is not None and ftype == field[0]:
|
||||
fname = field[1]
|
||||
fspec = field[2]
|
||||
val = self.read_val(ftype, fspec)
|
||||
setattr(obj, fname, val)
|
||||
else:
|
||||
self.skip(ftype)
|
||||
self.read_field_end()
|
||||
self.read_struct_end()
|
||||
|
||||
def read_val(self, ttype, spec=None):
|
||||
if ttype == TType.BOOL:
|
||||
return self.read_bool()
|
||||
|
||||
elif ttype == TType.BYTE:
|
||||
return self.read_byte()
|
||||
|
||||
elif ttype in (TType.I16, TType.I32, TType.I64):
|
||||
return self.read_int()
|
||||
|
||||
elif ttype == TType.DOUBLE:
|
||||
return self.read_double()
|
||||
|
||||
elif ttype == TType.STRING:
|
||||
return self.read_string()
|
||||
|
||||
elif ttype in (TType.LIST, TType.SET):
|
||||
if isinstance(spec, tuple):
|
||||
v_type, v_spec = spec[0], spec[1]
|
||||
else:
|
||||
v_type, v_spec = spec, None
|
||||
result = []
|
||||
r_type, sz = self.read_collection_begin()
|
||||
|
||||
for i in range(sz):
|
||||
result.append(self.read_val(v_type, v_spec))
|
||||
|
||||
self.read_collection_end()
|
||||
return result
|
||||
|
||||
elif ttype == TType.MAP:
|
||||
if isinstance(spec[0], int):
|
||||
k_type = spec[0]
|
||||
k_spec = None
|
||||
else:
|
||||
k_type, k_spec = spec[0]
|
||||
|
||||
if isinstance(spec[1], int):
|
||||
v_type = spec[1]
|
||||
v_spec = None
|
||||
else:
|
||||
v_type, v_spec = spec[1]
|
||||
|
||||
result = {}
|
||||
sk_type, sv_type, sz = self.read_map_begin()
|
||||
if sk_type != k_type or sv_type != v_type:
|
||||
for _ in range(sz):
|
||||
self.skip(sk_type)
|
||||
self.skip(sv_type)
|
||||
self.read_collection_end()
|
||||
return {}
|
||||
|
||||
for i in range(sz):
|
||||
k_val = self.read_val(k_type, k_spec)
|
||||
v_val = self.read_val(v_type, v_spec)
|
||||
result[k_val] = v_val
|
||||
self.read_collection_end()
|
||||
return result
|
||||
|
||||
elif ttype == TType.STRUCT:
|
||||
obj = spec()
|
||||
self.read_struct(obj)
|
||||
return obj
|
||||
|
||||
def _write_size(self, i32):
|
||||
write_varint(self.trans, i32)
|
||||
|
||||
def _write_field_header(self, type, fid):
|
||||
delta = fid - self._last_fid
|
||||
if 0 < delta <= 15:
|
||||
self.write_ubyte(delta << 4 | type)
|
||||
else:
|
||||
self.write_byte(type)
|
||||
self.write_i16(fid)
|
||||
self._last_fid = fid
|
||||
|
||||
def write_message_begin(self, name, type, seqid):
|
||||
self.write_ubyte(self.PROTOCOL_ID)
|
||||
self.write_ubyte(self.VERSION | (type << self.TYPE_SHIFT_AMOUNT))
|
||||
write_varint(self.trans, seqid)
|
||||
self.write_string(name)
|
||||
|
||||
def write_message_end(self):
|
||||
pass
|
||||
|
||||
def write_field_stop(self):
|
||||
self.write_byte(0)
|
||||
|
||||
def write_field_begin(self, name, type, fid):
|
||||
if type == TType.BOOL:
|
||||
self._bool_fid = fid
|
||||
else:
|
||||
self._write_field_header(CTYPES[type], fid)
|
||||
|
||||
def write_field_end(self):
|
||||
pass
|
||||
|
||||
def write_struct_begin(self):
|
||||
self._structs.append(self._last_fid)
|
||||
self._last_fid = 0
|
||||
|
||||
def write_struct_end(self):
|
||||
self._last_fid = self._structs.pop()
|
||||
|
||||
def write_collection_begin(self, etype, size):
|
||||
if size <= 14:
|
||||
self.write_ubyte(size << 4 | CTYPES[etype])
|
||||
else:
|
||||
self.write_ubyte(0xf0 | CTYPES[etype])
|
||||
self._write_size(size)
|
||||
|
||||
def write_map_begin(self, ktype, vtype, size):
|
||||
if size == 0:
|
||||
self.write_byte(0)
|
||||
else:
|
||||
self._write_size(size)
|
||||
self.write_ubyte(CTYPES[ktype] << 4 | CTYPES[vtype])
|
||||
|
||||
def write_collection_end(self):
|
||||
pass
|
||||
|
||||
def write_ubyte(self, byte):
|
||||
self.trans.write(pack('!B', byte))
|
||||
|
||||
def write_byte(self, byte):
|
||||
self.trans.write(pack('!b', byte))
|
||||
|
||||
def write_bool(self, bool):
|
||||
if self._bool_fid and self._bool_fid > self._last_fid \
|
||||
and self._bool_fid - self._last_fid <= 15:
|
||||
if bool:
|
||||
ctype = CompactType.TRUE
|
||||
else:
|
||||
ctype = CompactType.FALSE
|
||||
self._write_field_header(ctype, self._bool_fid)
|
||||
else:
|
||||
if bool:
|
||||
self.write_byte(CompactType.TRUE)
|
||||
else:
|
||||
self.write_byte(CompactType.FALSE)
|
||||
|
||||
def write_i16(self, i16):
|
||||
write_varint(self.trans, make_zig_zag(i16, 16))
|
||||
|
||||
def write_i32(self, i32):
|
||||
write_varint(self.trans, make_zig_zag(i32, 32))
|
||||
|
||||
def write_i64(self, i64):
|
||||
write_varint(self.trans, make_zig_zag(i64, 64))
|
||||
|
||||
def write_double(self, dub):
|
||||
self.trans.write(pack('<d', dub))
|
||||
|
||||
def write_string(self, s):
|
||||
if not isinstance(s, bytes):
|
||||
s = s.encode('utf-8')
|
||||
self._write_size(len(s))
|
||||
self.trans.write(s)
|
||||
|
||||
def write_struct(self, obj):
|
||||
self.write_struct_begin()
|
||||
|
||||
for field in obj.thrift_spec:
|
||||
if field is None:
|
||||
continue
|
||||
fspec = obj.thrift_spec[field]
|
||||
if len(fspec) == 3:
|
||||
ftype, fname, freq = fspec
|
||||
f_container_spec = None
|
||||
else:
|
||||
ftype, fname, f_container_spec, f_req = fspec
|
||||
val = getattr(obj, fname)
|
||||
if val is None:
|
||||
continue
|
||||
|
||||
self.write_field_begin(fname, ftype, field)
|
||||
self.write_val(ftype, val, f_container_spec)
|
||||
self.write_field_end()
|
||||
self.write_field_stop()
|
||||
self.write_struct_end()
|
||||
|
||||
def write_val(self, ttype, val, spec=None):
|
||||
|
||||
if ttype == TType.BOOL:
|
||||
self.write_bool(val)
|
||||
|
||||
elif ttype == TType.BYTE:
|
||||
self.write_byte(val)
|
||||
|
||||
elif ttype == TType.I16:
|
||||
self.write_i16(val)
|
||||
|
||||
elif ttype == TType.I32:
|
||||
self.write_i32(val)
|
||||
|
||||
elif ttype == TType.I64:
|
||||
self.write_i64(val)
|
||||
|
||||
elif ttype == TType.DOUBLE:
|
||||
self.write_double(val)
|
||||
|
||||
elif ttype == TType.STRING:
|
||||
self.write_string(val)
|
||||
|
||||
elif ttype == TType.LIST or ttype == TType.SET:
|
||||
if isinstance(spec, tuple):
|
||||
e_type, t_spec = spec[0], spec[1]
|
||||
else:
|
||||
e_type, t_spec = spec, None
|
||||
|
||||
val_len = len(val)
|
||||
self.write_collection_begin(e_type, val_len)
|
||||
for e_val in val:
|
||||
self.write_val(e_type, e_val, t_spec)
|
||||
self.write_collection_end()
|
||||
|
||||
elif ttype == TType.MAP:
|
||||
if isinstance(spec[0], int):
|
||||
k_type = spec[0]
|
||||
k_spec = None
|
||||
else:
|
||||
k_type, k_spec = spec[0]
|
||||
|
||||
if isinstance(spec[1], int):
|
||||
v_type = spec[1]
|
||||
v_spec = None
|
||||
else:
|
||||
v_type, v_spec = spec[1]
|
||||
|
||||
self.write_map_begin(k_type, v_type, len(val))
|
||||
for k in iter(val):
|
||||
self.write_val(k_type, k, k_spec)
|
||||
self.write_val(v_type, val[k], v_spec)
|
||||
self.write_collection_end()
|
||||
|
||||
elif ttype == TType.STRUCT:
|
||||
self.write_struct(val)
|
||||
|
||||
def skip(self, ttype):
|
||||
if ttype == TType.STOP:
|
||||
return
|
||||
|
||||
elif ttype == TType.BOOL:
|
||||
self.read_bool()
|
||||
|
||||
elif ttype == TType.BYTE:
|
||||
self.read_byte()
|
||||
|
||||
elif ttype in (TType.I16, TType.I32, TType.I64):
|
||||
from_zig_zag(read_varint(self.trans))
|
||||
|
||||
elif ttype == TType.DOUBLE:
|
||||
self.read_double()
|
||||
|
||||
elif ttype == TType.STRING:
|
||||
self.read_string()
|
||||
|
||||
elif ttype == TType.STRUCT:
|
||||
name = self.read_struct_begin()
|
||||
while True:
|
||||
(name, ttype, id) = self.read_field_begin()
|
||||
if ttype == TType.STOP:
|
||||
break
|
||||
self.skip(ttype)
|
||||
self.read_field_end()
|
||||
self.read_struct_end()
|
||||
|
||||
elif ttype == TType.MAP:
|
||||
ktype, vtype, size = self.read_map_begin()
|
||||
for i in range(size):
|
||||
self.skip(ktype)
|
||||
self.skip(vtype)
|
||||
self.read_collection_end()
|
||||
|
||||
elif ttype == TType.SET:
|
||||
etype, size = self.read_collection_begin()
|
||||
for i in range(size):
|
||||
self.skip(etype)
|
||||
self.read_collection_end()
|
||||
|
||||
elif ttype == TType.LIST:
|
||||
etype, size = self.read_collection_begin()
|
||||
for i in range(size):
|
||||
self.skip(etype)
|
||||
self.read_collection_end()
|
||||
|
||||
|
||||
class TCompactProtocolFactory(object):
|
||||
def __init__(self, decode_response=True):
|
||||
self.decode_response = decode_response
|
||||
|
||||
def get_protocol(self, trans):
|
||||
return TCompactProtocol(trans, decode_response=self.decode_response)
|
||||
@@ -1,496 +0,0 @@
|
||||
from cpython cimport
|
||||
|
||||
bool
|
||||
from libc.stdint cimport
|
||||
|
||||
int16_t, int32_t, int64_t
|
||||
from libc.stdlib cimport
|
||||
|
||||
free, malloc
|
||||
|
||||
from thriftpy.transport.cybase cimport
|
||||
|
||||
CyTransportBase, STACK_STRING_LEN
|
||||
from ..thrift import TDecodeException
|
||||
|
||||
cdef extern from "endian_port.h":
|
||||
int16_t htobe16(int16_t n)
|
||||
int32_t htobe32(int32_t n)
|
||||
int64_t htobe64(int64_t n)
|
||||
int16_t be16toh(int16_t n)
|
||||
int32_t be32toh(int32_t n)
|
||||
int64_t be64toh(int64_t n)
|
||||
|
||||
DEF VERSION_MASK = -65536
|
||||
DEF VERSION_1 = -2147418112
|
||||
DEF TYPE_MASK = 0x000000ff
|
||||
|
||||
ctypedef enum TType:
|
||||
T_STOP = 0,
|
||||
T_VOID = 1,
|
||||
T_BOOL = 2,
|
||||
T_BYTE = 3,
|
||||
T_I08 = 3,
|
||||
T_I16 = 6,
|
||||
T_I32 = 8,
|
||||
T_U64 = 9,
|
||||
T_I64 = 10,
|
||||
T_DOUBLE = 4,
|
||||
T_STRING = 11,
|
||||
T_UTF7 = 11,
|
||||
T_NARY = 11
|
||||
T_STRUCT = 12,
|
||||
T_MAP = 13,
|
||||
T_SET = 14,
|
||||
T_LIST = 15,
|
||||
T_UTF8 = 16,
|
||||
T_UTF16 = 17
|
||||
|
||||
class ProtocolError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
cdef inline char read_i08(CyTransportBase buf) except? -1:
|
||||
cdef char data = 0
|
||||
buf.c_read(1, &data)
|
||||
return data
|
||||
|
||||
|
||||
cdef inline int16_t read_i16(CyTransportBase buf) except? -1:
|
||||
cdef char data[2]
|
||||
buf.c_read(2, data)
|
||||
return be16toh((<int16_t*>data)[0])
|
||||
|
||||
|
||||
cdef inline int32_t read_i32(CyTransportBase buf) except? -1:
|
||||
cdef char data[4]
|
||||
buf.c_read(4, data)
|
||||
return be32toh((<int32_t*>data)[0])
|
||||
|
||||
|
||||
cdef inline int64_t read_i64(CyTransportBase buf) except? -1:
|
||||
cdef char data[8]
|
||||
buf.c_read(8, data)
|
||||
return be64toh((<int64_t*>data)[0])
|
||||
|
||||
|
||||
cdef inline int write_i08(CyTransportBase buf, char val) except -1:
|
||||
buf.c_write(&val, 1)
|
||||
return 0
|
||||
|
||||
|
||||
cdef inline int write_i16(CyTransportBase buf, int16_t val) except -1:
|
||||
val = htobe16(val)
|
||||
buf.c_write(<char*>(&val), 2)
|
||||
return 0
|
||||
|
||||
|
||||
cdef inline int write_i32(CyTransportBase buf, int32_t val) except -1:
|
||||
val = htobe32(val)
|
||||
buf.c_write(<char*>(&val), 4)
|
||||
return 0
|
||||
|
||||
|
||||
cdef inline int write_i64(CyTransportBase buf, int64_t val) except -1:
|
||||
val = htobe64(val)
|
||||
buf.c_write(<char*>(&val), 8)
|
||||
return 0
|
||||
|
||||
|
||||
cdef inline int write_double(CyTransportBase buf, double val) except -1:
|
||||
cdef int64_t v = htobe64((<int64_t*>(&val))[0])
|
||||
buf.c_write(<char*>(&v), 8)
|
||||
return 0
|
||||
|
||||
|
||||
cdef inline write_list(CyTransportBase buf, object val, spec):
|
||||
cdef TType e_type
|
||||
cdef int val_len
|
||||
|
||||
if isinstance(spec, int):
|
||||
e_type = spec
|
||||
e_spec = None
|
||||
else:
|
||||
e_type = spec[0]
|
||||
e_spec = spec[1]
|
||||
|
||||
val_len = len(val)
|
||||
write_i08(buf, e_type)
|
||||
write_i32(buf, val_len)
|
||||
|
||||
for e_val in val:
|
||||
c_write_val(buf, e_type, e_val, e_spec)
|
||||
|
||||
|
||||
cdef inline write_string(CyTransportBase buf, bytes val):
|
||||
cdef int val_len = len(val)
|
||||
write_i32(buf, val_len)
|
||||
|
||||
buf.c_write(<char*>val, val_len)
|
||||
|
||||
|
||||
cdef inline write_dict(CyTransportBase buf, object val, spec):
|
||||
cdef int val_len
|
||||
cdef TType v_type, k_type
|
||||
|
||||
key = spec[0]
|
||||
if isinstance(key, int):
|
||||
k_type = key
|
||||
k_spec = None
|
||||
else:
|
||||
k_type = key[0]
|
||||
k_spec = key[1]
|
||||
|
||||
value = spec[1]
|
||||
if isinstance(value, int):
|
||||
v_type = value
|
||||
v_spec = None
|
||||
else:
|
||||
v_type = value[0]
|
||||
v_spec = value[1]
|
||||
|
||||
val_len = len(val)
|
||||
|
||||
write_i08(buf, k_type)
|
||||
write_i08(buf, v_type)
|
||||
write_i32(buf, val_len)
|
||||
|
||||
for k, v in val.items():
|
||||
c_write_val(buf, k_type, k, k_spec)
|
||||
c_write_val(buf, v_type, v, v_spec)
|
||||
|
||||
|
||||
cdef inline read_struct(CyTransportBase buf, obj, decode_response=True):
|
||||
cdef dict field_specs = obj.thrift_spec
|
||||
cdef int fid
|
||||
cdef TType field_type, ttype
|
||||
cdef tuple field_spec
|
||||
cdef str name
|
||||
|
||||
while True:
|
||||
field_type = <TType>read_i08(buf)
|
||||
if field_type == T_STOP:
|
||||
break
|
||||
|
||||
fid = read_i16(buf)
|
||||
if fid not in field_specs:
|
||||
skip(buf, field_type)
|
||||
continue
|
||||
|
||||
field_spec = field_specs[fid]
|
||||
ttype = field_spec[0]
|
||||
if field_type != ttype:
|
||||
skip(buf, field_type)
|
||||
continue
|
||||
|
||||
name = field_spec[1]
|
||||
if len(field_spec) <= 3:
|
||||
spec = None
|
||||
else:
|
||||
spec = field_spec[2]
|
||||
|
||||
setattr(obj, name, c_read_val(buf, ttype, spec, decode_response))
|
||||
|
||||
return obj
|
||||
|
||||
|
||||
cdef inline write_struct(CyTransportBase buf, obj):
|
||||
cdef int fid
|
||||
cdef TType f_type
|
||||
cdef dict thrift_spec = obj.thrift_spec
|
||||
cdef tuple field_spec
|
||||
cdef str f_name
|
||||
|
||||
for fid, field_spec in thrift_spec.items():
|
||||
f_type = field_spec[0]
|
||||
f_name = field_spec[1]
|
||||
if len(field_spec) <= 3:
|
||||
container_spec = None
|
||||
else:
|
||||
container_spec = field_spec[2]
|
||||
|
||||
v = getattr(obj, f_name)
|
||||
if v is None:
|
||||
continue
|
||||
|
||||
write_i08(buf, f_type)
|
||||
write_i16(buf, fid)
|
||||
try:
|
||||
c_write_val(buf, f_type, v, container_spec)
|
||||
except (TypeError, AttributeError, AssertionError, OverflowError):
|
||||
raise TDecodeException(obj.__class__.__name__, fid, f_name, v,
|
||||
f_type, container_spec)
|
||||
|
||||
write_i08(buf, T_STOP)
|
||||
|
||||
|
||||
cdef inline c_read_binary(CyTransportBase buf, int32_t size):
|
||||
cdef char string_val[STACK_STRING_LEN]
|
||||
|
||||
if size > STACK_STRING_LEN:
|
||||
data = <char*>malloc(size)
|
||||
buf.c_read(size, data)
|
||||
py_data = data[:size]
|
||||
free(data)
|
||||
else:
|
||||
buf.c_read(size, string_val)
|
||||
py_data = string_val[:size]
|
||||
|
||||
return py_data
|
||||
|
||||
|
||||
cdef inline c_read_string(CyTransportBase buf, int32_t size):
|
||||
py_data = c_read_binary(buf, size)
|
||||
try:
|
||||
return py_data.decode("utf-8")
|
||||
except:
|
||||
return py_data
|
||||
|
||||
|
||||
cdef c_read_val(CyTransportBase buf, TType ttype, spec=None,
|
||||
decode_response=True):
|
||||
cdef int size
|
||||
cdef int64_t n
|
||||
cdef TType v_type, k_type, orig_type, orig_key_type
|
||||
|
||||
if ttype == T_BOOL:
|
||||
return <bint>read_i08(buf)
|
||||
|
||||
elif ttype == T_I08:
|
||||
return read_i08(buf)
|
||||
|
||||
elif ttype == T_I16:
|
||||
return read_i16(buf)
|
||||
|
||||
elif ttype == T_I32:
|
||||
return read_i32(buf)
|
||||
|
||||
elif ttype == T_I64:
|
||||
return read_i64(buf)
|
||||
|
||||
elif ttype == T_DOUBLE:
|
||||
n = read_i64(buf)
|
||||
return (<double*>(&n))[0]
|
||||
|
||||
elif ttype == T_STRING:
|
||||
size = read_i32(buf)
|
||||
if decode_response:
|
||||
return c_read_string(buf, size)
|
||||
else:
|
||||
return c_read_binary(buf, size)
|
||||
|
||||
elif ttype == T_SET or ttype == T_LIST:
|
||||
if isinstance(spec, int):
|
||||
v_type = spec
|
||||
v_spec = None
|
||||
else:
|
||||
v_type = spec[0]
|
||||
v_spec = spec[1]
|
||||
|
||||
orig_type = <TType>read_i08(buf)
|
||||
size = read_i32(buf)
|
||||
|
||||
if orig_type != v_type:
|
||||
for _ in range(size):
|
||||
skip(buf, orig_type)
|
||||
return []
|
||||
|
||||
return [c_read_val(buf, v_type, v_spec, decode_response)
|
||||
for _ in range(size)]
|
||||
|
||||
elif ttype == T_MAP:
|
||||
key = spec[0]
|
||||
if isinstance(key, int):
|
||||
k_type = key
|
||||
k_spec = None
|
||||
else:
|
||||
k_type = key[0]
|
||||
k_spec = key[1]
|
||||
|
||||
value = spec[1]
|
||||
if isinstance(value, int):
|
||||
v_type = value
|
||||
v_spec = None
|
||||
else:
|
||||
v_type = value[0]
|
||||
v_spec = value[1]
|
||||
|
||||
orig_key_type = <TType>read_i08(buf)
|
||||
orig_type = <TType>read_i08(buf)
|
||||
size = read_i32(buf)
|
||||
|
||||
if orig_key_type != k_type or orig_type != v_type:
|
||||
for _ in range(size):
|
||||
skip(buf, orig_key_type)
|
||||
skip(buf, orig_type)
|
||||
return {}
|
||||
|
||||
return {c_read_val(buf, k_type, k_spec, decode_response): c_read_val(buf, v_type, v_spec, decode_response)
|
||||
for _ in range(size)}
|
||||
|
||||
elif ttype == T_STRUCT:
|
||||
return read_struct(buf, spec(), decode_response)
|
||||
|
||||
|
||||
cdef c_write_val(CyTransportBase buf, TType ttype, val, spec=None):
|
||||
if ttype == T_BOOL:
|
||||
write_i08(buf, 1 if val else 0)
|
||||
|
||||
elif ttype == T_I08:
|
||||
write_i08(buf, val)
|
||||
|
||||
elif ttype == T_I16:
|
||||
write_i16(buf, val)
|
||||
|
||||
elif ttype == T_I32:
|
||||
write_i32(buf, val)
|
||||
|
||||
elif ttype == T_I64:
|
||||
write_i64(buf, val)
|
||||
|
||||
elif ttype == T_DOUBLE:
|
||||
write_double(buf, val)
|
||||
|
||||
elif ttype == T_STRING:
|
||||
if not isinstance(val, bytes):
|
||||
try:
|
||||
val = val.encode("utf-8")
|
||||
except Exception:
|
||||
pass
|
||||
write_string(buf, val)
|
||||
|
||||
elif ttype == T_SET or ttype == T_LIST:
|
||||
write_list(buf, val, spec)
|
||||
|
||||
elif ttype == T_MAP:
|
||||
write_dict(buf, val, spec)
|
||||
|
||||
elif ttype == T_STRUCT:
|
||||
write_struct(buf, val)
|
||||
|
||||
|
||||
cpdef skip(CyTransportBase buf, TType ttype):
|
||||
cdef TType v_type, k_type, f_type
|
||||
cdef int size
|
||||
|
||||
if ttype == T_BOOL or ttype == T_I08:
|
||||
read_i08(buf)
|
||||
elif ttype == T_I16:
|
||||
read_i16(buf)
|
||||
elif ttype == T_I32:
|
||||
read_i32(buf)
|
||||
elif ttype == T_I64 or ttype == T_DOUBLE:
|
||||
read_i64(buf)
|
||||
elif ttype == T_STRING:
|
||||
size = read_i32(buf)
|
||||
c_read_binary(buf, size)
|
||||
elif ttype == T_SET or ttype == T_LIST:
|
||||
v_type = <TType>read_i08(buf)
|
||||
size = read_i32(buf)
|
||||
for _ in range(size):
|
||||
skip(buf, v_type)
|
||||
elif ttype == T_MAP:
|
||||
k_type = <TType>read_i08(buf)
|
||||
v_type = <TType>read_i08(buf)
|
||||
size = read_i32(buf)
|
||||
for _ in range(size):
|
||||
skip(buf, k_type)
|
||||
skip(buf, v_type)
|
||||
elif ttype == T_STRUCT:
|
||||
while 1:
|
||||
f_type = <TType>read_i08(buf)
|
||||
if f_type == T_STOP:
|
||||
break
|
||||
read_i16(buf)
|
||||
skip(buf, f_type)
|
||||
|
||||
|
||||
def read_val(CyTransportBase buf, TType ttype, decode_response=True):
|
||||
return c_read_val(buf, ttype, None, decode_response)
|
||||
|
||||
|
||||
def write_val(CyTransportBase buf, TType ttype, val, spec=None):
|
||||
c_write_val(buf, ttype, val, spec)
|
||||
|
||||
|
||||
cdef class TCyBinaryProtocol(object):
|
||||
cdef public CyTransportBase trans
|
||||
cdef public bool strict_read
|
||||
cdef public bool strict_write
|
||||
cdef public bool decode_response
|
||||
|
||||
def __init__(self, trans, strict_read=True, strict_write=True,
|
||||
decode_response=True):
|
||||
self.trans = trans
|
||||
self.strict_read = strict_read
|
||||
self.strict_write = strict_write
|
||||
self.decode_response = decode_response
|
||||
|
||||
def skip(self, ttype):
|
||||
skip(self.trans, <TType>(ttype))
|
||||
|
||||
def read_message_begin(self):
|
||||
cdef int32_t size, version, seqid
|
||||
cdef TType ttype
|
||||
|
||||
size = read_i32(self.trans)
|
||||
if size < 0:
|
||||
version = size & VERSION_MASK
|
||||
if version != VERSION_1:
|
||||
raise ProtocolError('invalid version %d' % version)
|
||||
|
||||
name = c_read_val(self.trans, T_STRING)
|
||||
ttype = <TType>(size & TYPE_MASK)
|
||||
else:
|
||||
if self.strict_read:
|
||||
raise ProtocolError('No protocol version header')
|
||||
|
||||
name = c_read_string(self.trans, size)
|
||||
ttype = <TType>(read_i08(self.trans))
|
||||
|
||||
seqid = read_i32(self.trans)
|
||||
|
||||
return name, ttype, seqid
|
||||
|
||||
def read_message_end(self):
|
||||
pass
|
||||
|
||||
def write_message_begin(self, name, TType ttype, int32_t seqid):
|
||||
cdef int32_t version = VERSION_1 | ttype
|
||||
if self.strict_write:
|
||||
write_i32(self.trans, version)
|
||||
c_write_val(self.trans, T_STRING, name)
|
||||
else:
|
||||
c_write_val(self.trans, T_STRING, name)
|
||||
write_i08(self.trans, ttype)
|
||||
|
||||
write_i32(self.trans, seqid)
|
||||
|
||||
def write_message_end(self):
|
||||
self.trans.c_flush()
|
||||
|
||||
def read_struct(self, obj):
|
||||
try:
|
||||
return read_struct(self.trans, obj, self.decode_response)
|
||||
except Exception:
|
||||
self.trans.clean()
|
||||
raise
|
||||
|
||||
def write_struct(self, obj):
|
||||
try:
|
||||
write_struct(self.trans, obj)
|
||||
except Exception:
|
||||
self.trans.clean()
|
||||
raise
|
||||
|
||||
|
||||
class TCyBinaryProtocolFactory(object):
|
||||
def __init__(self, strict_read=True, strict_write=True,
|
||||
decode_response=True):
|
||||
self.strict_read = strict_read
|
||||
self.strict_write = strict_write
|
||||
self.decode_response = decode_response
|
||||
|
||||
def get_protocol(self, trans):
|
||||
return TCyBinaryProtocol(
|
||||
trans, self.strict_read, self.strict_write, self.decode_response)
|
||||
@@ -1,42 +0,0 @@
|
||||
|
||||
#if defined(__APPLE__)
|
||||
|
||||
#include <libkern/OSByteOrder.h>
|
||||
|
||||
#define htobe16(x) OSSwapHostToBigInt16(x)
|
||||
#define htobe32(x) OSSwapHostToBigInt32(x)
|
||||
#define htobe64(x) OSSwapHostToBigInt64(x)
|
||||
#define be16toh(x) OSSwapBigToHostInt16(x)
|
||||
#define be32toh(x) OSSwapBigToHostInt32(x)
|
||||
#define be64toh(x) OSSwapBigToHostInt64(x)
|
||||
|
||||
#else
|
||||
|
||||
#include <endian.h>
|
||||
#include <byteswap.h>
|
||||
|
||||
#ifndef htobe16
|
||||
#define htobe16(x) bswap_16(x)
|
||||
#endif
|
||||
|
||||
#ifndef htobe32
|
||||
#define htobe32(x) bswap_32(x)
|
||||
#endif
|
||||
|
||||
#ifndef htobe64
|
||||
#define htobe64(x) bswap_64(x)
|
||||
#endif
|
||||
|
||||
#ifndef be16toh
|
||||
#define be16toh(x) bswap_16(x)
|
||||
#endif
|
||||
|
||||
#ifndef be32toh
|
||||
#define be32toh(x) bswap_32(x)
|
||||
#endif
|
||||
|
||||
#ifndef be64toh
|
||||
#define be64toh(x) bswap_64(x)
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -1,19 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
from ..thrift import TException
|
||||
|
||||
|
||||
class TProtocolException(TException):
|
||||
"""Custom Protocol Exception class"""
|
||||
|
||||
UNKNOWN = 0
|
||||
INVALID_DATA = 1
|
||||
NEGATIVE_SIZE = 2
|
||||
SIZE_LIMIT = 3
|
||||
BAD_VERSION = 4
|
||||
|
||||
def __init__(self, type=UNKNOWN, message=None):
|
||||
TException.__init__(self, message)
|
||||
self.type = type
|
||||
@@ -1,213 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import struct
|
||||
|
||||
import json
|
||||
from thriftpy.thrift import TType
|
||||
from .exc import TProtocolException
|
||||
|
||||
INTEGER = (TType.BYTE, TType.I16, TType.I32, TType.I64)
|
||||
FLOAT = (TType.DOUBLE,)
|
||||
|
||||
VERSION = 1
|
||||
|
||||
|
||||
def json_value(ttype, val, spec=None):
|
||||
if ttype in INTEGER or ttype in FLOAT or ttype == TType.STRING:
|
||||
return val
|
||||
|
||||
if ttype == TType.BOOL:
|
||||
return True if val else False
|
||||
|
||||
if ttype == TType.STRUCT:
|
||||
return struct_to_json(val)
|
||||
|
||||
if ttype in (TType.SET, TType.LIST):
|
||||
return list_to_json(val, spec)
|
||||
|
||||
if ttype == TType.MAP:
|
||||
return map_to_json(val, spec)
|
||||
|
||||
|
||||
def obj_value(ttype, val, spec=None):
|
||||
if ttype in INTEGER:
|
||||
return int(val)
|
||||
|
||||
if ttype in FLOAT:
|
||||
return float(val)
|
||||
|
||||
if ttype in (TType.STRING, TType.BOOL):
|
||||
return val
|
||||
|
||||
if ttype == TType.STRUCT:
|
||||
return struct_to_obj(val, spec())
|
||||
|
||||
if ttype in (TType.SET, TType.LIST):
|
||||
return list_to_obj(val, spec)
|
||||
|
||||
if ttype == TType.MAP:
|
||||
return map_to_obj(val, spec)
|
||||
|
||||
|
||||
def map_to_obj(val, spec):
|
||||
res = {}
|
||||
if isinstance(spec[0], int):
|
||||
key_type, key_spec = spec[0], None
|
||||
else:
|
||||
key_type, key_spec = spec[0]
|
||||
|
||||
if isinstance(spec[1], int):
|
||||
value_type, value_spec = spec[1], None
|
||||
else:
|
||||
value_type, value_spec = spec[1]
|
||||
|
||||
for v in val:
|
||||
res[obj_value(key_type, v["key"], key_spec)] = obj_value(
|
||||
value_type, v["value"], value_spec)
|
||||
|
||||
return res
|
||||
|
||||
|
||||
def map_to_json(val, spec):
|
||||
res = []
|
||||
if isinstance(spec[0], int):
|
||||
key_type = spec[0]
|
||||
key_spec = None
|
||||
else:
|
||||
key_type, key_spec = spec[0]
|
||||
|
||||
if isinstance(spec[1], int):
|
||||
value_type = spec[1]
|
||||
value_spec = None
|
||||
else:
|
||||
value_type, value_spec = spec[1]
|
||||
|
||||
for k, v in val.items():
|
||||
res.append({"key": json_value(key_type, k, key_spec),
|
||||
"value": json_value(value_type, v, value_spec)})
|
||||
|
||||
return res
|
||||
|
||||
|
||||
def list_to_obj(val, spec):
|
||||
if isinstance(spec, tuple):
|
||||
elem_type, type_spec = spec
|
||||
else:
|
||||
elem_type, type_spec = spec, None
|
||||
|
||||
return [obj_value(elem_type, i, type_spec) for i in val]
|
||||
|
||||
|
||||
def list_to_json(val, spec):
|
||||
if isinstance(spec, tuple):
|
||||
elem_type, type_spec = spec
|
||||
else:
|
||||
elem_type, type_spec = spec, None
|
||||
|
||||
return [json_value(elem_type, i, type_spec) for i in val]
|
||||
|
||||
|
||||
def struct_to_json(val):
|
||||
outobj = {}
|
||||
for fid, field_spec in val.thrift_spec.items():
|
||||
field_type, field_name = field_spec[:2]
|
||||
|
||||
if len(field_spec) <= 3:
|
||||
field_type_spec = None
|
||||
else:
|
||||
field_type_spec = field_spec[2]
|
||||
|
||||
v = getattr(val, field_name)
|
||||
if v is None:
|
||||
continue
|
||||
|
||||
outobj[field_name] = json_value(field_type, v, field_type_spec)
|
||||
|
||||
return outobj
|
||||
|
||||
|
||||
def struct_to_obj(val, obj):
|
||||
for fid, field_spec in obj.thrift_spec.items():
|
||||
field_type, field_name = field_spec[:2]
|
||||
|
||||
if len(field_spec) <= 3:
|
||||
field_type_spec = None
|
||||
else:
|
||||
field_type_spec = field_spec[2]
|
||||
|
||||
if field_name in val:
|
||||
setattr(obj, field_name,
|
||||
obj_value(field_type, val[field_name], field_type_spec))
|
||||
|
||||
return obj
|
||||
|
||||
|
||||
class TJSONProtocol(object):
|
||||
"""A JSON protocol.
|
||||
|
||||
The message in the transport are encoded as this: 4 bytes represents
|
||||
the length of the json object and immediately followed by the json object.
|
||||
|
||||
'\x00\x00\x00+' '{"payload": {}, "metadata": {"version": 1}}'
|
||||
|
||||
the 4 bytes are the bytes representation of an integer and is encoded in
|
||||
big-endian.
|
||||
"""
|
||||
def __init__(self, trans):
|
||||
self.trans = trans
|
||||
self._meta = {"version": VERSION}
|
||||
self._data = None
|
||||
|
||||
def _write_len(self, x):
|
||||
self.trans.write(struct.pack('!I', int(x)))
|
||||
|
||||
def _read_len(self):
|
||||
l = self.trans.read(4)
|
||||
return struct.unpack('!I', l)[0]
|
||||
|
||||
def read_message_begin(self):
|
||||
size = self._read_len()
|
||||
self._data = json.loads(self.trans.read(size).decode("utf-8"))
|
||||
metadata = self._data["metadata"]
|
||||
|
||||
version = int(metadata["version"])
|
||||
if version != VERSION:
|
||||
raise TProtocolException(
|
||||
type=TProtocolException.BAD_VERSION,
|
||||
message="Bad version in read_message_begin:{}".format(version))
|
||||
|
||||
return metadata["name"], metadata["ttype"], metadata["seqid"]
|
||||
|
||||
def read_message_end(self):
|
||||
pass
|
||||
|
||||
def write_message_begin(self, name, ttype, seqid):
|
||||
self._meta.update({"name": name, "ttype": ttype, "seqid": seqid})
|
||||
|
||||
def write_message_end(self):
|
||||
pass
|
||||
|
||||
def read_struct(self, obj):
|
||||
if not self._data:
|
||||
size = self._read_len()
|
||||
self._data = json.loads(self.trans.read(size).decode("utf-8"))
|
||||
|
||||
res = struct_to_obj(self._data["payload"], obj)
|
||||
self._data = None
|
||||
return res
|
||||
|
||||
def write_struct(self, obj):
|
||||
data = json.dumps({
|
||||
"metadata": self._meta,
|
||||
"payload": struct_to_json(obj)
|
||||
})
|
||||
|
||||
self._write_len(len(data))
|
||||
self.trans.write(data.encode("utf-8"))
|
||||
|
||||
|
||||
class TJSONProtocolFactory(object):
|
||||
def get_protocol(self, trans):
|
||||
return TJSONProtocol(trans)
|
||||
@@ -1,34 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from thriftpy.thrift import TMultiplexedProcessor, TMessageType
|
||||
|
||||
|
||||
class TMultiplexedProtocol(object):
|
||||
"""Multiplex the protocol by prepend service name to api for every api call.
|
||||
Can be used together with all original protocols.
|
||||
"""
|
||||
|
||||
def __init__(self, proto, service_name):
|
||||
self.service_name = service_name
|
||||
self._proto = proto
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._proto, name)
|
||||
|
||||
def write_message_begin(self, name, ttype, seqid):
|
||||
if ttype in (TMessageType.CALL, TMessageType.ONEWAY):
|
||||
self._proto.write_message_begin(
|
||||
self.service_name + TMultiplexedProcessor.SEPARATOR + name,
|
||||
ttype, seqid)
|
||||
else:
|
||||
self._proto.write_message_begin(name, ttype, seqid)
|
||||
|
||||
|
||||
class TMultiplexedProtocolFactory(object):
|
||||
def __init__(self, proto_factory, service_name):
|
||||
self._proto_factory = proto_factory
|
||||
self.service_name = service_name
|
||||
|
||||
def get_protocol(self, trans):
|
||||
proto = self._proto_factory.get_protocol(trans)
|
||||
return TMultiplexedProtocol(proto, self.service_name)
|
||||
@@ -1,81 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import contextlib
|
||||
import warnings
|
||||
|
||||
from thriftpy.protocol import TBinaryProtocolFactory
|
||||
from thriftpy.server import TThreadedServer
|
||||
from thriftpy.thrift import TProcessor, TClient
|
||||
from thriftpy.transport import (
|
||||
TBufferedTransportFactory,
|
||||
TServerSocket,
|
||||
TSocket,
|
||||
)
|
||||
|
||||
|
||||
def make_client(service, host="localhost", port=9090, unix_socket=None,
|
||||
proto_factory=TBinaryProtocolFactory(),
|
||||
trans_factory=TBufferedTransportFactory(),
|
||||
timeout=None):
|
||||
if unix_socket:
|
||||
socket = TSocket(unix_socket=unix_socket)
|
||||
elif host and port:
|
||||
socket = TSocket(host, port, socket_timeout=timeout)
|
||||
else:
|
||||
raise ValueError("Either host/port or unix_socket must be provided.")
|
||||
|
||||
transport = trans_factory.get_transport(socket)
|
||||
protocol = proto_factory.get_protocol(transport)
|
||||
transport.open()
|
||||
return TClient(service, protocol)
|
||||
|
||||
|
||||
def make_server(service, handler,
|
||||
host="localhost", port=9090, unix_socket=None,
|
||||
proto_factory=TBinaryProtocolFactory(),
|
||||
trans_factory=TBufferedTransportFactory()):
|
||||
processor = TProcessor(service, handler)
|
||||
if unix_socket:
|
||||
server_socket = TServerSocket(unix_socket=unix_socket)
|
||||
elif host and port:
|
||||
server_socket = TServerSocket(host=host, port=port)
|
||||
else:
|
||||
raise ValueError("Either host/port or unix_socket must be provided.")
|
||||
|
||||
server = TThreadedServer(processor, server_socket,
|
||||
iprot_factory=proto_factory,
|
||||
itrans_factory=trans_factory)
|
||||
return server
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def client_context(service, host="localhost", port=9090, unix_socket=None,
|
||||
proto_factory=TBinaryProtocolFactory(),
|
||||
trans_factory=TBufferedTransportFactory(),
|
||||
timeout=3000, socket_timeout=3000, connect_timeout=None):
|
||||
if timeout:
|
||||
warnings.warn("`timeout` deprecated, use `socket_timeout` and "
|
||||
"`connect_timeout` instead.")
|
||||
socket_timeout = connect_timeout = timeout
|
||||
|
||||
if unix_socket:
|
||||
socket = TSocket(unix_socket=unix_socket,
|
||||
connect_timeout=connect_timeout,
|
||||
socket_timeout=socket_timeout)
|
||||
elif host and port:
|
||||
socket = TSocket(host, port,
|
||||
connect_timeout=connect_timeout,
|
||||
socket_timeout=socket_timeout)
|
||||
else:
|
||||
raise ValueError("Either host/port or unix_socket must be provided.")
|
||||
|
||||
try:
|
||||
transport = trans_factory.get_transport(socket)
|
||||
protocol = proto_factory.get_protocol(transport)
|
||||
transport.open()
|
||||
yield TClient(service, protocol)
|
||||
|
||||
finally:
|
||||
transport.close()
|
||||
@@ -1,104 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import logging
|
||||
import threading
|
||||
|
||||
from thriftpy.protocol import TBinaryProtocolFactory
|
||||
from thriftpy.transport import (
|
||||
TBufferedTransportFactory,
|
||||
TTransportException
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TServer(object):
|
||||
def __init__(self, processor, trans,
|
||||
itrans_factory=None, iprot_factory=None,
|
||||
otrans_factory=None, oprot_factory=None):
|
||||
self.processor = processor
|
||||
self.trans = trans
|
||||
|
||||
self.itrans_factory = itrans_factory or TBufferedTransportFactory()
|
||||
self.iprot_factory = iprot_factory or TBinaryProtocolFactory()
|
||||
self.otrans_factory = otrans_factory or self.itrans_factory
|
||||
self.oprot_factory = oprot_factory or self.iprot_factory
|
||||
|
||||
def serve(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
class TSimpleServer(TServer):
|
||||
"""Simple single-threaded server that just pumps around one transport."""
|
||||
|
||||
def __init__(self, *args):
|
||||
TServer.__init__(self, *args)
|
||||
self.closed = False
|
||||
|
||||
def serve(self):
|
||||
self.trans.listen()
|
||||
while True:
|
||||
client = self.trans.accept()
|
||||
itrans = self.itrans_factory.get_transport(client)
|
||||
otrans = self.otrans_factory.get_transport(client)
|
||||
iprot = self.iprot_factory.get_protocol(itrans)
|
||||
oprot = self.oprot_factory.get_protocol(otrans)
|
||||
try:
|
||||
while not self.closed:
|
||||
self.processor.process(iprot, oprot)
|
||||
except TTransportException:
|
||||
pass
|
||||
except Exception as x:
|
||||
logger.exception(x)
|
||||
|
||||
itrans.close()
|
||||
otrans.close()
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
class TThreadedServer(TServer):
|
||||
"""Threaded server that spawns a new thread per each connection."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.daemon = kwargs.pop("daemon", False)
|
||||
TServer.__init__(self, *args, **kwargs)
|
||||
self.closed = False
|
||||
|
||||
def serve(self):
|
||||
self.trans.listen()
|
||||
while not self.closed:
|
||||
try:
|
||||
client = self.trans.accept()
|
||||
t = threading.Thread(target=self.handle, args=(client,))
|
||||
t.setDaemon(self.daemon)
|
||||
t.start()
|
||||
except KeyboardInterrupt:
|
||||
raise
|
||||
except Exception as x:
|
||||
logger.exception(x)
|
||||
|
||||
def handle(self, client):
|
||||
itrans = self.itrans_factory.get_transport(client)
|
||||
otrans = self.otrans_factory.get_transport(client)
|
||||
iprot = self.iprot_factory.get_protocol(itrans)
|
||||
oprot = self.oprot_factory.get_protocol(otrans)
|
||||
try:
|
||||
while True:
|
||||
self.processor.process(iprot, oprot)
|
||||
except TTransportException:
|
||||
pass
|
||||
except Exception as x:
|
||||
logger.exception(x)
|
||||
|
||||
itrans.close()
|
||||
otrans.close()
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
@@ -1,393 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
thriftpy.thrift
|
||||
~~~~~~~~~~~~~~~~~~
|
||||
|
||||
Thrift simplified.
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import functools
|
||||
|
||||
from ._compat import init_func_generator, with_metaclass
|
||||
|
||||
|
||||
def args2kwargs(thrift_spec, *args):
|
||||
arg_names = [item[1][1] for item in sorted(thrift_spec.items())]
|
||||
return dict(zip(arg_names, args))
|
||||
|
||||
|
||||
def parse_spec(ttype, spec=None):
|
||||
name_map = TType._VALUES_TO_NAMES
|
||||
|
||||
def _type(s):
|
||||
return parse_spec(*s) if isinstance(s, tuple) else name_map[s]
|
||||
|
||||
if spec is None:
|
||||
return name_map[ttype]
|
||||
|
||||
if ttype == TType.STRUCT:
|
||||
return spec.__name__
|
||||
|
||||
if ttype in (TType.LIST, TType.SET):
|
||||
return "%s<%s>" % (name_map[ttype], _type(spec))
|
||||
|
||||
if ttype == TType.MAP:
|
||||
return "MAP<%s, %s>" % (_type(spec[0]), _type(spec[1]))
|
||||
|
||||
|
||||
class TType(object):
|
||||
STOP = 0
|
||||
VOID = 1
|
||||
BOOL = 2
|
||||
BYTE = 3
|
||||
I08 = 3
|
||||
DOUBLE = 4
|
||||
I16 = 6
|
||||
I32 = 8
|
||||
I64 = 10
|
||||
STRING = 11
|
||||
UTF7 = 11
|
||||
BINARY = 11 # This here just for parsing. For all purposes, it's a string
|
||||
STRUCT = 12
|
||||
MAP = 13
|
||||
SET = 14
|
||||
LIST = 15
|
||||
UTF8 = 16
|
||||
UTF16 = 17
|
||||
|
||||
_VALUES_TO_NAMES = {
|
||||
STOP: 'STOP',
|
||||
VOID: 'VOID',
|
||||
BOOL: 'BOOL',
|
||||
BYTE: 'BYTE',
|
||||
I08: 'BYTE',
|
||||
DOUBLE: 'DOUBLE',
|
||||
I16: 'I16',
|
||||
I32: 'I32',
|
||||
I64: 'I64',
|
||||
STRING: 'STRING',
|
||||
UTF7: 'STRING',
|
||||
BINARY: 'STRING',
|
||||
STRUCT: 'STRUCT',
|
||||
MAP: 'MAP',
|
||||
SET: 'SET',
|
||||
LIST: 'LIST',
|
||||
UTF8: 'UTF8',
|
||||
UTF16: 'UTF16'
|
||||
}
|
||||
|
||||
|
||||
class TMessageType(object):
|
||||
CALL = 1
|
||||
REPLY = 2
|
||||
EXCEPTION = 3
|
||||
ONEWAY = 4
|
||||
|
||||
|
||||
class TPayloadMeta(type):
|
||||
|
||||
def __new__(cls, name, bases, attrs):
|
||||
if "default_spec" in attrs:
|
||||
attrs["__init__"] = init_func_generator(attrs.pop("default_spec"))
|
||||
return super(TPayloadMeta, cls).__new__(cls, name, bases, attrs)
|
||||
|
||||
|
||||
def gen_init(cls, thrift_spec=None, default_spec=None):
|
||||
if thrift_spec is not None:
|
||||
cls.thrift_spec = thrift_spec
|
||||
|
||||
if default_spec is not None:
|
||||
cls.__init__ = init_func_generator(default_spec)
|
||||
return cls
|
||||
|
||||
|
||||
class TPayload(with_metaclass(TPayloadMeta, object)):
|
||||
|
||||
__hash__ = None
|
||||
|
||||
def read(self, iprot):
|
||||
iprot.read_struct(self)
|
||||
|
||||
def write(self, oprot):
|
||||
oprot.write_struct(self)
|
||||
|
||||
def __repr__(self):
|
||||
l = ['%s=%r' % (key, value) for key, value in self.__dict__.items()]
|
||||
return '%s(%s)' % (self.__class__.__name__, ', '.join(l))
|
||||
|
||||
def __str__(self):
|
||||
return repr(self)
|
||||
|
||||
def __eq__(self, other):
|
||||
return isinstance(other, self.__class__) and \
|
||||
self.__dict__ == other.__dict__
|
||||
|
||||
def __ne__(self, other):
|
||||
return not self.__eq__(other)
|
||||
|
||||
|
||||
class TClient(object):
|
||||
|
||||
def __init__(self, service, iprot, oprot=None):
|
||||
self._service = service
|
||||
self._iprot = self._oprot = iprot
|
||||
if oprot is not None:
|
||||
self._oprot = oprot
|
||||
self._seqid = 0
|
||||
|
||||
def __getattr__(self, _api):
|
||||
if _api in self._service.thrift_services:
|
||||
return functools.partial(self._req, _api)
|
||||
|
||||
raise AttributeError("{} instance has no attribute '{}'".format(
|
||||
self.__class__.__name__, _api))
|
||||
|
||||
def __dir__(self):
|
||||
return self._service.thrift_services
|
||||
|
||||
def _req(self, _api, *args, **kwargs):
|
||||
_kw = args2kwargs(getattr(self._service, _api + "_args").thrift_spec,
|
||||
*args)
|
||||
kwargs.update(_kw)
|
||||
result_cls = getattr(self._service, _api + "_result")
|
||||
|
||||
self._send(_api, **kwargs)
|
||||
# wait result only if non-oneway
|
||||
if not getattr(result_cls, "oneway"):
|
||||
return self._recv(_api)
|
||||
|
||||
def _send(self, _api, **kwargs):
|
||||
self._oprot.write_message_begin(_api, TMessageType.CALL, self._seqid)
|
||||
args = getattr(self._service, _api + "_args")()
|
||||
for k, v in kwargs.items():
|
||||
setattr(args, k, v)
|
||||
args.write(self._oprot)
|
||||
self._oprot.write_message_end()
|
||||
self._oprot.trans.flush()
|
||||
|
||||
def _recv(self, _api):
|
||||
fname, mtype, rseqid = self._iprot.read_message_begin()
|
||||
if mtype == TMessageType.EXCEPTION:
|
||||
x = TApplicationException()
|
||||
x.read(self._iprot)
|
||||
self._iprot.read_message_end()
|
||||
raise x
|
||||
result = getattr(self._service, _api + "_result")()
|
||||
result.read(self._iprot)
|
||||
self._iprot.read_message_end()
|
||||
|
||||
if hasattr(result, "success") and result.success is not None:
|
||||
return result.success
|
||||
|
||||
# void api without throws
|
||||
if len(result.thrift_spec) == 0:
|
||||
return
|
||||
|
||||
# check throws
|
||||
for k, v in result.__dict__.items():
|
||||
if k != "success" and v:
|
||||
raise v
|
||||
|
||||
# no throws & not void api
|
||||
if hasattr(result, "success"):
|
||||
raise TApplicationException(TApplicationException.MISSING_RESULT)
|
||||
|
||||
def close(self):
|
||||
self._iprot.trans.close()
|
||||
if self._iprot != self._oprot:
|
||||
self._oprot.trans.close()
|
||||
|
||||
|
||||
class TProcessor(object):
|
||||
"""Base class for procsessor, which works on two streams."""
|
||||
|
||||
def __init__(self, service, handler):
|
||||
self._service = service
|
||||
self._handler = handler
|
||||
|
||||
def process_in(self, iprot):
|
||||
api, type, seqid = iprot.read_message_begin()
|
||||
if api not in self._service.thrift_services:
|
||||
iprot.skip(TType.STRUCT)
|
||||
iprot.read_message_end()
|
||||
return api, seqid, TApplicationException(TApplicationException.UNKNOWN_METHOD), None # noqa
|
||||
|
||||
args = getattr(self._service, api + "_args")()
|
||||
args.read(iprot)
|
||||
iprot.read_message_end()
|
||||
result = getattr(self._service, api + "_result")()
|
||||
|
||||
# convert kwargs to args
|
||||
api_args = [args.thrift_spec[k][1] for k in sorted(args.thrift_spec)]
|
||||
|
||||
def call():
|
||||
f = getattr(self._handler, api)
|
||||
return f(*(args.__dict__[k] for k in api_args))
|
||||
|
||||
return api, seqid, result, call
|
||||
|
||||
def send_exception(self, oprot, api, exc, seqid):
|
||||
oprot.write_message_begin(api, TMessageType.EXCEPTION, seqid)
|
||||
exc.write(oprot)
|
||||
oprot.write_message_end()
|
||||
oprot.trans.flush()
|
||||
|
||||
def send_result(self, oprot, api, result, seqid):
|
||||
oprot.write_message_begin(api, TMessageType.REPLY, seqid)
|
||||
result.write(oprot)
|
||||
oprot.write_message_end()
|
||||
oprot.trans.flush()
|
||||
|
||||
def handle_exception(self, e, result):
|
||||
for k in sorted(result.thrift_spec):
|
||||
if result.thrift_spec[k][1] == "success":
|
||||
continue
|
||||
|
||||
_, exc_name, exc_cls, _ = result.thrift_spec[k]
|
||||
if isinstance(e, exc_cls):
|
||||
setattr(result, exc_name, e)
|
||||
break
|
||||
else:
|
||||
raise e
|
||||
|
||||
def process(self, iprot, oprot):
|
||||
api, seqid, result, call = self.process_in(iprot)
|
||||
|
||||
if isinstance(result, TApplicationException):
|
||||
return self.send_exception(oprot, api, result, seqid)
|
||||
|
||||
try:
|
||||
result.success = call()
|
||||
except Exception as e:
|
||||
# raise if api don't have throws
|
||||
self.handle_exception(e, result)
|
||||
|
||||
if not result.oneway:
|
||||
self.send_result(oprot, api, result, seqid)
|
||||
|
||||
|
||||
class TMultiplexedProcessor(TProcessor):
|
||||
SEPARATOR = ":"
|
||||
|
||||
def __init__(self):
|
||||
self.processors = {}
|
||||
|
||||
def register_processor(self, service_name, processor):
|
||||
if service_name in self.processors:
|
||||
raise TApplicationException(
|
||||
type=TApplicationException.INTERNAL_ERROR,
|
||||
message='processor for `{0}` already registered'
|
||||
.format(service_name))
|
||||
self.processors[service_name] = processor
|
||||
|
||||
def process_in(self, iprot):
|
||||
api, type, seqid = iprot.read_message_begin()
|
||||
if type not in (TMessageType.CALL, TMessageType.ONEWAY):
|
||||
raise TException("TMultiplex protocol only supports CALL & ONEWAY")
|
||||
if TMultiplexedProcessor.SEPARATOR not in api:
|
||||
raise TException("Service name not found in message. "
|
||||
"You should use TMultiplexedProtocol in client.")
|
||||
|
||||
service_name, api = api.split(TMultiplexedProcessor.SEPARATOR)
|
||||
if service_name not in self.processors:
|
||||
iprot.skip(TType.STRUCT)
|
||||
iprot.read_message_end()
|
||||
e = TApplicationException(TApplicationException.UNKNOWN_METHOD)
|
||||
return api, seqid, e, None
|
||||
|
||||
proc = self.processors[service_name]
|
||||
args = getattr(proc._service, api + "_args")()
|
||||
args.read(iprot)
|
||||
iprot.read_message_end()
|
||||
result = getattr(proc._service, api + "_result")()
|
||||
|
||||
# convert kwargs to args
|
||||
api_args = [args.thrift_spec[k][1] for k in sorted(args.thrift_spec)]
|
||||
|
||||
def call():
|
||||
f = getattr(proc._handler, api)
|
||||
return f(*(args.__dict__[k] for k in api_args))
|
||||
|
||||
return api, seqid, result, call
|
||||
|
||||
|
||||
class TProcessorFactory(object):
|
||||
|
||||
def __init__(self, processor_class, *args, **kwargs):
|
||||
self.args = args
|
||||
self.kwargs = kwargs
|
||||
|
||||
self.processor_class = processor_class
|
||||
|
||||
def get_processor(self):
|
||||
return self.processor_class(*self.args, **self.kwargs)
|
||||
|
||||
|
||||
class TException(TPayload, Exception):
|
||||
"""Base class for all thrift exceptions."""
|
||||
|
||||
def __hash__(self):
|
||||
return id(self)
|
||||
|
||||
def __eq__(self, other):
|
||||
return id(self) == id(other)
|
||||
|
||||
|
||||
class TDecodeException(TException):
|
||||
def __init__(self, name, fid, field, value, ttype, spec=None):
|
||||
self.struct_name = name
|
||||
self.fid = fid
|
||||
self.field = field
|
||||
self.value = value
|
||||
|
||||
self.type_repr = parse_spec(ttype, spec)
|
||||
|
||||
def __str__(self):
|
||||
return (
|
||||
"Field '%s(%s)' of '%s' needs type '%s', "
|
||||
"but the value is `%r`"
|
||||
) % (self.field, self.fid, self.struct_name, self.type_repr,
|
||||
self.value)
|
||||
|
||||
|
||||
class TApplicationException(TException):
|
||||
"""Application level thrift exceptions."""
|
||||
|
||||
thrift_spec = {
|
||||
1: (TType.STRING, 'message', False),
|
||||
2: (TType.I32, 'type', False),
|
||||
}
|
||||
|
||||
UNKNOWN = 0
|
||||
UNKNOWN_METHOD = 1
|
||||
INVALID_MESSAGE_TYPE = 2
|
||||
WRONG_METHOD_NAME = 3
|
||||
BAD_SEQUENCE_ID = 4
|
||||
MISSING_RESULT = 5
|
||||
INTERNAL_ERROR = 6
|
||||
PROTOCOL_ERROR = 7
|
||||
|
||||
def __init__(self, type=UNKNOWN, message=None):
|
||||
super(TApplicationException, self).__init__()
|
||||
self.type = type
|
||||
self.message = message
|
||||
|
||||
def __str__(self):
|
||||
if self.message:
|
||||
return self.message
|
||||
|
||||
if self.type == self.UNKNOWN_METHOD:
|
||||
return 'Unknown method'
|
||||
elif self.type == self.INVALID_MESSAGE_TYPE:
|
||||
return 'Invalid message type'
|
||||
elif self.type == self.WRONG_METHOD_NAME:
|
||||
return 'Wrong method name'
|
||||
elif self.type == self.BAD_SEQUENCE_ID:
|
||||
return 'Bad sequence ID'
|
||||
elif self.type == self.MISSING_RESULT:
|
||||
return 'Missing result'
|
||||
else:
|
||||
return 'Default (unknown) TApplicationException'
|
||||
@@ -1,231 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
>>> pingpong = thriftpy.load("pingpong.thrift")
|
||||
>>>
|
||||
>>> class Dispatcher(object):
|
||||
>>> def ping(self):
|
||||
>>> return "pong"
|
||||
|
||||
>>> server = make_server(pingpong.PingPong, Dispatcher())
|
||||
>>> server.listen(6000)
|
||||
>>> client = ioloop.IOLoop.current().run_sync(
|
||||
lambda: make_client(pingpong.PingPong, '127.0.0.1', 6000))
|
||||
>>> ioloop.IOLoop.current().run_sync(client.ping)
|
||||
'pong'
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import logging
|
||||
import socket
|
||||
import struct
|
||||
import toro
|
||||
from contextlib import contextmanager
|
||||
from datetime import timedelta
|
||||
from io import BytesIO
|
||||
from tornado import tcpserver, ioloop, iostream, gen
|
||||
|
||||
# TODO need TCyTornadoStreamTransport to work with cython binary protocol
|
||||
from .protocol.binary import TBinaryProtocolFactory
|
||||
from .thrift import TApplicationException, TProcessor, TClient
|
||||
from .transport import TTransportException, TTransportBase
|
||||
from .transport.memory import TMemoryBuffer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TTornadoStreamTransport(TTransportBase):
|
||||
"""a framed, buffered transport over a Tornado stream"""
|
||||
DEFAULT_CONNECT_TIMEOUT = timedelta(seconds=1)
|
||||
DEFAULT_READ_TIMEOUT = timedelta(seconds=1)
|
||||
|
||||
def __init__(self, host, port, stream=None, io_loop=None, ssl_options=None,
|
||||
read_timeout=DEFAULT_READ_TIMEOUT):
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.io_loop = io_loop or ioloop.IOLoop.current()
|
||||
self.read_timeout = read_timeout
|
||||
self.is_queuing_reads = False
|
||||
self.read_queue = []
|
||||
self.__wbuf = BytesIO()
|
||||
self._read_lock = toro.Lock()
|
||||
self.ssl_options = ssl_options
|
||||
|
||||
# servers provide a ready-to-go stream
|
||||
self.stream = stream
|
||||
if self.stream is not None:
|
||||
self._set_close_callback()
|
||||
|
||||
def with_timeout(self, timeout, future):
|
||||
return gen.with_timeout(timeout, future, self.io_loop)
|
||||
|
||||
@gen.coroutine
|
||||
def open(self, timeout=DEFAULT_CONNECT_TIMEOUT):
|
||||
logger.debug('socket connecting')
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM, 0)
|
||||
if self.ssl_options is None:
|
||||
self.stream = iostream.IOStream(sock)
|
||||
else:
|
||||
self.stream = iostream.SSLIOStream(sock, ssl_options=self.ssl_options)
|
||||
|
||||
try:
|
||||
yield self.with_timeout(timeout, self.stream.connect(
|
||||
(self.host, self.port)))
|
||||
except (socket.error, OSError, IOError):
|
||||
message = 'could not connect to {}:{}'.format(self.host, self.port)
|
||||
raise TTransportException(
|
||||
type=TTransportException.NOT_OPEN,
|
||||
message=message)
|
||||
|
||||
self._set_close_callback()
|
||||
raise gen.Return(self)
|
||||
|
||||
def _set_close_callback(self):
|
||||
self.stream.set_close_callback(self.close)
|
||||
|
||||
def close(self):
|
||||
# don't raise if we intend to close
|
||||
self.stream.set_close_callback(None)
|
||||
self.stream.close()
|
||||
|
||||
def read(self, _):
|
||||
# The generated code for Tornado shouldn't do individual reads -- only
|
||||
# frames at a time
|
||||
assert False, "you're doing it wrong"
|
||||
|
||||
@contextmanager
|
||||
def io_exception_context(self):
|
||||
try:
|
||||
yield
|
||||
except (socket.error, OSError, IOError) as e:
|
||||
raise TTransportException(
|
||||
type=TTransportException.END_OF_FILE,
|
||||
message=str(e))
|
||||
except iostream.StreamBufferFullError as e:
|
||||
raise TTransportException(
|
||||
type=TTransportException.UNKNOWN,
|
||||
message=str(e))
|
||||
except gen.TimeoutError as e:
|
||||
raise TTransportException(
|
||||
type=TTransportException.TIMED_OUT,
|
||||
message=str(e))
|
||||
|
||||
@gen.coroutine
|
||||
def read_frame(self):
|
||||
# IOStream processes reads one at a time
|
||||
with (yield self._read_lock.acquire()):
|
||||
with self.io_exception_context():
|
||||
frame_header = yield self._read_bytes(4)
|
||||
if len(frame_header) == 0:
|
||||
raise iostream.StreamClosedError(
|
||||
'Read zero bytes from stream')
|
||||
frame_length, = struct.unpack('!i', frame_header)
|
||||
logger.debug('received frame header, frame length = %d',
|
||||
frame_length)
|
||||
frame = yield self._read_bytes(frame_length)
|
||||
logger.debug('received frame payload: %r', frame)
|
||||
raise gen.Return(frame)
|
||||
|
||||
def _read_bytes(self, n):
|
||||
return self.with_timeout(self.read_timeout, self.stream.read_bytes(n))
|
||||
|
||||
def write(self, buf):
|
||||
self.__wbuf.write(buf)
|
||||
|
||||
def flush(self):
|
||||
frame = self.__wbuf.getvalue()
|
||||
# reset wbuf before write/flush to preserve state on underlying failure
|
||||
frame_length = struct.pack('!i', len(frame))
|
||||
self.__wbuf = BytesIO()
|
||||
with self.io_exception_context():
|
||||
return self.stream.write(frame_length + frame)
|
||||
|
||||
|
||||
class TTornadoServer(tcpserver.TCPServer):
|
||||
def __init__(self, processor, iprot_factory, oprot_factory=None,
|
||||
transport_read_timeout=TTornadoStreamTransport.DEFAULT_READ_TIMEOUT, # noqa
|
||||
*args, **kwargs):
|
||||
super(TTornadoServer, self).__init__(*args, **kwargs)
|
||||
|
||||
self._processor = processor
|
||||
self._iprot_factory = iprot_factory
|
||||
self._oprot_factory = (oprot_factory if oprot_factory is not None
|
||||
else iprot_factory)
|
||||
self.transport_read_timeout = transport_read_timeout
|
||||
|
||||
@gen.coroutine
|
||||
def handle_stream(self, stream, address):
|
||||
host, port = address
|
||||
trans = TTornadoStreamTransport(
|
||||
host=host, port=port, stream=stream,
|
||||
io_loop=self.io_loop, read_timeout=self.transport_read_timeout)
|
||||
try:
|
||||
oprot = self._oprot_factory.get_protocol(trans)
|
||||
iprot = self._iprot_factory.get_protocol(TMemoryBuffer())
|
||||
|
||||
while not trans.stream.closed():
|
||||
# TODO: maybe read multiple frames in advance for concurrency
|
||||
try:
|
||||
frame = yield trans.read_frame()
|
||||
except TTransportException as e:
|
||||
if e.type == TTransportException.END_OF_FILE:
|
||||
break
|
||||
else:
|
||||
raise
|
||||
|
||||
iprot.trans.setvalue(frame)
|
||||
api, seqid, result, call = self._processor.process_in(iprot)
|
||||
if isinstance(result, TApplicationException):
|
||||
self._processor.send_exception(oprot, api, result, seqid)
|
||||
else:
|
||||
try:
|
||||
result.success = yield gen.maybe_future(call())
|
||||
except Exception as e:
|
||||
# raise if api don't have throws
|
||||
self._processor.handle_exception(e, result)
|
||||
|
||||
self._processor.send_result(oprot, api, result, seqid)
|
||||
except Exception:
|
||||
logger.exception('thrift exception in handle_stream')
|
||||
trans.close()
|
||||
|
||||
logger.info('client disconnected %s:%d', host, port)
|
||||
|
||||
|
||||
class TTornadoClient(TClient):
|
||||
@gen.coroutine
|
||||
def _recv(self, api):
|
||||
frame = yield self._oprot.trans.read_frame()
|
||||
self._iprot.trans.setvalue(frame)
|
||||
result = super(TTornadoClient, self)._recv(api)
|
||||
raise gen.Return(result)
|
||||
|
||||
def close(self):
|
||||
self._oprot.trans.close()
|
||||
|
||||
|
||||
def make_server(
|
||||
service, handler, proto_factory=TBinaryProtocolFactory(),
|
||||
io_loop=None, ssl_options=None,
|
||||
transport_read_timeout=TTornadoStreamTransport.DEFAULT_READ_TIMEOUT):
|
||||
processor = TProcessor(service, handler)
|
||||
server = TTornadoServer(processor, iprot_factory=proto_factory,
|
||||
transport_read_timeout=transport_read_timeout,
|
||||
io_loop=io_loop, ssl_options=ssl_options)
|
||||
return server
|
||||
|
||||
|
||||
@gen.coroutine
|
||||
def make_client(
|
||||
service, host, port, proto_factory=TBinaryProtocolFactory(),
|
||||
io_loop=None, ssl_options=None,
|
||||
connect_timeout=TTornadoStreamTransport.DEFAULT_CONNECT_TIMEOUT,
|
||||
read_timeout=TTornadoStreamTransport.DEFAULT_READ_TIMEOUT):
|
||||
transport = TTornadoStreamTransport(host, port, io_loop=io_loop, ssl_options=ssl_options,
|
||||
read_timeout=read_timeout)
|
||||
iprot = proto_factory.get_protocol(TMemoryBuffer())
|
||||
oprot = proto_factory.get_protocol(transport)
|
||||
yield transport.open(connect_timeout)
|
||||
client = TTornadoClient(service, iprot, oprot)
|
||||
raise gen.Return(client)
|
||||
@@ -1,89 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
from thriftpy._compat import CYTHON
|
||||
from ..thrift import TType, TException
|
||||
|
||||
|
||||
def readall(read_fn, sz):
|
||||
buff = b''
|
||||
have = 0
|
||||
while have < sz:
|
||||
chunk = read_fn(sz - have)
|
||||
have += len(chunk)
|
||||
buff += chunk
|
||||
|
||||
if len(chunk) == 0:
|
||||
raise TTransportException(TTransportException.END_OF_FILE,
|
||||
"End of file reading from transport")
|
||||
|
||||
return buff
|
||||
|
||||
|
||||
class TTransportBase(object):
|
||||
"""Base class for Thrift transport layer."""
|
||||
|
||||
def _read(self, sz):
|
||||
raise NotImplementedError
|
||||
|
||||
def read(self, sz):
|
||||
return readall(self._read, sz)
|
||||
|
||||
|
||||
class TTransportException(TException):
|
||||
"""Custom Transport Exception class"""
|
||||
|
||||
thrift_spec = {
|
||||
1: (TType.STRING, 'message'),
|
||||
2: (TType.I32, 'type'),
|
||||
}
|
||||
|
||||
UNKNOWN = 0
|
||||
NOT_OPEN = 1
|
||||
ALREADY_OPEN = 2
|
||||
TIMED_OUT = 3
|
||||
END_OF_FILE = 4
|
||||
|
||||
def __init__(self, type=UNKNOWN, message=None):
|
||||
super(TTransportException, self).__init__()
|
||||
self.type = type
|
||||
self.message = message
|
||||
|
||||
|
||||
# Avoid recursive import
|
||||
from .socket import TSocket, TServerSocket # noqa
|
||||
from .sslsocket import TSSLSocket, TSSLServerSocket # noqa
|
||||
from ._ssl import create_thriftpy_context # noqa
|
||||
from .buffered import TBufferedTransport, TBufferedTransportFactory # noqa
|
||||
from .framed import TFramedTransport, TFramedTransportFactory # noqa
|
||||
from .memory import TMemoryBuffer # noqa
|
||||
|
||||
if CYTHON:
|
||||
from .buffered import TCyBufferedTransport, TCyBufferedTransportFactory
|
||||
from .framed import TCyFramedTransport, TCyFramedTransportFactory
|
||||
from .memory import TCyMemoryBuffer
|
||||
|
||||
# enable cython binary by default for CPython.
|
||||
TMemoryBuffer = TCyMemoryBuffer # noqa
|
||||
TBufferedTransport = TCyBufferedTransport # noqa
|
||||
TBufferedTransportFactory = TCyBufferedTransportFactory # noqa
|
||||
TFramedTransport = TCyFramedTransport # noqa
|
||||
TFramedTransportFactory = TCyFramedTransportFactory # noqa
|
||||
else:
|
||||
# disable cython binary protocol for PYPY since it's slower.
|
||||
TCyMemoryBuffer = TMemoryBuffer
|
||||
TCyBufferedTransport = TBufferedTransport
|
||||
TCyBufferedTransportFactory = TBufferedTransportFactory
|
||||
TCyFramedTransport = TFramedTransport
|
||||
TCyFramedTransportFactory = TFramedTransportFactory
|
||||
|
||||
__all__ = [
|
||||
"TSocket", "TServerSocket",
|
||||
"TSSLSocket", "TSSLServerSocket", "create_thriftpy_context",
|
||||
"TTransportBase", "TTransportException",
|
||||
"TMemoryBuffer", "TFramedTransport", "TFramedTransportFactory",
|
||||
"TBufferedTransport", "TBufferedTransportFactory", "TCyMemoryBuffer",
|
||||
"TCyBufferedTransport", "TCyBufferedTransportFactory",
|
||||
"TCyFramedTransport", "TCyFramedTransportFactory"
|
||||
]
|
||||
@@ -1,174 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
"""
|
||||
The codes in this ssl compat lib were inspired by urllib3.utils.ssl_ module.
|
||||
"""
|
||||
|
||||
import ssl
|
||||
import warnings
|
||||
|
||||
from .._compat import MODERN_SSL
|
||||
|
||||
try:
|
||||
from ssl import (
|
||||
OP_NO_SSLv2, OP_NO_SSLv3, OP_NO_COMPRESSION,
|
||||
OP_CIPHER_SERVER_PREFERENCE, OP_SINGLE_DH_USE, OP_SINGLE_ECDH_USE
|
||||
)
|
||||
except ImportError:
|
||||
OP_NO_SSLv2 = 0x1000000
|
||||
OP_NO_SSLv3 = 0x2000000
|
||||
OP_NO_COMPRESSION = 0x20000
|
||||
OP_CIPHER_SERVER_PREFERENCE = 0x400000
|
||||
OP_SINGLE_DH_USE = 0x100000
|
||||
OP_SINGLE_ECDH_USE = 0x80000
|
||||
|
||||
|
||||
# Disable weak or insecure ciphers by default
|
||||
# (OpenSSL's default setting is 'DEFAULT:!aNULL:!eNULL')
|
||||
# Enable a better set of ciphers by default
|
||||
# This list has been explicitly chosen to:
|
||||
# * Prefer cipher suites that offer perfect forward secrecy (DHE/ECDHE)
|
||||
# * Prefer ECDHE over DHE for better performance
|
||||
# * Prefer any AES-GCM over any AES-CBC for better performance and security
|
||||
# * Then Use HIGH cipher suites as a fallback
|
||||
# * Then Use 3DES as fallback which is secure but slow
|
||||
# * Disable NULL authentication, NULL encryption, and MD5 MACs for security
|
||||
# reasons
|
||||
DEFAULT_CIPHERS = (
|
||||
'ECDH+AESGCM:DH+AESGCM:ECDH+AES256:DH+AES256:ECDH+AES128:DH+AES:ECDH+HIGH:'
|
||||
'DH+HIGH:ECDH+3DES:DH+3DES:RSA+AESGCM:RSA+AES:RSA+HIGH:RSA+3DES:!aNULL:'
|
||||
'!eNULL:!MD5'
|
||||
)
|
||||
|
||||
# Restricted and more secure ciphers for the server side
|
||||
# This list has been explicitly chosen to:
|
||||
# * Prefer cipher suites that offer perfect forward secrecy (DHE/ECDHE)
|
||||
# * Prefer ECDHE over DHE for better performance
|
||||
# * Prefer any AES-GCM over any AES-CBC for better performance and security
|
||||
# * Then Use HIGH cipher suites as a fallback
|
||||
# * Then Use 3DES as fallback which is secure but slow
|
||||
# * Disable NULL authentication, NULL encryption, MD5 MACs, DSS, and RC4 for
|
||||
# security reasons
|
||||
RESTRICTED_SERVER_CIPHERS = (
|
||||
'ECDH+AESGCM:DH+AESGCM:ECDH+AES256:DH+AES256:ECDH+AES128:DH+AES:ECDH+HIGH:'
|
||||
'DH+HIGH:ECDH+3DES:DH+3DES:RSA+AESGCM:RSA+AES:RSA+HIGH:RSA+3DES:!aNULL:'
|
||||
'!eNULL:!MD5:!DSS:!RC4'
|
||||
)
|
||||
|
||||
|
||||
class InsecurePlatformWarning(Warning):
|
||||
"""Warned when certain SSL configuration is not available on a platform.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
try:
|
||||
from ssl import SSLContext
|
||||
except ImportError:
|
||||
import sys
|
||||
|
||||
class SSLContext(object):
|
||||
supports_set_ciphers = ((2, 7) <= sys.version_info < (3,) or
|
||||
(3, 2) <= sys.version_info)
|
||||
|
||||
def __init__(self, protocol_version):
|
||||
self.protocol = protocol_version
|
||||
# Use default values from a real SSLContext
|
||||
self.check_hostname = False
|
||||
self.verify_mode = ssl.CERT_NONE
|
||||
self.ca_certs = None
|
||||
self.options = 0
|
||||
self.certfile = None
|
||||
self.keyfile = None
|
||||
self.ciphers = None
|
||||
|
||||
def load_cert_chain(self, certfile=None, keyfile=None):
|
||||
self.certfile = certfile
|
||||
self.keyfile = keyfile
|
||||
|
||||
def load_verify_locations(self, cafile=None, capath=None):
|
||||
if capath is not None:
|
||||
raise OSError("CA directories not supported in older Pythons")
|
||||
self.ca_certs = cafile
|
||||
|
||||
def set_ciphers(self, cipher_suite):
|
||||
if not self.supports_set_ciphers:
|
||||
raise TypeError(
|
||||
"Your version of Python does not support setting "
|
||||
"a custom cipher suite. Please upgrade to Python "
|
||||
"2.7, 3.2, or later if you need this functionality."
|
||||
)
|
||||
self.ciphers = cipher_suite
|
||||
|
||||
def wrap_socket(self, socket, server_hostname=None, server_side=False):
|
||||
warnings.warn(
|
||||
"A true SSLContext object is not available. This prevents "
|
||||
"urllib3 from configuring SSL appropriately and may cause "
|
||||
"certain SSL connections to fail.",
|
||||
InsecurePlatformWarning
|
||||
)
|
||||
kwargs = {
|
||||
"keyfile": self.keyfile,
|
||||
"certfile": self.certfile,
|
||||
"ca_certs": self.ca_certs,
|
||||
"cert_reqs": self.verify_mode,
|
||||
"ssl_version": self.protocol,
|
||||
"server_side": server_side,
|
||||
}
|
||||
|
||||
if self.supports_set_ciphers:
|
||||
# Platform-specific: Python 2.7+
|
||||
return ssl.wrap_socket(socket, ciphers=self.ciphers, **kwargs)
|
||||
else:
|
||||
# Platform-specific: Python 2.6
|
||||
return ssl.wrap_socket(socket, **kwargs)
|
||||
|
||||
|
||||
def create_thriftpy_context(server_side=False, ciphers=None):
|
||||
"""Backport create_default_context for older python versions.
|
||||
|
||||
The SSLContext has some default security options, you can disable them
|
||||
manually, for example::
|
||||
|
||||
from thriftpy.transport import _ssl
|
||||
context = _ssl.create_thriftpy_context()
|
||||
context.options &= ~_ssl.OP_NO_SSLv3
|
||||
|
||||
You can do the same to enable compression.
|
||||
"""
|
||||
if MODERN_SSL:
|
||||
if server_side:
|
||||
context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
|
||||
else:
|
||||
context = ssl.create_default_context(ssl.Purpose.SERVER_AUTH)
|
||||
|
||||
if ciphers:
|
||||
context.set_ciphers(ciphers)
|
||||
|
||||
else:
|
||||
context = SSLContext(ssl.PROTOCOL_SSLv23)
|
||||
context.options |= OP_NO_SSLv2
|
||||
context.options |= OP_NO_SSLv3
|
||||
context.options |= OP_NO_COMPRESSION
|
||||
|
||||
# server/client default options
|
||||
if server_side:
|
||||
context.options |= OP_CIPHER_SERVER_PREFERENCE
|
||||
context.options |= OP_SINGLE_DH_USE
|
||||
context.options |= OP_SINGLE_ECDH_USE
|
||||
else:
|
||||
context.verify_mode = ssl.CERT_REQUIRED
|
||||
# context.check_hostname = True
|
||||
warnings.warn(
|
||||
"ssl check hostname support disabled, upgrade your python",
|
||||
InsecurePlatformWarning)
|
||||
|
||||
# Platform-specific: Python 2.6
|
||||
if getattr(context, 'supports_set_ciphers', True):
|
||||
if ciphers:
|
||||
context.set_ciphers(ciphers)
|
||||
else:
|
||||
warnings.warn("ssl ciphers support disabled, upgrade your python",
|
||||
InsecurePlatformWarning)
|
||||
|
||||
return context
|
||||
@@ -1,62 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
from io import BytesIO
|
||||
|
||||
from thriftpy._compat import CYTHON
|
||||
from .. import TTransportBase
|
||||
|
||||
|
||||
class TBufferedTransport(TTransportBase):
|
||||
"""Class that wraps another transport and buffers its I/O.
|
||||
|
||||
The implementation uses a (configurable) fixed-size read buffer
|
||||
but buffers all writes until a flush is performed.
|
||||
"""
|
||||
DEFAULT_BUFFER = 4096
|
||||
|
||||
def __init__(self, trans, buf_size=DEFAULT_BUFFER):
|
||||
self._trans = trans
|
||||
self._wbuf = BytesIO()
|
||||
self._rbuf = BytesIO(b"")
|
||||
self._buf_size = buf_size
|
||||
|
||||
def is_open(self):
|
||||
return self._trans.is_open()
|
||||
|
||||
def open(self):
|
||||
return self._trans.open()
|
||||
|
||||
def close(self):
|
||||
return self._trans.close()
|
||||
|
||||
def _read(self, sz):
|
||||
ret = self._rbuf.read(sz)
|
||||
if len(ret) != 0:
|
||||
return ret
|
||||
|
||||
self._rbuf = BytesIO(self._trans.read(max(sz, self._buf_size)))
|
||||
return self._rbuf.read(sz)
|
||||
|
||||
def write(self, buf):
|
||||
self._wbuf.write(buf)
|
||||
|
||||
def flush(self):
|
||||
out = self._wbuf.getvalue()
|
||||
# reset wbuf before write/flush to preserve state on underlying failure
|
||||
self._wbuf = BytesIO()
|
||||
self._trans.write(out)
|
||||
self._trans.flush()
|
||||
|
||||
def getvalue(self):
|
||||
return self._trans.getvalue()
|
||||
|
||||
|
||||
class TBufferedTransportFactory(object):
|
||||
def get_transport(self, trans):
|
||||
return TBufferedTransport(trans)
|
||||
|
||||
|
||||
if CYTHON:
|
||||
from .cybuffered import TCyBufferedTransport, TCyBufferedTransportFactory # noqa
|
||||
@@ -1,90 +0,0 @@
|
||||
from thriftpy.transport.cybase cimport (
|
||||
TCyBuffer,
|
||||
CyTransportBase,
|
||||
DEFAULT_BUFFER
|
||||
)
|
||||
|
||||
from .. import TTransportException
|
||||
|
||||
DEF MIN_BUFFER_SIZE = 1024
|
||||
|
||||
|
||||
cdef class TCyBufferedTransport(CyTransportBase):
|
||||
"""binary reader/writer"""
|
||||
|
||||
cdef:
|
||||
TCyBuffer rbuf, wbuf
|
||||
|
||||
def __init__(self, trans, int buf_size=DEFAULT_BUFFER):
|
||||
if buf_size < MIN_BUFFER_SIZE:
|
||||
raise Exception("buffer too small")
|
||||
|
||||
self.trans = trans
|
||||
self.rbuf = TCyBuffer(buf_size)
|
||||
self.wbuf = TCyBuffer(buf_size)
|
||||
|
||||
def clean(self):
|
||||
self.rbuf.clean()
|
||||
self.wbuf.clean()
|
||||
|
||||
def is_open(self):
|
||||
return self.trans.is_open()
|
||||
|
||||
def open(self):
|
||||
return self.trans.open()
|
||||
|
||||
def close(self):
|
||||
return self.trans.close()
|
||||
|
||||
def write(self, bytes data):
|
||||
cdef int sz = len(data)
|
||||
return self.c_write(data, sz)
|
||||
|
||||
def read(self, int sz):
|
||||
return self.get_string(sz)
|
||||
|
||||
def flush(self):
|
||||
return self.c_flush()
|
||||
|
||||
cdef c_write(self, const char *data, int sz):
|
||||
cdef:
|
||||
int cap = self.wbuf.buf_size - self.wbuf.data_size
|
||||
int r
|
||||
|
||||
if cap < sz:
|
||||
self.c_flush()
|
||||
|
||||
r = self.wbuf.write(sz, data)
|
||||
if r == -1:
|
||||
raise MemoryError("Write to buffer error")
|
||||
|
||||
cdef c_read(self, int sz, char* out):
|
||||
if sz <= 0:
|
||||
return 0
|
||||
|
||||
self.read_trans(sz, out)
|
||||
return sz
|
||||
|
||||
cdef read_trans(self, int sz, char *out):
|
||||
cdef int i = self.rbuf.read_trans(self.trans, sz, out)
|
||||
if i == -1:
|
||||
raise TTransportException(TTransportException.END_OF_FILE,
|
||||
"End of file reading from transport")
|
||||
elif i == -2:
|
||||
raise MemoryError("grow read buffer fail")
|
||||
|
||||
cdef c_flush(self):
|
||||
cdef bytes data
|
||||
if self.wbuf.data_size > 0:
|
||||
data = self.wbuf.buf[:self.wbuf.data_size]
|
||||
self.trans.write(data)
|
||||
self.trans.flush()
|
||||
self.wbuf.clean()
|
||||
|
||||
def getvalue(self):
|
||||
return self.trans.getvalue()
|
||||
|
||||
|
||||
class TCyBufferedTransportFactory(object):
|
||||
def get_transport(self, trans):
|
||||
return TCyBufferedTransport(trans)
|
||||
@@ -1,24 +0,0 @@
|
||||
cdef enum:
|
||||
DEFAULT_BUFFER = 4096
|
||||
STACK_STRING_LEN = 4096
|
||||
|
||||
cdef class TCyBuffer(object):
|
||||
cdef:
|
||||
char *buf
|
||||
int cur, buf_size, data_size
|
||||
|
||||
void move_to_start(self)
|
||||
void clean(self)
|
||||
int write(self, int sz, const char *value)
|
||||
int grow(self, int min_size)
|
||||
read_trans(self, trans, int sz, char *out)
|
||||
|
||||
|
||||
cdef class CyTransportBase(object):
|
||||
cdef object trans
|
||||
|
||||
cdef c_read(self, int sz, char* out)
|
||||
cdef c_write(self, char* data, int sz)
|
||||
cdef c_flush(self)
|
||||
|
||||
cdef get_string(self, int sz)
|
||||
@@ -1,138 +0,0 @@
|
||||
from libc.stdlib cimport malloc, free
|
||||
from libc.string cimport memcpy, memmove
|
||||
|
||||
|
||||
cdef class TCyBuffer(object):
|
||||
def __cinit__(self, buf_size):
|
||||
self.buf = <char*>malloc(buf_size)
|
||||
self.buf_size = buf_size
|
||||
self.cur = 0
|
||||
self.data_size = 0
|
||||
|
||||
def __dealloc__(self):
|
||||
if self.buf != NULL:
|
||||
free(self.buf)
|
||||
self.buf = NULL
|
||||
|
||||
cdef void move_to_start(self):
|
||||
memmove(self.buf, self.buf + self.cur, self.data_size)
|
||||
self.cur = 0
|
||||
|
||||
cdef void clean(self):
|
||||
self.cur = 0
|
||||
self.data_size = 0
|
||||
|
||||
cdef int write(self, int sz, const char *value):
|
||||
cdef:
|
||||
int cap = self.buf_size - self.data_size
|
||||
int remain = cap - self.cur
|
||||
|
||||
if sz <= 0:
|
||||
return 0
|
||||
|
||||
if remain < sz:
|
||||
self.move_to_start()
|
||||
|
||||
# recompute remain spaces
|
||||
remain = cap - self.cur
|
||||
|
||||
if remain < sz:
|
||||
if self.grow(sz - remain + self.buf_size) != 0:
|
||||
return -1
|
||||
|
||||
memcpy(self.buf + self.cur + self.data_size, value, sz)
|
||||
self.data_size += sz
|
||||
|
||||
return sz
|
||||
|
||||
cdef read_trans(self, trans, int sz, char *out):
|
||||
cdef int cap, new_data_len
|
||||
|
||||
if sz <= 0:
|
||||
return 0
|
||||
|
||||
if self.data_size < sz:
|
||||
if self.buf_size < sz:
|
||||
if self.grow(sz) != 0:
|
||||
return -2 # grow buffer error
|
||||
|
||||
cap = self.buf_size - self.data_size
|
||||
|
||||
new_data = trans.read(cap)
|
||||
new_data_len = len(new_data)
|
||||
|
||||
while new_data_len + self.data_size < sz:
|
||||
more = trans.read(cap - new_data_len)
|
||||
more_len = len(more)
|
||||
if more_len <= 0:
|
||||
return -1 # end of file error
|
||||
|
||||
new_data += more
|
||||
new_data_len += more_len
|
||||
|
||||
if cap - self.cur < new_data_len:
|
||||
self.move_to_start()
|
||||
|
||||
memcpy(self.buf + self.cur + self.data_size, <char*>new_data,
|
||||
new_data_len)
|
||||
self.data_size += new_data_len
|
||||
|
||||
memcpy(out, self.buf + self.cur, sz)
|
||||
self.cur += sz
|
||||
self.data_size -= sz
|
||||
|
||||
return sz
|
||||
|
||||
cdef int grow(self, int min_size):
|
||||
if min_size <= self.buf_size:
|
||||
return 0
|
||||
|
||||
cdef int multiples = min_size / self.buf_size
|
||||
if min_size % self.buf_size != 0:
|
||||
multiples += 1
|
||||
|
||||
cdef int new_size = self.buf_size * multiples
|
||||
cdef char *new_buf = <char*>malloc(new_size)
|
||||
if new_buf == NULL:
|
||||
return -1
|
||||
memcpy(new_buf + self.cur, self.buf + self.cur, self.data_size)
|
||||
free(self.buf)
|
||||
self.buf_size = new_size
|
||||
self.buf = new_buf
|
||||
return 0
|
||||
|
||||
|
||||
cdef class CyTransportBase(object):
|
||||
cdef c_read(self, int sz, char* out):
|
||||
pass
|
||||
|
||||
cdef c_write(self, char* data, int sz):
|
||||
pass
|
||||
|
||||
cdef c_flush(self):
|
||||
pass
|
||||
|
||||
def clean(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def sock(self):
|
||||
if not self.trans:
|
||||
return
|
||||
return getattr(self.trans, 'sock', None)
|
||||
|
||||
cdef get_string(self, int sz):
|
||||
cdef:
|
||||
char out[STACK_STRING_LEN]
|
||||
char *dy_out
|
||||
|
||||
if sz > STACK_STRING_LEN:
|
||||
dy_out = <char*>malloc(sz)
|
||||
try:
|
||||
size = self.c_read(sz, dy_out)
|
||||
return dy_out[:size]
|
||||
finally:
|
||||
free(dy_out)
|
||||
else:
|
||||
size = self.c_read(sz, out)
|
||||
return out[:size]
|
||||
@@ -1,74 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import struct
|
||||
from io import BytesIO
|
||||
|
||||
from thriftpy._compat import CYTHON
|
||||
from .. import TTransportBase, readall
|
||||
from ..buffered import TBufferedTransport
|
||||
|
||||
|
||||
class TFramedTransport(TTransportBase):
|
||||
"""Class that wraps another transport and frames its I/O when writing."""
|
||||
def __init__(self, trans):
|
||||
self._trans = trans
|
||||
self._rbuf = BytesIO()
|
||||
self._wbuf = BytesIO()
|
||||
|
||||
def is_open(self):
|
||||
return self._trans.is_open()
|
||||
|
||||
def open(self):
|
||||
return self._trans.open()
|
||||
|
||||
def close(self):
|
||||
return self._trans.close()
|
||||
|
||||
def read(self, sz):
|
||||
# Important: don't attempt to read the next frame if the caller
|
||||
# doesn't actually need any data.
|
||||
if sz == 0:
|
||||
return b''
|
||||
|
||||
ret = self._rbuf.read(sz)
|
||||
if len(ret) != 0:
|
||||
return ret
|
||||
|
||||
self.read_frame()
|
||||
return self._rbuf.read(sz)
|
||||
|
||||
def read_frame(self):
|
||||
buff = readall(self._trans.read, 4)
|
||||
sz, = struct.unpack('!i', buff)
|
||||
frame = readall(self._trans.read, sz)
|
||||
self._rbuf = BytesIO(frame)
|
||||
|
||||
def write(self, buf):
|
||||
self._wbuf.write(buf)
|
||||
|
||||
def flush(self):
|
||||
# reset wbuf before write/flush to preserve state on underlying failure
|
||||
out = self._wbuf.getvalue()
|
||||
self._wbuf = BytesIO()
|
||||
|
||||
# N.B.: Doing this string concatenation is WAY cheaper than making
|
||||
# two separate calls to the underlying socket object. Socket writes in
|
||||
# Python turn out to be REALLY expensive, but it seems to do a pretty
|
||||
# good job of managing string buffer operations without excessive
|
||||
# copies
|
||||
self._trans.write(struct.pack("!i", len(out)) + out)
|
||||
self._trans.flush()
|
||||
|
||||
def getvalue(self):
|
||||
return self._trans.getvalue()
|
||||
|
||||
|
||||
class TFramedTransportFactory(object):
|
||||
def get_transport(self, trans):
|
||||
return TBufferedTransport(TFramedTransport(trans))
|
||||
|
||||
|
||||
if CYTHON:
|
||||
from .cyframed import TCyFramedTransport, TCyFramedTransportFactory # noqa
|
||||
@@ -1,133 +0,0 @@
|
||||
from libc.stdint cimport
|
||||
|
||||
int32_t
|
||||
from libc.stdlib cimport
|
||||
|
||||
malloc, free
|
||||
from libc.string cimport
|
||||
|
||||
memcpy
|
||||
|
||||
from thriftpy.transport.cybase cimport
|
||||
|
||||
(
|
||||
TCyBuffer,
|
||||
CyTransportBase,
|
||||
DEFAULT_BUFFER,
|
||||
STACK_STRING_LEN
|
||||
)
|
||||
from .. import TTransportException
|
||||
|
||||
|
||||
cdef extern from "../../protocol/cybin/endian_port.h":
|
||||
int32_t be32toh(int32_t n)
|
||||
int32_t htobe32(int32_t n)
|
||||
|
||||
|
||||
cdef class TCyFramedTransport(CyTransportBase):
|
||||
cdef:
|
||||
TCyBuffer rbuf, rframe_buf, wframe_buf
|
||||
|
||||
def __init__(self, trans, int buf_size=DEFAULT_BUFFER):
|
||||
self.trans = trans
|
||||
self.rbuf = TCyBuffer(buf_size)
|
||||
self.rframe_buf = TCyBuffer(buf_size)
|
||||
self.wframe_buf = TCyBuffer(buf_size)
|
||||
|
||||
cdef read_trans(self, int sz, char *out):
|
||||
cdef int i = self.rbuf.read_trans(self.trans, sz, out)
|
||||
if i == -1:
|
||||
raise TTransportException(TTransportException.END_OF_FILE,
|
||||
"End of file reading from transport")
|
||||
elif i == -2:
|
||||
raise MemoryError("grow buffer fail")
|
||||
|
||||
cdef write_rframe_buffer(self, const char *data, int sz):
|
||||
cdef int r = self.rframe_buf.write(sz, data)
|
||||
if r == -1:
|
||||
raise MemoryError("Write to buffer error")
|
||||
|
||||
cdef c_read(self, int sz, char *out):
|
||||
if sz <= 0:
|
||||
return 0
|
||||
|
||||
while self.rframe_buf.data_size < sz:
|
||||
self.read_frame()
|
||||
|
||||
memcpy(out, self.rframe_buf.buf + self.rframe_buf.cur, sz)
|
||||
self.rframe_buf.cur += sz
|
||||
self.rframe_buf.data_size -= sz
|
||||
|
||||
return sz
|
||||
|
||||
cdef c_write(self, const char *data, int sz):
|
||||
cdef int r = self.wframe_buf.write(sz, data)
|
||||
if r == -1:
|
||||
raise MemoryError("Write to buffer error")
|
||||
|
||||
cdef read_frame(self):
|
||||
cdef:
|
||||
char frame_len[4]
|
||||
char stack_frame[STACK_STRING_LEN]
|
||||
char *dy_frame
|
||||
int32_t frame_size
|
||||
|
||||
self.read_trans(4, frame_len)
|
||||
frame_size = be32toh((<int32_t*>frame_len)[0])
|
||||
|
||||
if frame_size <= 0:
|
||||
raise TTransportException("No frame.", TTransportException.UNKNOWN)
|
||||
|
||||
if frame_size <= STACK_STRING_LEN:
|
||||
self.read_trans(frame_size, stack_frame)
|
||||
self.write_rframe_buffer(stack_frame, frame_size)
|
||||
else:
|
||||
dy_frame = <char*>malloc(frame_size)
|
||||
try:
|
||||
self.read_trans(frame_size, dy_frame)
|
||||
self.write_rframe_buffer(dy_frame, frame_size)
|
||||
finally:
|
||||
free(dy_frame)
|
||||
|
||||
cdef c_flush(self):
|
||||
cdef:
|
||||
bytes data
|
||||
char *size_str
|
||||
|
||||
if self.wframe_buf.data_size > 0:
|
||||
data = self.wframe_buf.buf[:self.wframe_buf.data_size]
|
||||
size = htobe32(self.wframe_buf.data_size)
|
||||
size_str = <char*>(&size)
|
||||
|
||||
self.trans.write(size_str[:4] + data)
|
||||
self.trans.flush()
|
||||
self.wframe_buf.clean()
|
||||
|
||||
def read(self, int sz):
|
||||
return self.get_string(sz)
|
||||
|
||||
def write(self, bytes data):
|
||||
cdef int sz = len(data)
|
||||
self.c_write(data, sz)
|
||||
|
||||
def flush(self):
|
||||
self.c_flush()
|
||||
|
||||
def is_open(self):
|
||||
return self.trans.is_open()
|
||||
|
||||
def open(self):
|
||||
return self.trans.open()
|
||||
|
||||
def close(self):
|
||||
return self.trans.close()
|
||||
|
||||
def clean(self):
|
||||
self.rbuf.clean()
|
||||
self.rframe_buf.clean()
|
||||
self.wframe_buf.clean()
|
||||
|
||||
|
||||
class TCyFramedTransportFactory(object):
|
||||
def get_transport(self, trans):
|
||||
return TCyFramedTransport(trans)
|
||||
@@ -1,57 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
from io import BytesIO
|
||||
|
||||
from thriftpy._compat import CYTHON
|
||||
from .. import TTransportBase
|
||||
|
||||
|
||||
class TMemoryBuffer(TTransportBase):
|
||||
"""Wraps a BytesIO object as a TTransport."""
|
||||
|
||||
def __init__(self, value=None):
|
||||
"""value -- a value as the initial value in the BytesIO object.
|
||||
|
||||
If value is set, the transport can be read first.
|
||||
"""
|
||||
self._buffer = BytesIO(value) if value is not None else BytesIO()
|
||||
self._pos = 0
|
||||
|
||||
def is_open(self):
|
||||
return not self._buffer.closed
|
||||
|
||||
def open(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
self._buffer.close()
|
||||
|
||||
def read(self, sz):
|
||||
return self._read(sz)
|
||||
|
||||
def _read(self, sz):
|
||||
orig_pos = self._buffer.tell()
|
||||
self._buffer.seek(self._pos)
|
||||
res = self._buffer.read(sz)
|
||||
self._buffer.seek(orig_pos)
|
||||
self._pos += len(res)
|
||||
return res
|
||||
|
||||
def write(self, buf):
|
||||
self._buffer.write(buf)
|
||||
|
||||
def flush(self):
|
||||
pass
|
||||
|
||||
def getvalue(self):
|
||||
return self._buffer.getvalue()
|
||||
|
||||
def setvalue(self, value):
|
||||
self._buffer = BytesIO(value)
|
||||
self._pos = 0
|
||||
|
||||
|
||||
if CYTHON:
|
||||
from .cymemory import TCyMemoryBuffer # noqa
|
||||
@@ -1,98 +0,0 @@
|
||||
from libc.stdlib cimport
|
||||
|
||||
malloc, free
|
||||
from libc.string cimport
|
||||
|
||||
memcpy
|
||||
|
||||
from thriftpy.transport.cybase cimport
|
||||
|
||||
(
|
||||
TCyBuffer,
|
||||
CyTransportBase,
|
||||
DEFAULT_BUFFER,
|
||||
)
|
||||
|
||||
def to_bytes(s):
|
||||
try:
|
||||
return s.encode("utf-8")
|
||||
except Exception:
|
||||
return s
|
||||
|
||||
|
||||
cdef class TCyMemoryBuffer(CyTransportBase):
|
||||
cdef TCyBuffer buf
|
||||
|
||||
def __init__(self, value=b'', int buf_size=DEFAULT_BUFFER):
|
||||
self.trans = None
|
||||
self.buf = TCyBuffer(buf_size)
|
||||
|
||||
if value:
|
||||
self.setvalue(value)
|
||||
|
||||
cdef c_read(self, int sz, char* out):
|
||||
if self.buf.data_size < sz:
|
||||
sz = self.buf.data_size
|
||||
|
||||
if sz <= 0:
|
||||
out[0] = '\0'
|
||||
else:
|
||||
memcpy(out, self.buf.buf + self.buf.cur, sz)
|
||||
self.buf.cur += sz
|
||||
self.buf.data_size -= sz
|
||||
|
||||
return sz
|
||||
|
||||
cdef c_write(self, const char* data, int sz):
|
||||
cdef int r = self.buf.write(sz, data)
|
||||
if r == -1:
|
||||
raise MemoryError("Write to memory error")
|
||||
|
||||
cdef _getvalue(self):
|
||||
cdef char *out
|
||||
cdef int size = self.buf.data_size
|
||||
|
||||
if size <= 0:
|
||||
return b''
|
||||
|
||||
out = <char*>malloc(size)
|
||||
try:
|
||||
memcpy(out, self.buf.buf + self.buf.cur, size)
|
||||
return out[:size]
|
||||
finally:
|
||||
free(out)
|
||||
|
||||
cdef _setvalue(self, int sz, const char *value):
|
||||
self.buf.clean()
|
||||
self.buf.write(sz, value)
|
||||
|
||||
def read(self, sz):
|
||||
return self.get_string(sz)
|
||||
|
||||
def write(self, data):
|
||||
data = to_bytes(data)
|
||||
|
||||
cdef int sz = len(data)
|
||||
return self.c_write(data, sz)
|
||||
|
||||
def is_open(self):
|
||||
return True
|
||||
|
||||
def open(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
def flush(self):
|
||||
pass
|
||||
|
||||
def clean(self):
|
||||
self.buf.clean()
|
||||
|
||||
def getvalue(self):
|
||||
return self._getvalue()
|
||||
|
||||
def setvalue(self, value):
|
||||
value = to_bytes(value)
|
||||
self._setvalue(len(value), value)
|
||||
@@ -1,217 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import, division
|
||||
|
||||
import errno
|
||||
import os
|
||||
import socket
|
||||
import struct
|
||||
import sys
|
||||
|
||||
from . import TTransportException
|
||||
|
||||
|
||||
class TSocket(object):
|
||||
"""Socket implementation for client side."""
|
||||
|
||||
def __init__(self, host=None, port=None, unix_socket=None,
|
||||
sock=None, socket_family=socket.AF_INET,
|
||||
socket_timeout=3000, connect_timeout=None):
|
||||
"""Initialize a TSocket
|
||||
|
||||
TSocket can be initialized in 3 ways:
|
||||
* host + port. can configure to use AF_INET/AF_INET6
|
||||
* unix_socket
|
||||
* socket. should pass already opened socket here.
|
||||
|
||||
@param host(str) The host to connect to.
|
||||
@param port(int) The (TCP) port to connect to.
|
||||
@param unix_socket(str) The filename of a unix socket to connect to.
|
||||
@param sock(socket) Initialize with opened socket directly.
|
||||
If this param used, the host, port and unix_socket params will
|
||||
be ignored.
|
||||
@param socket_family(str) socket.AF_INET or socket.AF_INET6. only
|
||||
take effect when using host/port
|
||||
@param socket_timeout socket timeout in ms
|
||||
@param connect_timeout connect timeout in ms, only used in
|
||||
connection, will be set to socket_timeout if not set.
|
||||
"""
|
||||
if sock:
|
||||
self.sock = sock
|
||||
elif unix_socket:
|
||||
self.unix_socket = unix_socket
|
||||
self.host = None
|
||||
self.port = None
|
||||
self.sock = None
|
||||
else:
|
||||
self.unix_socket = None
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.sock = None
|
||||
|
||||
self.socket_family = socket_family
|
||||
self.socket_timeout = socket_timeout / 1000 if socket_timeout else None
|
||||
self.connect_timeout = connect_timeout / 1000 if connect_timeout \
|
||||
else self.socket_timeout
|
||||
|
||||
def _init_sock(self):
|
||||
if self.unix_socket:
|
||||
_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
else:
|
||||
_sock = socket.socket(self.socket_family, socket.SOCK_STREAM)
|
||||
_sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
|
||||
|
||||
# socket options
|
||||
linger = struct.pack('ii', 0, 0)
|
||||
_sock.setsockopt(socket.SOL_SOCKET, socket.SO_LINGER, linger)
|
||||
_sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
|
||||
|
||||
self.sock = _sock
|
||||
|
||||
def set_handle(self, sock):
|
||||
self.sock = sock
|
||||
|
||||
def set_timeout(self, ms):
|
||||
"""Backward compat api, will bind the timeout to both connect_timeout
|
||||
and socket_timeout.
|
||||
"""
|
||||
self.socket_timeout = ms / 1000 if (ms and ms > 0) else None
|
||||
self.connect_timeout = self.socket_timeout
|
||||
|
||||
if self.sock is not None:
|
||||
self.sock.settimeout(self.socket_timeout)
|
||||
|
||||
def is_open(self):
|
||||
return bool(self.sock)
|
||||
|
||||
def open(self):
|
||||
self._init_sock()
|
||||
|
||||
addr = self.unix_socket or (self.host, self.port)
|
||||
|
||||
try:
|
||||
if self.connect_timeout:
|
||||
self.sock.settimeout(self.connect_timeout)
|
||||
|
||||
self.sock.connect(addr)
|
||||
|
||||
if self.socket_timeout:
|
||||
self.sock.settimeout(self.socket_timeout)
|
||||
|
||||
except (socket.error, OSError):
|
||||
raise TTransportException(
|
||||
type=TTransportException.NOT_OPEN,
|
||||
message="Could not connect to %s" % str(addr))
|
||||
|
||||
def read(self, sz):
|
||||
try:
|
||||
buff = self.sock.recv(sz)
|
||||
except socket.error as e:
|
||||
if (e.args[0] == errno.ECONNRESET and
|
||||
(sys.platform == 'darwin' or
|
||||
sys.platform.startswith('freebsd'))):
|
||||
# freebsd and Mach don't follow POSIX semantic of recv
|
||||
# and fail with ECONNRESET if peer performed shutdown.
|
||||
# See corresponding comment and code in TSocket::read()
|
||||
# in lib/cpp/src/transport/TSocket.cpp.
|
||||
self.close()
|
||||
# Trigger the check to raise the END_OF_FILE exception below.
|
||||
buff = ''
|
||||
else:
|
||||
raise
|
||||
|
||||
if len(buff) == 0:
|
||||
raise TTransportException(type=TTransportException.END_OF_FILE,
|
||||
message='TSocket read 0 bytes')
|
||||
return buff
|
||||
|
||||
def write(self, buff):
|
||||
self.sock.sendall(buff)
|
||||
|
||||
def flush(self):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
if not self.sock:
|
||||
return
|
||||
|
||||
try:
|
||||
self.sock.shutdown(socket.SHUT_RDWR)
|
||||
self.sock.close()
|
||||
except (socket.error, OSError):
|
||||
pass
|
||||
|
||||
|
||||
class TServerSocket(object):
|
||||
"""Socket implementation for server side."""
|
||||
|
||||
def __init__(self, host=None, port=None, unix_socket=None,
|
||||
socket_family=socket.AF_INET, client_timeout=3000,
|
||||
backlog=128):
|
||||
"""Initialize a TServerSocket
|
||||
|
||||
TSocket can be initialized in 2 ways:
|
||||
* host + port. can configure to use AF_INET/AF_INET6
|
||||
* unix_socket
|
||||
|
||||
@param host(str) The host to connect to
|
||||
@param port(int) The (TCP) port to connect to
|
||||
@param unix_socket(str) The filename of a unix socket to connect to
|
||||
@param socket_family(str) socket.AF_INET or socket.AF_INET6. only
|
||||
take effect when using host/port
|
||||
@param client_timeout client socket timeout
|
||||
@param backlog backlog for server socket
|
||||
"""
|
||||
|
||||
if unix_socket:
|
||||
self.unix_socket = unix_socket
|
||||
self.host = None
|
||||
self.port = None
|
||||
else:
|
||||
self.unix_socket = None
|
||||
self.host = host
|
||||
self.port = port
|
||||
|
||||
self.socket_family = socket_family
|
||||
self.client_timeout = client_timeout / 1000 if client_timeout else None
|
||||
self.backlog = backlog
|
||||
|
||||
def _init_sock(self):
|
||||
if self.unix_socket:
|
||||
# try remove the sock file it already exists
|
||||
_sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
try:
|
||||
_sock.connect(self.unix_socket)
|
||||
except (socket.error, OSError) as err:
|
||||
if err.args[0] == errno.ECONNREFUSED:
|
||||
os.unlink(self.unix_socket)
|
||||
else:
|
||||
_sock = socket.socket(self.socket_family, socket.SOCK_STREAM)
|
||||
|
||||
_sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
if hasattr(socket, "SO_REUSEPORT"):
|
||||
_sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1)
|
||||
_sock.settimeout(None)
|
||||
self.sock = _sock
|
||||
|
||||
def listen(self):
|
||||
self._init_sock()
|
||||
|
||||
addr = self.unix_socket or (self.host, self.port)
|
||||
self.sock.bind(addr)
|
||||
self.sock.listen(self.backlog)
|
||||
|
||||
def accept(self):
|
||||
client, _ = self.sock.accept()
|
||||
client.settimeout(self.client_timeout)
|
||||
return TSocket(sock=client)
|
||||
|
||||
def close(self):
|
||||
if not self.sock:
|
||||
return
|
||||
|
||||
try:
|
||||
self.sock.shutdown(socket.SHUT_RDWR)
|
||||
self.sock.close()
|
||||
except (socket.error, OSError):
|
||||
pass
|
||||
@@ -1,122 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import struct
|
||||
|
||||
from ._ssl import (
|
||||
create_thriftpy_context,
|
||||
RESTRICTED_SERVER_CIPHERS,
|
||||
DEFAULT_CIPHERS
|
||||
)
|
||||
from .socket import TSocket, TServerSocket
|
||||
|
||||
|
||||
class TSSLSocket(TSocket):
|
||||
"""SSL socket implementation for client side
|
||||
"""
|
||||
|
||||
def __init__(self, host, port, socket_family=socket.AF_INET,
|
||||
socket_timeout=3000, connect_timeout=None,
|
||||
ssl_context=None, validate=True,
|
||||
cafile=None, capath=None, certfile=None, keyfile=None,
|
||||
ciphers=DEFAULT_CIPHERS):
|
||||
"""Initialize a TSSLSocket
|
||||
|
||||
@param validate(bool) Set to False to disable SSL certificate
|
||||
validation and hostname validation. Default enabled.
|
||||
@param cafile(str) Path to a file of concatenated CA
|
||||
certificates in PEM format.
|
||||
@param capath(str) path to a directory containing several CA
|
||||
certificates in PEM format, following an OpenSSL specific layout.
|
||||
@param certfile(str) The certfile string must be the path to a
|
||||
single file in PEM format containing the certificate as well as
|
||||
any number of CA certificates needed to establish the
|
||||
certificate’s authenticity.
|
||||
@param keyfile(str) The keyfile string, if not present,
|
||||
the private key will be taken from certfile as well.
|
||||
@param ciphers(list<str>) The cipher suites to allow
|
||||
@param ssl_context(SSLContext) Customize the SSLContext, can be used
|
||||
to persist SSLContext object. Caution it's easy to get wrong, only
|
||||
use if you know what you're doing.
|
||||
|
||||
The `host` must be the same with server if validate enabled.
|
||||
"""
|
||||
super(TSSLSocket, self).__init__(
|
||||
host=host, port=port, socket_family=socket_family,
|
||||
connect_timeout=connect_timeout, socket_timeout=socket_timeout)
|
||||
|
||||
if ssl_context:
|
||||
self.ssl_context = ssl_context
|
||||
else:
|
||||
self.ssl_context = create_thriftpy_context(server_side=False,
|
||||
ciphers=ciphers)
|
||||
|
||||
if cafile or capath:
|
||||
self.ssl_context.load_verify_locations(cafile=cafile,
|
||||
capath=capath)
|
||||
|
||||
if certfile:
|
||||
self.ssl_context.load_cert_chain(certfile, keyfile=keyfile)
|
||||
|
||||
if not validate:
|
||||
self.ssl_context.check_hostname = False
|
||||
self.ssl_context.verify_mode = ssl.CERT_NONE
|
||||
|
||||
def _init_sock(self):
|
||||
_sock = socket.socket(self.socket_family, socket.SOCK_STREAM)
|
||||
_sock = self.ssl_context.wrap_socket(_sock,
|
||||
server_hostname=self.host)
|
||||
# socket options
|
||||
linger = struct.pack('ii', 0, 0)
|
||||
_sock.setsockopt(socket.SOL_SOCKET, socket.SO_LINGER, linger)
|
||||
_sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
|
||||
_sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
|
||||
self.sock = _sock
|
||||
|
||||
|
||||
class TSSLServerSocket(TServerSocket):
|
||||
"""SSL implementation of TServerSocket
|
||||
"""
|
||||
|
||||
def __init__(self, host, port, socket_family=socket.AF_INET,
|
||||
client_timeout=3000, backlog=128,
|
||||
ssl_context=None, certfile='cert.pem',
|
||||
ciphers=RESTRICTED_SERVER_CIPHERS):
|
||||
"""Initialize a TSSLServerSocket
|
||||
|
||||
@param certfile(str) The server cert pem filename
|
||||
@param ciphers(list<str>) The cipher suites to allow
|
||||
@param ssl_context(SSLContext) Customize the SSLContext, can be used
|
||||
to persist SSLContext object. Caution it's easy to get wrong, only
|
||||
use if you know what you're doing.
|
||||
"""
|
||||
super(TSSLServerSocket, self).__init__(
|
||||
host=host, port=port, socket_family=socket_family,
|
||||
client_timeout=client_timeout, backlog=backlog)
|
||||
|
||||
if ssl_context:
|
||||
self.ssl_context = ssl_context
|
||||
else:
|
||||
if not os.access(certfile, os.R_OK):
|
||||
raise IOError('No such certfile found: %s' % certfile)
|
||||
|
||||
self.ssl_context = create_thriftpy_context(server_side=True,
|
||||
ciphers=ciphers)
|
||||
self.ssl_context.load_cert_chain(certfile=certfile)
|
||||
|
||||
def accept(self):
|
||||
sock, _ = self.sock.accept()
|
||||
try:
|
||||
ssl_sock = self.ssl_context.wrap_socket(sock, server_side=True)
|
||||
except ssl.SSLError:
|
||||
# failed handshake/ssl wrap, close socket to client
|
||||
sock.shutdown(socket.SHUT_RDWR)
|
||||
sock.close()
|
||||
raise
|
||||
else:
|
||||
ssl_sock.settimeout(self.client_timeout)
|
||||
return TSocket(sock=ssl_sock)
|
||||
@@ -1,37 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import absolute_import
|
||||
|
||||
import binascii
|
||||
|
||||
from .protocol.binary import TBinaryProtocolFactory
|
||||
from .transport import TMemoryBuffer
|
||||
|
||||
|
||||
def serialize(thrift_object, proto_factory=TBinaryProtocolFactory()):
|
||||
transport = TMemoryBuffer()
|
||||
protocol = proto_factory.get_protocol(transport)
|
||||
thrift_object.write(protocol)
|
||||
protocol.write_message_end()
|
||||
return transport.getvalue()
|
||||
|
||||
|
||||
def deserialize(thrift_object, buf, proto_factory=TBinaryProtocolFactory()):
|
||||
transport = TMemoryBuffer(buf)
|
||||
protocol = proto_factory.get_protocol(transport)
|
||||
thrift_object.read(protocol)
|
||||
return thrift_object
|
||||
|
||||
|
||||
def hexlify(byte_array, delimeter=' '):
|
||||
s = binascii.hexlify(byte_array).decode('utf-8')
|
||||
return delimeter.join(a+b for a, b in zip(s[::2], s[1::2]))
|
||||
|
||||
|
||||
def hexprint(byte_array, delimeter=' ', count=10):
|
||||
print("Bytes:")
|
||||
print(byte_array)
|
||||
|
||||
print("\nHex:")
|
||||
g = hexlify(byte_array, delimeter).split(delimeter)
|
||||
print('\n'.join(' '.join(g[i:i+10]) for i in range(0, len(g), 10)))
|
||||
@@ -0,0 +1 @@
|
||||
#Thriftpy with dependencies
|
||||
Vendored
Vendored
Vendored
Vendored
Vendored
Vendored
Vendored
@@ -5,6 +5,7 @@
|
||||
<content url="file://$MODULE_DIR$/helpers">
|
||||
<sourceFolder url="file://$MODULE_DIR$/helpers" isTestSource="false" />
|
||||
<sourceFolder url="file://$MODULE_DIR$/helpers/pydev" isTestSource="false" packagePrefix="pydev" />
|
||||
<sourceFolder url="file://$MODULE_DIR$/helpers/third_party/thriftpy" isTestSource="false" />
|
||||
<excludeFolder url="file://$MODULE_DIR$/helpers/pydev/build" />
|
||||
</content>
|
||||
<orderEntry type="inheritedJdk" />
|
||||
|
||||
@@ -20,12 +20,14 @@ import com.intellij.execution.configurations.GeneralCommandLine;
|
||||
import com.intellij.execution.configurations.ParamsGroup;
|
||||
import com.intellij.openapi.projectRoots.Sdk;
|
||||
import com.intellij.openapi.util.io.FileUtil;
|
||||
import com.intellij.util.containers.ContainerUtil;
|
||||
import com.jetbrains.python.psi.LanguageLevel;
|
||||
import com.jetbrains.python.sdk.PythonEnvUtil;
|
||||
import com.jetbrains.python.sdk.PythonSdkType;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
|
||||
import java.io.File;
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
@@ -41,12 +43,12 @@ public enum PythonHelper implements HelperPackage {
|
||||
COVERAGEPY("coveragepy", ""),
|
||||
COVERAGE("coverage_runner", "run_coverage"),
|
||||
DEBUGGER("pydev", "pydevd"),
|
||||
|
||||
|
||||
ATTACH_DEBUGGER("pydev/pydevd_attach_to_process/attach_pydevd.py"),
|
||||
|
||||
CONSOLE("pydev", "pydevconsole"),
|
||||
CONSOLE("pydev", "pydevconsole", HelperDependency.THRIFTPY),
|
||||
RUN_IN_CONSOLE("pydev", "pydev_run_in_console"),
|
||||
PROFILER("profiler", "run_profiler"),
|
||||
PROFILER("profiler", "run_profiler", HelperDependency.THRIFTPY),
|
||||
|
||||
LOAD_ENTRY_POINT("pycharm", "pycharm_load_entry_point"),
|
||||
|
||||
@@ -91,41 +93,49 @@ public enum PythonHelper implements HelperPackage {
|
||||
public static final String PY2_HELPER_DEPENDENCIES_DIR = "py2only";
|
||||
|
||||
@NotNull
|
||||
private static PathHelperPackage findModule(String moduleEntryPoint, String path, boolean asModule) {
|
||||
private static PathHelperPackage findModule(String moduleEntryPoint, String path, boolean asModule, String[] thirdPartyDependencies) {
|
||||
List<HelperDependency> dependencies = HelperDependency.findThirdPartyDependencies(thirdPartyDependencies);
|
||||
|
||||
if (getHelperFile(path + ".zip").isFile()) {
|
||||
return new ModuleHelperPackage(moduleEntryPoint, path + ".zip");
|
||||
return new ModuleHelperPackage(moduleEntryPoint, path + ".zip", dependencies);
|
||||
}
|
||||
|
||||
if (!asModule && new File(getHelperFile(path), moduleEntryPoint + ".py").isFile()) {
|
||||
return new ScriptPythonHelper(moduleEntryPoint + ".py", getHelperFile(path));
|
||||
return new ScriptPythonHelper(moduleEntryPoint + ".py", getHelperFile(path), dependencies);
|
||||
}
|
||||
|
||||
return new ModuleHelperPackage(moduleEntryPoint, path);
|
||||
return new ModuleHelperPackage(moduleEntryPoint, path, dependencies);
|
||||
}
|
||||
|
||||
private final PathHelperPackage myModule;
|
||||
|
||||
PythonHelper(String pythonPath, String moduleName) {
|
||||
this(pythonPath, moduleName, false);
|
||||
PythonHelper(String pythonPath, String moduleName, String... dependencies) {
|
||||
this(pythonPath, moduleName, false, dependencies);
|
||||
}
|
||||
|
||||
PythonHelper(String pythonPath, String moduleName, boolean asModule) {
|
||||
myModule = findModule(moduleName, pythonPath, asModule);
|
||||
PythonHelper(String pythonPath, String moduleName, boolean asModule, String... dependencies) {
|
||||
myModule = findModule(moduleName, pythonPath, asModule, dependencies);
|
||||
}
|
||||
|
||||
PythonHelper(String helperScript) {
|
||||
myModule = new ScriptPythonHelper(helperScript, getHelpersRoot());
|
||||
myModule = new ScriptPythonHelper(helperScript, getHelpersRoot(), Collections.emptyList());
|
||||
}
|
||||
|
||||
public abstract static class PathHelperPackage implements HelperPackage {
|
||||
protected final File myPath;
|
||||
@NotNull
|
||||
protected final List<HelperDependency> myDependencies;
|
||||
|
||||
PathHelperPackage(String path) {
|
||||
PathHelperPackage(String path, @NotNull List<HelperDependency> dependencies) {
|
||||
myPath = new File(path);
|
||||
myDependencies = dependencies;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void addToPythonPath(@NotNull Map<String, String> environment) {
|
||||
// at first add dependencies
|
||||
myDependencies.forEach(dependency -> dependency.addToPythonPath(environment));
|
||||
// then add helper script
|
||||
PythonEnvUtil.addToPythonPath(environment, getPythonPathEntry());
|
||||
}
|
||||
|
||||
@@ -174,8 +184,8 @@ public enum PythonHelper implements HelperPackage {
|
||||
public static class ModuleHelperPackage extends PathHelperPackage {
|
||||
private final String myModuleName;
|
||||
|
||||
public ModuleHelperPackage(String moduleName, String relativePath) {
|
||||
super(getHelperFile(relativePath).getAbsolutePath());
|
||||
public ModuleHelperPackage(String moduleName, String relativePath, @NotNull List<HelperDependency> dependencies) {
|
||||
super(getHelperFile(relativePath).getAbsolutePath(), dependencies);
|
||||
this.myModuleName = moduleName;
|
||||
}
|
||||
|
||||
@@ -200,8 +210,8 @@ public enum PythonHelper implements HelperPackage {
|
||||
public static class ScriptPythonHelper extends PathHelperPackage {
|
||||
private final String myPythonPath;
|
||||
|
||||
public ScriptPythonHelper(String script, File pythonPath) {
|
||||
super(new File(pythonPath, script).getAbsolutePath());
|
||||
public ScriptPythonHelper(String script, File pythonPath, @NotNull List<HelperDependency> dependencies) {
|
||||
super(new File(pythonPath, script).getAbsolutePath(), dependencies);
|
||||
myPythonPath = pythonPath.getAbsolutePath();
|
||||
}
|
||||
|
||||
@@ -218,7 +228,38 @@ public enum PythonHelper implements HelperPackage {
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
private static class HelperDependency {
|
||||
private static final String THRIFTPY = "thriftpy";
|
||||
|
||||
@NotNull
|
||||
private final String myPythonPath;
|
||||
|
||||
private HelperDependency(@NotNull String pythonPath) {myPythonPath = pythonPath;}
|
||||
|
||||
public void addToPythonPath(@NotNull Map<String, String> environment) {
|
||||
PythonEnvUtil.addToPythonPath(environment, myPythonPath);
|
||||
}
|
||||
|
||||
@NotNull
|
||||
private static List<HelperDependency> findThirdPartyDependencies(String... dependencies) {
|
||||
if (dependencies == null) {
|
||||
return Collections.emptyList();
|
||||
}
|
||||
return ContainerUtil.map(dependencies, s -> getThirdPartyDependency(s));
|
||||
}
|
||||
|
||||
@NotNull
|
||||
private static HelperDependency getThirdPartyDependency(@NotNull String name) {
|
||||
String path = new File(getHelpersThirdPartyDir(), name).getAbsolutePath();
|
||||
return new HelperDependency(path);
|
||||
}
|
||||
|
||||
@NotNull
|
||||
private static File getHelpersThirdPartyDir() {
|
||||
return getHelperFile("third_party");
|
||||
}
|
||||
}
|
||||
|
||||
@NotNull
|
||||
@Override
|
||||
public String getPythonPathEntry() {
|
||||
@@ -252,5 +293,4 @@ public enum PythonHelper implements HelperPackage {
|
||||
public GeneralCommandLine newCommandLine(@NotNull Sdk pythonSdk, @NotNull List<String> parameters) {
|
||||
return myModule.newCommandLine(pythonSdk, parameters);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user