Visitar URL original
Infer native calling convention flags from argument metadata by youknowone 路 Pull Request #9005 路 RustPython/RustPython 路 GitHub
Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 3 additions & 9 deletions crates/derive-impl/src/pyclass.rs
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
use super::Diagnostic;
use crate::util::{
ALL_ALLOWED_NAMES, ClassItemMeta, ContentItem, ContentItemInner, ErrorVec, ExceptionItemMeta,
ItemMeta, ItemMetaInner, ItemNursery, SimpleItemMeta, infer_native_call_flags,
internal_doc_tokens, pyclass_ident_and_attrs, pyexception_ident_and_attrs,
ItemMeta, ItemMetaInner, ItemNursery, SimpleItemMeta, internal_doc_tokens,
pyclass_ident_and_attrs, pyexception_ident_and_attrs,
};
use core::str::FromStr;
use proc_macro2::{Delimiter, Group, Span, TokenStream, TokenTree};
Expand Down Expand Up @@ -1307,9 +1307,6 @@ where
_ => None,
}
};
let drop_first_typed = usize::from(implicit_self.is_some());
let call_flags = infer_native_call_flags(func.sig(), drop_first_typed);

// Add #[allow(non_snake_case)] for setter methods like set___name__
let method_name = ident.to_string();
if method_name.starts_with("set_") && method_name.contains("__") {
Expand Down Expand Up @@ -1341,7 +1338,6 @@ where
raw,
coexist,
attr_name: self.inner.attr_name,
call_flags,
});
Ok(())
}
Expand Down Expand Up @@ -1538,7 +1534,6 @@ struct MethodNurseryItem {
doc: TokenStream,
doc_body_pending: TokenStream,
attr_name: AttrName,
call_flags: TokenStream,
}

