use crate::{
args::{ArgPosIter, ArgValues},
bytecode::VM,
defer_drop_mut,
exception_private::{ExcType, RunResult, SimpleException},
expressions::Identifier,
heap::{HeapData, HeapGuard},
heap_traits::DropWithHeap,
intern::{Interns, StringId},
resource::ResourceTracker,
types::{Dict, allocate_tuple},
value::Value,
};
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub(crate) struct Signature {
pos_args: Option<Vec<StringId>>,
pos_defaults_count: usize,
args: Option<Vec<StringId>>,
arg_defaults_count: usize,
var_args: Option<StringId>,
kwargs: Option<Vec<StringId>>,
kwarg_default_map: Option<Vec<Option<usize>>>,
var_kwargs: Option<StringId>,
bind_mode: BindMode,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
enum BindMode {
#[default]
Simple,
SimpleWithDefaults,
Complex,
}
impl Signature {
#[expect(clippy::too_many_arguments)]
pub fn new(
pos_args: Vec<StringId>,
pos_defaults_count: usize,
args: Vec<StringId>,
arg_defaults_count: usize,
var_args: Option<StringId>,
kwargs: Vec<StringId>,
kwarg_default_map: Vec<Option<usize>>,
var_kwargs: Option<StringId>,
) -> Self {
let pos_args = if pos_args.is_empty() { None } else { Some(pos_args) };
let has_kwonly = !kwargs.is_empty();
let kwargs = if has_kwonly { Some(kwargs) } else { None };
let bind_mode = if pos_args.is_none()
&& pos_defaults_count == 0
&& arg_defaults_count == 0
&& var_args.is_none()
&& kwargs.is_none()
&& var_kwargs.is_none()
{
BindMode::Simple
} else if pos_args.is_none()
&& var_args.is_none()
&& kwargs.is_none()
&& var_kwargs.is_none()
&& arg_defaults_count > 0
{
BindMode::SimpleWithDefaults
} else {
BindMode::Complex
};
Self {
pos_args,
pos_defaults_count,
args: if args.is_empty() { None } else { Some(args) },
arg_defaults_count,
var_args,
kwargs,
kwarg_default_map: if has_kwonly { Some(kwarg_default_map) } else { None },
var_kwargs,
bind_mode,
}
}
pub fn bind(
&self,
args: ArgValues,
defaults: &[Value],
vm: &mut VM<'_, impl ResourceTracker>,
func_name: Identifier,
namespace: &mut Vec<Value>,
) -> RunResult<()> {
let (pos_iter, keyword_args) = args.into_parts();
let n_kwargs = keyword_args.len();
let namespace_base = namespace.len();
if n_kwargs == 0 && matches!(self.bind_mode, BindMode::Simple | BindMode::SimpleWithDefaults) {
keyword_args.drop_with_heap(vm);
match pos_iter {
ArgPosIter::Empty => {}
ArgPosIter::One(a) => {
namespace.push(a);
}
ArgPosIter::Two([a1, a2]) => {
namespace.push(a1);
namespace.push(a2);
}
ArgPosIter::Vec(args) => {
namespace.extend(args);
}
}
let actual_count = namespace.len() - namespace_base;
let param_count = self.param_count();
if actual_count == param_count {
return Ok(());
} else if self.bind_mode == BindMode::SimpleWithDefaults {
let required = self.required_positional_count();
if actual_count >= required && actual_count < param_count {
let defaults_needed = param_count - actual_count;
let defaults_start = self.arg_defaults_count - defaults_needed;
for i in 0..defaults_needed {
namespace.push(defaults[defaults_start + i].clone_with_heap(vm));
}
return Ok(());
}
}
return self.wrong_arg_count_error(actual_count, vm.interns, func_name);
}
let keyword_args = keyword_args.into_iter();
defer_drop_mut!(keyword_args, vm);
let mut pos_iter_guard = HeapGuard::new(pos_iter, vm);
let (pos_iter, vm) = pos_iter_guard.as_parts_mut();
let pos_param_count = self.pos_arg_count();
let arg_param_count = self.arg_count();
let total_positional_params = pos_param_count + arg_param_count;
let positional_count = pos_iter.len();
let positional_overflow = self.max_positional_count().is_some_and(|max| positional_count > max);
let var_args_offset = usize::from(self.var_args.is_some());
namespace.resize_with(namespace.len() + self.total_slots(), || Value::Undefined);
let mut bound_params: u64 = 0;
for (i, slot) in namespace[namespace_base..].iter_mut().enumerate().take(pos_param_count) {
if let Some(val) = pos_iter.next() {
*slot = val;
bound_params |= 1 << i;
}
}
for (i, slot) in namespace[namespace_base..]
.iter_mut()
.enumerate()
.take(total_positional_params)
.skip(pos_param_count)
{
if let Some(val) = pos_iter.next() {
*slot = val;
bound_params |= 1 << i;
}
}
if self.var_args.is_some() {
namespace[namespace_base + total_positional_params] = allocate_tuple(pos_iter.collect(), vm.heap)?;
}
let mut excess_kwargs_guard = HeapGuard::new(self.var_kwargs.is_some().then(Dict::new), vm);
let (excess_kwargs, vm) = excess_kwargs_guard.as_parts_mut();
'kwargs: for (key, value) in keyword_args {
let mut key_guard = HeapGuard::new(key, vm);
let (key, vm) = key_guard.as_parts_mut();
let mut value_guard = HeapGuard::new(value, vm);
let vm = value_guard.heap();
let Some(keyword_name) = key.as_either_str(vm.heap) else {
return Err(ExcType::type_error("keywords must be strings"));
};
if let Some(pos_args) = &self.pos_args
&& let Some(¶m_id) = pos_args
.iter()
.find(|&¶m_id| keyword_name.matches(param_id, vm.interns))
{
let func = vm.interns.get_str(func_name.name_id);
let param = vm.interns.get_str(param_id);
return Err(ExcType::type_error_positional_only(func, param));
}
if let Some(args) = &self.args {
for (i, ¶m_id) in args.iter().enumerate() {
if keyword_name.matches(param_id, vm.interns) {
let ns_idx = pos_param_count + i;
if (bound_params & (1 << ns_idx)) != 0 {
let func = vm.interns.get_str(func_name.name_id);
let param = vm.interns.get_str(param_id);
return Err(ExcType::type_error_duplicate_arg(func, param));
}
let (value, _) = value_guard.into_parts();
namespace[namespace_base + ns_idx] = value;
bound_params |= 1 << ns_idx;
continue 'kwargs;
}
}
}
if let Some(kwargs) = &self.kwargs {
for (i, ¶m_id) in kwargs.iter().enumerate() {
if keyword_name.matches(param_id, vm.interns) {
let ns_idx = total_positional_params + var_args_offset + i;
let bit_idx = total_positional_params + i;
if (bound_params & (1 << bit_idx)) != 0 {
let func = vm.interns.get_str(func_name.name_id);
let param = vm.interns.get_str(param_id);
return Err(ExcType::type_error_duplicate_arg(func, param));
}
let (value, _) = value_guard.into_parts();
namespace[namespace_base + ns_idx] = value;
bound_params |= 1 << bit_idx;
continue 'kwargs;
}
}
}
if let Some(excess_kwargs) = excess_kwargs {
let (value, _) = value_guard.into_parts();
let (key, vm) = key_guard.into_parts();
excess_kwargs.set(key, value, vm)?;
continue 'kwargs;
}
let func = vm.interns.get_str(func_name.name_id);
let key_str = keyword_name.as_str(vm.interns);
return Err(ExcType::type_error_unexpected_keyword(func, key_str));
}
if positional_overflow {
let kwonly_given = (0..self.kwarg_count())
.filter(|i| (bound_params & (1 << (total_positional_params + i))) != 0)
.count();
let func = vm.interns.get_str(func_name.name_id);
return Err(ExcType::type_error_too_many_positional_range(
func,
self.required_positional_count(),
total_positional_params,
positional_count,
kwonly_given,
));
}
let mut default_idx = 0;
if self.pos_defaults_count > 0 {
let first_optional = pos_param_count - self.pos_defaults_count;
for i in first_optional..pos_param_count {
if (bound_params & (1 << i)) == 0 {
namespace[namespace_base + i] = defaults[default_idx + (i - first_optional)].clone_with_heap(vm);
bound_params |= 1 << i;
}
}
}
default_idx += self.pos_defaults_count;
if self.arg_defaults_count > 0 {
let first_optional = arg_param_count - self.arg_defaults_count;
for i in first_optional..arg_param_count {
let ns_idx = pos_param_count + i;
if (bound_params & (1 << ns_idx)) == 0 {
namespace[namespace_base + ns_idx] =
defaults[default_idx + (i - first_optional)].clone_with_heap(vm);
bound_params |= 1 << ns_idx;
}
}
}
default_idx += self.arg_defaults_count;
if let Some(ref default_map) = self.kwarg_default_map {
for (i, default_slot) in default_map.iter().enumerate() {
if let Some(slot_idx) = default_slot {
let bound_idx = total_positional_params + i;
let ns_idx = total_positional_params + var_args_offset + i;
if (bound_params & (1 << bound_idx)) == 0 {
namespace[namespace_base + ns_idx] = defaults[default_idx + slot_idx].clone_with_heap(vm);
bound_params |= 1 << bound_idx;
}
}
}
}
let mut missing_positional: Vec<&str> = Vec::new();
if let Some(ref pos_args) = self.pos_args {
let required_pos_only = pos_args.len().saturating_sub(self.pos_defaults_count);
for (i, ¶m_id) in pos_args.iter().enumerate() {
if i < required_pos_only && (bound_params & (1 << i)) == 0 {
missing_positional.push(vm.interns.get_str(param_id));
}
}
}
if let Some(ref args_params) = self.args {
let required_args = args_params.len().saturating_sub(self.arg_defaults_count);
for (i, ¶m_id) in args_params.iter().enumerate() {
if i < required_args && (bound_params & (1 << (pos_param_count + i))) == 0 {
missing_positional.push(vm.interns.get_str(param_id));
}
}
}
if !missing_positional.is_empty() {
let func = vm.interns.get_str(func_name.name_id);
return Err(ExcType::type_error_missing_positional_with_names(
func,
&missing_positional,
));
}
let mut missing_kwonly: Vec<&str> = Vec::new();
if let Some(ref kwargs_params) = self.kwargs {
let default_map = self.kwarg_default_map.as_ref();
for (i, ¶m_id) in kwargs_params.iter().enumerate() {
let has_default = default_map.and_then(|map| map.get(i)).is_some_and(Option::is_some);
if !has_default && (bound_params & (1 << (total_positional_params + i))) == 0 {
missing_kwonly.push(vm.interns.get_str(param_id));
}
}
}
if !missing_kwonly.is_empty() {
let func = vm.interns.get_str(func_name.name_id);
return Err(ExcType::type_error_missing_kwonly_with_names(func, &missing_kwonly));
}
let (excess_kwargs, vm) = excess_kwargs_guard.into_parts();
if let Some(excess_kwargs) = excess_kwargs {
let dict_id = vm.heap.allocate(HeapData::Dict(excess_kwargs))?;
let last_slot = namespace.len() - 1;
namespace[last_slot] = Value::Ref(dict_id);
}
Ok(())
}
pub fn param_count(&self) -> usize {
self.pos_arg_count() + self.arg_count() + self.kwarg_count()
}
pub fn total_slots(&self) -> usize {
let mut slots = self.param_count();
if self.var_args.is_some() {
slots += 1;
}
if self.var_kwargs.is_some() {
slots += 1;
}
slots
}
pub fn total_defaults_count(&self) -> usize {
self.pos_defaults_count + self.arg_defaults_count + self.kwarg_defaults_count()
}
#[inline]
fn required_positional_count(&self) -> usize {
self.pos_arg_count() + self.arg_count() - self.pos_defaults_count - self.arg_defaults_count
}
fn kwarg_defaults_count(&self) -> usize {
self.kwarg_default_map
.as_deref()
.map(|v| v.iter().filter(|&x| x.is_some()).count())
.unwrap_or_default()
}
fn pos_arg_count(&self) -> usize {
self.pos_args.as_ref().map_or(0, Vec::len)
}
fn arg_count(&self) -> usize {
self.args.as_ref().map_or(0, Vec::len)
}
fn kwarg_count(&self) -> usize {
self.kwargs.as_ref().map_or(0, Vec::len)
}
fn param_names(&self) -> impl Iterator<Item = StringId> + '_ {
let pos_args = self.pos_args.iter().flat_map(|v| v.iter().copied());
let args = self.args.iter().flat_map(|v| v.iter().copied());
let var_args = self.var_args.iter().copied();
let kwargs = self.kwargs.iter().flat_map(|v| v.iter().copied());
let var_kwargs = self.var_kwargs.iter().copied();
pos_args.chain(args).chain(var_args).chain(kwargs).chain(var_kwargs)
}
fn max_positional_count(&self) -> Option<usize> {
if self.var_args.is_some() {
None
} else {
Some(self.pos_arg_count() + self.arg_count())
}
}
fn wrong_arg_count_error<T>(&self, actual_count: usize, interns: &Interns, func_name: Identifier) -> RunResult<T> {
let name_str = interns.get_str(func_name.name_id);
let param_count = self.param_count();
let required = self.required_positional_count();
let msg = if actual_count < required {
let missing_names: Vec<String> = self
.param_names()
.take(required)
.skip(actual_count)
.map(|string_id| interns.get_str(string_id).to_string())
.collect();
let missing_refs: Vec<&str> = missing_names.iter().map(String::as_str).collect();
ExcType::missing_positional_msg(name_str, &missing_refs)
} else {
ExcType::too_many_positional_range_msg(name_str, required, param_count, actual_count, 0)
};
Err(SimpleException::new_msg(ExcType::TypeError, msg)
.with_position(func_name.position)
.into())
}
}