Add downcasting
This commit is contained in:
parent
eba88b32ba
commit
3c81157ede
6 changed files with 106 additions and 32 deletions
|
|
@ -28,10 +28,10 @@ class Namespace(object):
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return self.name
|
return self.name
|
||||||
|
|
||||||
class Type(object):
|
class _Type(object):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
class BuiltinTypes(Type):
|
class BuiltinTypes(_Type):
|
||||||
def __init__(self, name):
|
def __init__(self, name):
|
||||||
self.name = name
|
self.name = name
|
||||||
|
|
||||||
|
|
@ -47,7 +47,7 @@ Unsigned = BuiltinTypes('unsigned')
|
||||||
Bool = BuiltinTypes('bool')
|
Bool = BuiltinTypes('bool')
|
||||||
ConstStdString = BuiltinTypes('const std::string')
|
ConstStdString = BuiltinTypes('const std::string')
|
||||||
|
|
||||||
class Class(Type):
|
class Class(_Type):
|
||||||
format = 'O'
|
format = 'O'
|
||||||
|
|
||||||
def __init__(self, ns, *bases):
|
def __init__(self, ns, *bases):
|
||||||
|
|
@ -57,9 +57,11 @@ class Class(Type):
|
||||||
self.methods = []
|
self.methods = []
|
||||||
self.enums = []
|
self.enums = []
|
||||||
self.includes = set()
|
self.includes = set()
|
||||||
|
self.downcastables = set()
|
||||||
|
|
||||||
def __call__(self, defn):
|
def __call__(self, defn):
|
||||||
assert not self._is_defined
|
assert not self._is_defined
|
||||||
|
# process the definition in "defn"
|
||||||
self.name = defn.__name__
|
self.name = defn.__name__
|
||||||
for k, v in defn.__dict__.items():
|
for k, v in defn.__dict__.items():
|
||||||
if isinstance(v, Method):
|
if isinstance(v, Method):
|
||||||
|
|
@ -79,6 +81,12 @@ class Class(Type):
|
||||||
else:
|
else:
|
||||||
for i in v:
|
for i in v:
|
||||||
self.includes.add(i)
|
self.includes.add(i)
|
||||||
|
elif k == '_downcast_':
|
||||||
|
if isinstance(v, Class):
|
||||||
|
self.downcastables.add(v)
|
||||||
|
else:
|
||||||
|
for i in v:
|
||||||
|
self.downcastables.add(i)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def compile_cpp(self, writer):
|
def compile_cpp(self, writer):
|
||||||
|
|
@ -106,12 +114,11 @@ class Class(Type):
|
||||||
bases = ', '.join(x.name for x in self.bases)
|
bases = ', '.join(x.name for x in self.bases)
|
||||||
writer.println('@capsule.register_class')
|
writer.println('@capsule.register_class')
|
||||||
with writer.block('class %(clsname)s(%(bases)s):' % locals()):
|
with writer.block('class %(clsname)s(%(bases)s):' % locals()):
|
||||||
|
writer.println('_llvm_type_ = "%s"' % self.fullname)
|
||||||
for enum in self.enums:
|
for enum in self.enums:
|
||||||
enum.compile_py(writer)
|
enum.compile_py(writer)
|
||||||
for meth in self.methods:
|
for meth in self.methods:
|
||||||
meth.compile_py(writer)
|
meth.compile_py(writer)
|
||||||
if not self.enums and not self.methods:
|
|
||||||
writer.println('pass')
|
|
||||||
writer.println()
|
writer.println()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|
@ -178,6 +185,7 @@ class Enum(object):
|
||||||
|
|
||||||
def compile_py(self, writer):
|
def compile_py(self, writer):
|
||||||
with writer.block('class %s:' % self.name):
|
with writer.block('class %s:' % self.name):
|
||||||
|
writer.println('_llvm_type_ = "%s"' % self.fullname)
|
||||||
for v in self.value_names:
|
for v in self.value_names:
|
||||||
writer.println('%(v)s = "%(v)s"' % locals())
|
writer.println('%(v)s = "%(v)s"' % locals())
|
||||||
writer.println()
|
writer.println()
|
||||||
|
|
@ -340,7 +348,7 @@ class Constructor(StaticMethod):
|
||||||
ret = writer.declare(retty.fullname, stmt)
|
ret = writer.declare(retty.fullname, stmt)
|
||||||
return ret
|
return ret
|
||||||
|
|
||||||
class ref(Type):
|
class ref(_Type):
|
||||||
def __init__(self, element):
|
def __init__(self, element):
|
||||||
assert isinstance(element, Class), type(element)
|
assert isinstance(element, Class), type(element)
|
||||||
self.element = element
|
self.element = element
|
||||||
|
|
@ -369,7 +377,7 @@ class ref(Type):
|
||||||
return writer.declare(self.fullname, '*%s' % p)
|
return writer.declare(self.fullname, '*%s' % p)
|
||||||
|
|
||||||
|
|
||||||
class ptr(Type):
|
class ptr(_Type):
|
||||||
def __init__(self, element):
|
def __init__(self, element):
|
||||||
assert isinstance(element, Class)
|
assert isinstance(element, Class)
|
||||||
self.element = element
|
self.element = element
|
||||||
|
|
@ -393,7 +401,7 @@ class ptr(Type):
|
||||||
return writer.pycapsule_new(val, self.element.capsule_name,
|
return writer.pycapsule_new(val, self.element.capsule_name,
|
||||||
self.element.fullname)
|
self.element.fullname)
|
||||||
|
|
||||||
class cast(Type):
|
class cast(_Type):
|
||||||
format = 'O'
|
format = 'O'
|
||||||
|
|
||||||
def __init__(self, original, target):
|
def __init__(self, original, target):
|
||||||
|
|
@ -406,14 +414,14 @@ class cast(Type):
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def python_type(self):
|
def python_type(self):
|
||||||
if not isinstance(self.target, Type):
|
if not isinstance(self.target, _Type):
|
||||||
return self.target
|
return self.target
|
||||||
else:
|
else:
|
||||||
return self.original
|
return self.original
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def binding_type(self):
|
def binding_type(self):
|
||||||
if isinstance(self.target, Type):
|
if isinstance(self.target, _Type):
|
||||||
return self.target
|
return self.target
|
||||||
else:
|
else:
|
||||||
return self.original
|
return self.original
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,7 @@ def set_debug(enabled):
|
||||||
|
|
||||||
|
|
||||||
class WeakRef(ref):
|
class WeakRef(ref):
|
||||||
__slots__ = 'capsule', 'dtor'
|
__slots__ = 'capsule', 'dtor', 'owning'
|
||||||
|
|
||||||
_pyclasses = {}
|
_pyclasses = {}
|
||||||
_addr2obj = WeakValueDictionary()
|
_addr2obj = WeakValueDictionary()
|
||||||
|
|
@ -33,6 +33,7 @@ def classof(cap):
|
||||||
return _pyclasses[cls]
|
return _pyclasses[cls]
|
||||||
|
|
||||||
def _capsule_destructor(weak):
|
def _capsule_destructor(weak):
|
||||||
|
if weak.owning:
|
||||||
cap = weak.capsule
|
cap = weak.capsule
|
||||||
addr = _capsule.getPointer(cap)
|
addr = _capsule.getPointer(cap)
|
||||||
cls = _capsule.getClassName(cap)
|
cls = _capsule.getClassName(cap)
|
||||||
|
|
@ -40,6 +41,13 @@ def _capsule_destructor(weak):
|
||||||
weak.dtor(cap)
|
weak.dtor(cap)
|
||||||
del _owners[addr]
|
del _owners[addr]
|
||||||
|
|
||||||
|
def release_ownership(old):
|
||||||
|
addr = _capsule.getPointer(old)
|
||||||
|
if hasattr(old, '_delete_'):
|
||||||
|
oldweak = _owners[addr]
|
||||||
|
oldweak.owning = False # dis-own
|
||||||
|
del _owners[addr]
|
||||||
|
|
||||||
def wrap(cap):
|
def wrap(cap):
|
||||||
'''Wrap a PyCapsule with the corresponding Wrapper class.
|
'''Wrap a PyCapsule with the corresponding Wrapper class.
|
||||||
If `cap` is not a PyCapsule, returns `cap`
|
If `cap` is not a PyCapsule, returns `cap`
|
||||||
|
|
@ -60,13 +68,26 @@ def wrap(cap):
|
||||||
weak = WeakRef(obj, _capsule_destructor)
|
weak = WeakRef(obj, _capsule_destructor)
|
||||||
_owners[addr] = weak
|
_owners[addr] = weak
|
||||||
weak.capsule = cap
|
weak.capsule = cap
|
||||||
|
weak.owning = True
|
||||||
weak.dtor = cls._delete_
|
weak.dtor = cls._delete_
|
||||||
else:
|
else:
|
||||||
assert classof(obj._ptr) is classof(cap)
|
oldcls = classof(obj._ptr)
|
||||||
# Unset destructor for capsules that are repeated
|
newcls = classof(cap)
|
||||||
_capsule.setDestructor(cap, None)
|
if issubclass(oldcls, newcls):
|
||||||
|
# do auto downcast
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
assert oldcls is newcls
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
|
def downcast(old, new):
|
||||||
|
assert old is not new
|
||||||
|
oldcls = classof(old)
|
||||||
|
newcls = classof(new)
|
||||||
|
assert issubclass(newcls, oldcls)
|
||||||
|
release_ownership(old)
|
||||||
|
del _addr2obj[_capsule.getPointer(old)] # clear cache
|
||||||
|
return wrap(new)
|
||||||
|
|
||||||
def unwrap(obj):
|
def unwrap(obj):
|
||||||
'''Unwrap a Wrapper instance into the underlying PyCapsule.
|
'''Unwrap a Wrapper instance into the underlying PyCapsule.
|
||||||
|
|
@ -96,8 +117,3 @@ class Wrapper(object):
|
||||||
def _ptr(self):
|
def _ptr(self):
|
||||||
return self.__ptr
|
return self.__ptr
|
||||||
|
|
||||||
def _release_ownership(self):
|
|
||||||
_capsule.setDestructor(self._ptr, None)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -3,8 +3,13 @@ Wrapped the extra functions in _api.so
|
||||||
'''
|
'''
|
||||||
|
|
||||||
import capsule
|
import capsule
|
||||||
|
import _api
|
||||||
|
#
|
||||||
|
# Re-export the native API from the _api.extra and wrap the functions
|
||||||
|
#
|
||||||
|
|
||||||
def _wrapper(func):
|
def _wrapper(func):
|
||||||
|
"Wrap the re-exported functions"
|
||||||
def _core(*args):
|
def _core(*args):
|
||||||
unwrapped = map(capsule.unwrap, args)
|
unwrapped = map(capsule.unwrap, args)
|
||||||
ret = func(*unwrapped)
|
ret = func(*unwrapped)
|
||||||
|
|
@ -12,9 +17,27 @@ def _wrapper(func):
|
||||||
return _core
|
return _core
|
||||||
|
|
||||||
def _init(glob):
|
def _init(glob):
|
||||||
from _api import extra
|
for k, v in _api.extra.__dict__.items():
|
||||||
for k, v in extra.__dict__.items():
|
|
||||||
glob[k] = _wrapper(v)
|
glob[k] = _wrapper(v)
|
||||||
|
|
||||||
_init(globals())
|
_init(globals())
|
||||||
|
|
||||||
|
#
|
||||||
|
# Downcasts
|
||||||
|
#
|
||||||
|
|
||||||
|
def downcast(obj, cls):
|
||||||
|
if type(obj) is cls:
|
||||||
|
return obj
|
||||||
|
fromty = obj._llvm_type_
|
||||||
|
toty = cls._llvm_type_
|
||||||
|
fname = 'downcast_%s_to_%s' % (fromty, toty)
|
||||||
|
fname = fname.replace('::', '_')
|
||||||
|
try:
|
||||||
|
caster = getattr(_api, fname)
|
||||||
|
except AttributeError:
|
||||||
|
fmt = "Downcast from %s to %s is not supported"
|
||||||
|
raise TypeError(fmt % (fromty, toty))
|
||||||
|
old = capsule.unwrap(obj)
|
||||||
|
new = caster(old)
|
||||||
|
return capsule.downcast(old, new)
|
||||||
|
|
@ -72,6 +72,20 @@ def main():
|
||||||
print cls
|
print cls
|
||||||
units.append(cls)
|
units.append(cls)
|
||||||
|
|
||||||
|
# add extra stuffs
|
||||||
|
downcastlist = []
|
||||||
|
## add downcast
|
||||||
|
for cls in units:
|
||||||
|
if isinstance(cls, Class):
|
||||||
|
for bcls in cls.downcastables:
|
||||||
|
from_to = bcls.fullname, cls.fullname
|
||||||
|
name = 'downcast_%s_to_%s' % tuple(map(codegen.mangle, from_to))
|
||||||
|
fn = Function(namespaces[''], name, ptr(cls), ptr(bcls))
|
||||||
|
downcastlist.append((from_to, fn))
|
||||||
|
units.append(fn)
|
||||||
|
|
||||||
|
|
||||||
|
# generate cpp source
|
||||||
with open('%s.cpp' % outputfilename, 'w') as cppfile:
|
with open('%s.cpp' % outputfilename, 'w') as cppfile:
|
||||||
println = wrap_println_from_file(cppfile)
|
println = wrap_println_from_file(cppfile)
|
||||||
|
|
||||||
|
|
@ -87,6 +101,18 @@ def main():
|
||||||
println('#include "%s"' % inc)
|
println('#include "%s"' % inc)
|
||||||
println()
|
println()
|
||||||
|
|
||||||
|
# generate downcast
|
||||||
|
for ((fromty, toty), fn) in downcastlist:
|
||||||
|
name = fn.name
|
||||||
|
fmt = '''
|
||||||
|
static
|
||||||
|
%(toty)s* %(name)s(%(fromty)s* arg)
|
||||||
|
{
|
||||||
|
return typecast<%(toty)s>::from(arg);
|
||||||
|
}
|
||||||
|
'''
|
||||||
|
println(fmt % locals())
|
||||||
|
|
||||||
# write methods and method tables
|
# write methods and method tables
|
||||||
for u in units:
|
for u in units:
|
||||||
writer = codegen.CppCodeWriter(println)
|
writer = codegen.CppCodeWriter(println)
|
||||||
|
|
|
||||||
|
|
@ -40,6 +40,6 @@ static PyMethodDef extra_methodtable[] = {
|
||||||
#define method(func) { #func, (PyCFunction)func, METH_VARARGS, NULL }
|
#define method(func) { #func, (PyCFunction)func, METH_VARARGS, NULL }
|
||||||
method( make_raw_ostream_for_printing ),
|
method( make_raw_ostream_for_printing ),
|
||||||
method( make_small_vector_from_types ),
|
method( make_small_vector_from_types ),
|
||||||
{ NULL }
|
|
||||||
#undef method
|
#undef method
|
||||||
|
{ NULL }
|
||||||
};
|
};
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,4 @@
|
||||||
from binding import *
|
from binding import *
|
||||||
|
|
||||||
llvm = Namespace('llvm')
|
llvm = Namespace('llvm')
|
||||||
|
default = Namespace('')
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue