Allow retrieving argument attributes

This commit is contained in:
Siu Kwan Lam 2013-05-16 13:45:59 -05:00
commit 592f9c6988
3 changed files with 66 additions and 0 deletions

View file

@ -1372,6 +1372,8 @@ class GlobalVariable(GlobalValue):
class Argument(Value):
_type_ = api.llvm.Argument
_valid_attrs = frozenset([ATTR_BY_VAL, ATTR_NEST, ATTR_NO_ALIAS,
ATTR_NO_CAPTURE, ATTR_STRUCT_RET])
def add_attribute(self, attr):
context = api.llvm.getGlobalContext()
@ -1379,6 +1381,9 @@ class Argument(Value):
attrbldr.addAttribute(attr)
attrs = api.llvm.Attributes.get(context, attrbldr)
self._ptr.addAttr(attrs)
if attr not in self:
raise ValueError("Attribute %s is not valid for arg %s" %
(attr, self))
def remove_attribute(self, attr):
context = api.llvm.getGlobalContext()
@ -1400,6 +1405,39 @@ class Argument(Value):
alignment = property(_get_alignment,
_set_alignment)
def __contains__(self, attr):
if attr == ATTR_BY_VAL:
return self.has_by_val()
elif attr == ATTR_NEST:
return self.has_nest()
elif attr == ATTR_NO_ALIAS:
return self.has_no_alias()
elif attr == ATTR_NO_CAPTURE:
return self.has_no_capture()
elif attr == ATTR_STRUCT_RET:
return self.has_struct_ret()
else:
raise ValueError('invalid attribute for argument')
@property
def arg_no(self):
return self._ptr.getArgNo()
def has_by_val(self):
return self._ptr.hasByValAttr()
def has_nest(self):
return self._ptr.hasNestAttr()
def has_no_alias(self):
return self._ptr.hasNoAliasAttr()
def has_no_capture(self):
return self._ptr.hasNoCaptureAttr()
def has_struct_ret(self):
return self._ptr.hasStructRetAttr()
class Function(GlobalValue):
_type_ = api.llvm.Function

View file

@ -1283,6 +1283,24 @@ class TestTypeHash(TestCase):
tests.append(TestTypeHash)
# ---------------------------------------------------------------------------
class TestArgAttr(TestCase):
def test_arg_attr(self):
m = Module.new('oifjda')
vptr = Type.pointer(Type.float())
sptr = Type.pointer(Type.struct([]))
fnty = Type.function(Type.void(), [vptr] * 5)
func = m.add_function(fnty, 'foo')
attrs = [lc.ATTR_STRUCT_RET, lc.ATTR_BY_VAL, lc.ATTR_NEST,
lc.ATTR_NO_ALIAS, lc.ATTR_NO_CAPTURE]
for i, attr in enumerate(attrs):
arg = func.args[i]
self.assertEqual(i, arg.arg_no)
arg.add_attribute(attr)
self.assertTrue(attr in func.args[i])
tests.append(TestArgAttr)
# ---------------------------------------------------------------------------

View file

@ -13,3 +13,13 @@ class Argument:
removeAttr = Method(Void, ref(Attributes))
getParamAlignment = Method(cast(Unsigned, int))
getArgNo = Method(cast(Unsigned, int))
hasByValAttr = Method(cast(Bool, bool))
hasNestAttr = Method(cast(Bool, bool))
hasNoAliasAttr = Method(cast(Bool, bool))
hasNoCaptureAttr = Method(cast(Bool, bool))
hasStructRetAttr = Method(cast(Bool, bool))
if LLVM_VERSION > (3, 2):
hasReturnedAttr = Method(cast(Bool, bool))