PY-18029 Move thriftpy to third-party dependencies for Python helpers

This commit is contained in:
Alexander Koshevoy
2018-08-22 23:16:40 +03:00
parent 4c4e4d1a63
commit bdc5a3b50e
87 changed files with 61 additions and 11298 deletions
-5
View File
@@ -1,5 +0,0 @@
# PLY package
# Author: David Beazley (dave@dabeaz.com)
__version__ = '3.7'
__all__ = ['lex','yacc']
-908
View File
@@ -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)
-133
View File
@@ -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
-74
View File
@@ -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()
-11
View File
@@ -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"]
-126
View File
@@ -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
}
-35
View File
@@ -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)
-81
View File
@@ -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()
-104
View File
@@ -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
-393
View File
@@ -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'
-231
View File
@@ -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)
-37
View File
@@ -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)))
+1
View File
@@ -0,0 +1 @@
#Thriftpy with dependencies
+1
View File
@@ -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);
}
}