impl MethodNursery {
Expand Down Expand Up @@ -1578,15 +1573,14 @@ impl ToTokens for MethodNursery {
}
_ => unreachable!(),
};
let call_flags = &item.call_flags;
let coexist_flags = if item.coexist {
quote! { | rustpython_vm::function::PyMethodFlags::COEXIST.bits() }
} else {
quote! {}
};
let flags = quote! {
rustpython_vm::function::PyMethodFlags::from_bits_retain(
(#binding_flags).bits() | (#call_flags).bits() #coexist_flags
(#binding_flags).bits() #coexist_flags
)
};
// TODO: intern
Expand Down
10 changes: 3 additions & 7 deletions crates/derive-impl/src/pymodule.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,8 @@ use crate::error::Diagnostic;
use crate::pystructseq::PyStructSequenceMeta;
use crate::util::{
ALL_ALLOWED_NAMES, AttrItemMeta, AttributeExt, ClassItemMeta, ContentItem, ContentItemInner,
ErrorVec, ItemMeta, ItemNursery, ModuleItemMeta, SimpleItemMeta, infer_native_call_flags,
internal_doc_tokens, iter_use_idents, pyclass_ident_and_attrs,
ErrorVec, ItemMeta, ItemNursery, ModuleItemMeta, SimpleItemMeta, internal_doc_tokens,
iter_use_idents, pyclass_ident_and_attrs,
};
use core::str::FromStr;
use proc_macro2::{Delimiter, Group, TokenStream, TokenTree};
Expand Down Expand Up @@ -536,7 +536,6 @@ struct FunctionNurseryItem {
ident: Ident,
/// One internal doc per [`py_names`](Self::py_names) entry.
docs: Vec<TokenStream>,
call_flags: TokenStream,
}

impl FunctionNursery {
Expand Down Expand Up @@ -566,14 +565,13 @@ impl ToTokens for ValidatedFunctionNursery {
let ident = &item.ident;
let cfgs = &item.cfgs;
let cfgs = quote!(#(#cfgs)*);
let flags = &item.call_flags;
for (py_name, doc) in item.py_names.iter().zip(&item.docs) {
inner_tokens.extend(quote![
#cfgs
rustpython_vm::function::PyMethodDef::new_const(
#py_name,
#ident,
#flags,
rustpython_vm::function::PyMethodFlags::empty(),
#doc,
),
]);
Expand Down Expand Up @@ -711,14 +709,12 @@ impl ModuleItem for FunctionItem {
internal_doc_tokens(func.sig(), py_name, None, doc, None, Some("$module"))
})
.collect();
let call_flags = infer_native_call_flags(func.sig(), 0);

args.context.function_items.add_item(FunctionNurseryItem {
ident: ident.to_owned(),
py_names,
cfgs: args.cfgs.to_vec(),
docs,
call_flags,
});
Ok(())
}
Expand Down
72 changes: 0 additions & 72 deletions crates/derive-impl/src/util.rs
Original file line number Diff line number Diff line change
Expand Up @@ -936,75 +936,3 @@ pub(crate) fn internal_doc_tokens(
}
}
}

pub(crate) fn infer_native_call_flags(sig: &Signature, drop_first_typed: usize) -> TokenStream {
// Best-effort mapping of Rust function signatures to CPython-style
// METH_* calling convention flags used by CALL specialization.
let mut typed_args = Vec::new();
for arg in &sig.inputs {
let FnArg::Typed(typed) = arg else {
continue;
};
let ty_tokens = &typed.ty;
let ty = quote!(#ty_tokens).to_string().replace(' ', "");
// The interpreter supplies `vm` and `callee`; a Python call never
// passes them.
if (ty.starts_with('&') && ty.ends_with("VirtualMachine")) || ty.ends_with("Callee") {
continue;
}
typed_args.push(ty);
}

let mut user_args = typed_args.into_iter();
for _ in 0..drop_first_typed {
if user_args.next().is_none() {
break;
}
}

let mut has_keywords = false;
let mut variable_arity = false;
let mut fixed_positional = 0usize;

for ty in user_args {
let is_named = |name: &str| {
ty == name
|| ty.starts_with(&format!("{name}<"))
|| ty.contains(&format!("::{name}<"))
|| ty.ends_with(&format!("::{name}"))
};

if is_named("FuncArgs") {
has_keywords = true;
variable_arity = true;
continue;
}
if is_named("KwArgs") {
has_keywords = true;
variable_arity = true;
continue;
}
if is_named("PosArgs") || is_named("OptionalArg") || is_named("OptionalOption") {
variable_arity = true;
continue;
}
fixed_positional += 1;
}

if has_keywords {
quote! {
rustpython_vm::function::PyMethodFlags::from_bits_retain(
rustpython_vm::function::PyMethodFlags::FASTCALL.bits()
| rustpython_vm::function::PyMethodFlags::KEYWORDS.bits()
)
}
} else if variable_arity {
quote! { rustpython_vm::function::PyMethodFlags::FASTCALL }
} else {
match fixed_positional {
0 => quote! { rustpython_vm::function::PyMethodFlags::NOARGS },
1 => quote! { rustpython_vm::function::PyMethodFlags::O },
_ => quote! { rustpython_vm::function::PyMethodFlags::FASTCALL },
}
}
}
25 changes: 24 additions & 1 deletion crates/vm/src/function/builtin.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
use super::{Callee, FromArgs, FuncArgs};
use super::{Callee, FromArgs, FuncArgs, SigArg};
use crate::{
Py, PyPayload, PyRef, PyResult, VirtualMachine, convert::ToPyResult,
object::PyThreadingConstraint,
Expand Down Expand Up @@ -37,6 +37,10 @@ impl<F> PyNativeFn for F where
/// just pass an unconstrained generic type, e.g.
/// `fn foo<F, FKind>(f: F) where F: IntoPyNativeFn<FKind>`
pub trait IntoPyNativeFn<Kind>: Sized + PyThreadingConstraint + 'static {
/// Binding metadata of each argument, from which the calling convention
/// is inferred. A `&self` receiver is the `$self` marker.
const ARGS: &'static [SigArg] = &[SigArg::from_arg::<FuncArgs>("args")];

fn call(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult;

/// `IntoPyNativeFn::into_func()` generates a PyNativeFn that performs the
Expand Down Expand Up @@ -89,6 +93,8 @@ impl<F, T, R, VM> IntoPyNativeFn<(T, R, VM)> for F
where
F: PyNativeFnInternal<T, R, VM>,
{
const ARGS: &'static [SigArg] = F::ARGS;

#[inline(always)]
fn call(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
self.call_(vm, args, callee)
Expand All @@ -98,6 +104,7 @@ where
mod sealed {
use super::*;
pub trait PyNativeFnInternal<T, R, VM>: Sized + PyThreadingConstraint + 'static {
const ARGS: &'static [SigArg];
fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult;
}
}
Expand All @@ -124,6 +131,8 @@ macro_rules! into_py_native_fn_tuple {
$($T: FromArgs,)*
R: ToPyResult,
{
const ARGS: &'static [SigArg] = &[$(SigArg::from_arg::<$T>("")),*];

fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
let ($($n,)*) = args.bind_for::<($($T,)*)>(vm, callee)?;

Expand All @@ -138,6 +147,9 @@ macro_rules! into_py_native_fn_tuple {
$($T: FromArgs,)*
R: ToPyResult,
{
const ARGS: &'static [SigArg] =
&[SigArg::marker("$self") $(, SigArg::from_arg::<$T>(""))*];

fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
let (zelf, $($n,)*) = args.bind_for::<(PyRef<S>, $($T,)*)>(vm, callee)?;

Expand All @@ -152,6 +164,9 @@ macro_rules! into_py_native_fn_tuple {
$($T: FromArgs,)*
R: ToPyResult,
{
const ARGS: &'static [SigArg] =
&[SigArg::marker("$self") $(, SigArg::from_arg::<$T>(""))*];

fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
let (zelf, $($n,)*) = args.bind_for::<(PyRef<S>, $($T,)*)>(vm, callee)?;

Expand All @@ -165,6 +180,8 @@ macro_rules! into_py_native_fn_tuple {
$($T: FromArgs,)*
R: ToPyResult,
{
const ARGS: &'static [SigArg] = &[$(SigArg::from_arg::<$T>("")),*];

fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
let ($($n,)*) = args.bind_for::<($($T,)*)>(vm, callee)?;

Expand All @@ -179,6 +196,9 @@ macro_rules! into_py_native_fn_tuple {
$($T: FromArgs,)*
R: ToPyResult,
{
const ARGS: &'static [SigArg] =
&[SigArg::marker("$self") $(, SigArg::from_arg::<$T>(""))*];

fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
let (zelf, $($n,)*) = args.bind_for::<(PyRef<S>, $($T,)*)>(vm, callee)?;

Expand All @@ -193,6 +213,9 @@ macro_rules! into_py_native_fn_tuple {
$($T: FromArgs,)*
R: ToPyResult,
{
const ARGS: &'static [SigArg] =
&[SigArg::marker("$self") $(, SigArg::from_arg::<$T>(""))*];

fn call_(&self, vm: &VirtualMachine, args: FuncArgs, callee: Callee) -> PyResult {
let (zelf, $($n,)*) = args.bind_for::<(PyRef<S>, $($T,)*)>(vm, callee)?;

Expand Down
26 changes: 21 additions & 5 deletions crates/vm/src/function/method.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ use crate::{
descriptor::{PyClassMethodDescriptor, PyMethodDescriptor},
},
class::PyClassDef,
function::{IntoPyNativeFn, PyNativeFn},
function::{FuncArgs, IntoPyNativeFn, PyNativeFn, SigArg},
};

bitflags::bitflags! {
Expand Down Expand Up @@ -49,6 +49,22 @@ bitflags::bitflags! {
impl PyMethodFlags {
// FIXME: macro temp
pub const EMPTY: Self = Self::empty();

const CALL_CONVENTION: Self = Self::VARARGS
.union(Self::KEYWORDS)
.union(Self::NOARGS)
.union(Self::O)
.union(Self::FASTCALL);

/// Adds the calling convention inferred from `args` unless one is set.
/// For METHOD and CLASS the first argument is the bound receiver.
pub(crate) const fn with_call_convention(self, args: &[SigArg]) -> Self {
if self.intersects(Self::CALL_CONVENTION) {
return self;
}
let receiver = self.intersects(Self::METHOD.union(Self::CLASS));
self.union(super::signature::native_call_flags(args, receiver))
}
}

#[macro_export]
Expand Down Expand Up @@ -115,16 +131,16 @@ impl PyMethodDef {
}

#[inline]
pub const fn new_const<Kind>(
pub const fn new_const<Kind, F: IntoPyNativeFn<Kind>>(
name: &'static str,
func: impl IntoPyNativeFn<Kind>,
func: F,
flags: PyMethodFlags,
doc: super::ItemDoc,
) -> Self {
Self {
name,
func: super::static_func(func),
flags,
flags: flags.with_call_convention(F::ARGS),
#[cfg(feature = "doc")]
doc_off: doc.offset,
#[cfg(feature = "doc")]
Expand All @@ -145,7 +161,7 @@ impl PyMethodDef {
Self {
name,
func: super::static_raw_func(func),
flags,
flags: flags.with_call_convention(&[SigArg::from_arg::<FuncArgs>("args")]),
#[cfg(feature = "doc")]
doc_off: doc.offset,
#[cfg(feature = "doc")]
Expand Down
Loading
Loading