use std::{mem, ptr};
use crate::{
ResourceTracker,
args::{ArgPosIter, ArgValues, KwargsValues, KwargsValuesIter},
bytecode::VM,
exception_private::{ExcType, RunError, RunResult},
heap::{ContainsHeap, DropWithHeap, HeapGuard},
intern::{Interns, StringId},
value::{EitherStr, Value},
};
#[inline]
pub(crate) fn bind<const N: usize>(
spec: &'static ParamSpec,
bound: &mut Bound<N>,
args: ArgValues,
vm: &mut VM<'_, impl ResourceTracker>,
) -> RunResult<()> {
debug_assert!(ptr::eq(spec, bound.spec), "spec must match the bound's spec");
let args = match args {
ArgValues::Empty if spec.n_required_positional == 0 => {
return Ok(());
}
ArgValues::One(v) if spec.n_required_positional <= 1 && spec.n_positional >= 1 => {
bound.slots[0] = Some(v);
return Ok(());
}
ArgValues::Two(v1, v2) if spec.n_required_positional <= 2 && spec.n_positional >= 2 => {
bound.slots[0] = Some(v1);
bound.slots[1] = Some(v2);
return Ok(());
}
other => other,
};
bind_slow(spec, bound, args, vm)
}
fn bind_slow<const N: usize>(
spec: &'static ParamSpec,
bound: &mut Bound<N>,
args: ArgValues,
vm: &mut VM<'_, impl ResourceTracker>,
) -> RunResult<()> {
let (pos, kwargs) = args.into_parts();
let state = IterState {
pos,
kwargs: kwargs.into_iter(),
};
let mut guard = HeapGuard::new(state, vm);
let (state, vm) = guard.as_parts_mut();
if spec.kwargs_not_supported_yet && state.kwargs.len() > 0 {
return Err(ExcType::kwargs_not_implemented(spec.func_name));
}
let n_pos = state.pos.len();
let n_kw = state.kwargs.len();
if matches!(spec.family, ErrorFamily::Unpack)
&& let Some(err) = unpack_arity_error(spec, n_pos)
{
return Err(err);
}
if spec.at_most_total && n_pos + n_kw > spec.n_positional {
return Err(total_overflow_error(spec, n_pos + n_kw));
}
if spec.uses_c_method_arity() && n_pos < spec.n_required_pos_only {
return Err(ExcType::type_error_at_least_positional(
spec.func_name,
spec.n_required_pos_only,
n_pos,
));
}
let positional_overflow = n_pos > spec.n_positional && !spec.varargs;
if positional_overflow && !matches!(spec.family, ErrorFamily::Def) {
return Err(positional_overflow_error(spec, n_pos, n_kw));
}
for slot in bound.slots.iter_mut().take(n_pos.min(spec.n_positional)) {
*slot = state.pos.next();
}
if spec.varargs {
bound.varargs.extend(state.pos.by_ref());
}
for (key, value) in state.kwargs.by_ref() {
let Some(key_str) = key.as_either_str(vm.heap) else {
(key, value).drop_with_heap(vm);
return Err(ExcType::type_error_kwargs_nonstring_key());
};
match find_param(spec, &key_str, vm.interns) {
Some((_, param)) if matches!(param.kind, ParamKind::PosOnly) => {
(key, value).drop_with_heap(vm);
return Err(ExcType::type_error_positional_only(spec.func_name, param.name));
}
Some((idx, param)) => {
key.drop_with_heap(vm);
if bound.slots[idx].is_some() {
value.drop_with_heap(vm);
match duplicate_error(spec, idx, param) {
DuplicateOutcome::Raise(err) => return Err(err),
DuplicateOutcome::Defer(err) => {
let deferred = bound.deferred_mut();
if deferred.conflict.as_ref().is_none_or(|(ci, _)| idx < *ci) {
deferred.conflict = Some((idx, err));
}
}
}
} else {
bound.slots[idx] = Some(value);
}
}
None if spec.varkwargs => {
bound.varkwargs.push((key, value));
}
None if spec.family.defers_unknown_kwarg() => {
value.drop_with_heap(vm);
if bound.deferred.as_ref().is_none_or(|d| d.unknown.is_none()) {
let err = match spec.family {
ErrorFamily::C { .. } => ExcType::type_error_c_unexpected_keyword(key_str.as_str(vm.interns)),
_ => ExcType::type_error_unexpected_keyword(spec.func_name, key_str.as_str(vm.interns)),
};
bound.deferred_mut().unknown = Some(err);
}
key.drop_with_heap(vm);
}
None => {
value.drop_with_heap(vm);
let name = key_str.as_str(vm.interns).to_owned();
key.drop_with_heap(vm);
let err_name = spec.kwarg_error_name.unwrap_or(spec.func_name);
return Err(ExcType::type_error_unexpected_keyword(err_name, &name));
}
}
}
if positional_overflow {
let kwonly_given = bound.slots[spec.n_positional..].iter().filter(|s| s.is_some()).count();
return Err(ExcType::type_error_too_many_positional_range(
spec.func_name,
spec.n_required_positional,
spec.n_positional,
n_pos,
kwonly_given,
));
}
if spec.family.aggregates_missing() {
let missing = missing_names(spec, &bound.slots, 0, spec.n_positional);
if !missing.is_empty() {
return Err(ExcType::type_error_missing_positional_with_names(
spec.func_name,
&missing,
));
}
let missing = missing_names(spec, &bound.slots, spec.n_positional, N);
if !missing.is_empty() {
return Err(ExcType::type_error_missing_kwonly_with_names(spec.func_name, &missing));
}
}
Ok(())
}
#[expect(clippy::struct_excessive_bools, reason = "mirrors independent signature axes")]
pub(crate) struct ParamSpec {
pub func_name: &'static str,
pub family: ErrorFamily,
pub params: &'static [Param],
pub n_positional: usize,
pub n_required_positional: usize,
pub n_required_pos_only: usize,
pub varargs: bool,
pub varkwargs: bool,
pub at_most_total: bool,
pub kwargs_not_supported_yet: bool,
pub kwarg_error_name: Option<&'static str>,
}
impl ParamSpec {
fn uses_c_method_arity(&self) -> bool {
self.n_required_pos_only > 0 && !matches!(self.family, ErrorFamily::Def | ErrorFamily::Unpack)
}
}
pub(crate) struct Param {
pub name: &'static str,
pub kwarg_id: Option<StringId>,
pub kind: ParamKind,
pub required: bool,
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum ParamKind {
PosOnly,
PosOrKeyword,
KwOnly,
}
#[derive(Clone, Copy)]
pub(crate) enum ErrorFamily {
Def,
Clinic,
C { positional_pivot: bool },
CNamed,
Unpack,
}
impl ErrorFamily {
fn defers_unknown_kwarg(self) -> bool {
matches!(self, Self::C { .. } | Self::CNamed)
}
fn aggregates_missing(self) -> bool {
!matches!(self, Self::C { .. } | Self::CNamed)
}
}
pub(crate) struct Bound<const N: usize> {
spec: &'static ParamSpec,
slots: [Option<Value>; N],
varargs: Vec<Value>,
varkwargs: Vec<(Value, Value)>,
deferred: Option<Box<DeferredLeftovers>>,
}
#[derive(Default)]
struct DeferredLeftovers {
conflict: Option<(usize, RunError)>,
unknown: Option<RunError>,
}
impl<const N: usize> Bound<N> {
pub(crate) fn new(spec: &'static ParamSpec) -> Self {
debug_assert_eq!(N, spec.params.len(), "slot count must match spec params");
Self {
spec,
slots: [const { None }; N],
varargs: Vec::new(),
varkwargs: Vec::new(),
deferred: None,
}
}
fn deferred_mut(&mut self) -> &mut DeferredLeftovers {
self.deferred.get_or_insert_default()
}
pub(crate) fn require(&mut self, i: usize) -> RunResult<Value> {
match self.slots[i].take() {
Some(v) => Ok(v),
None => Err(self.missing_error(i)),
}
}
pub(crate) fn take(&mut self, i: usize) -> Option<Value> {
self.slots[i].take()
}
pub(crate) fn take_varargs(&mut self) -> Vec<Value> {
mem::take(&mut self.varargs)
}
pub(crate) fn take_varkwargs(&mut self) -> KwargsValues {
if self.varkwargs.is_empty() {
KwargsValues::Empty
} else {
KwargsValues::Pairs(mem::take(&mut self.varkwargs))
}
}
pub(crate) fn finish(&mut self) -> RunResult<()> {
match self.deferred.take() {
None => Ok(()),
Some(deferred) => match *deferred {
DeferredLeftovers {
conflict: Some((_, err)),
..
} => Err(err),
DeferredLeftovers { unknown: Some(err), .. } => Err(err),
DeferredLeftovers {
conflict: None,
unknown: None,
} => Ok(()),
},
}
}
#[cold]
fn missing_error(&self, i: usize) -> RunError {
let param = &self.spec.params[i];
match self.spec.family {
ErrorFamily::C { .. } if i < self.spec.n_positional => {
ExcType::type_error_c_missing_required(param.name, i + 1)
}
ErrorFamily::CNamed if i < self.spec.n_positional => {
ExcType::type_error_c_missing_required_named(self.spec.func_name, param.name, i + 1)
}
ErrorFamily::C { .. } | ErrorFamily::CNamed => {
ExcType::type_error_missing_kwonly_with_names(self.spec.func_name, &[param.name])
}
_ => ExcType::type_error_missing_positional_with_names(self.spec.func_name, &[param.name]),
}
}
}
impl<const N: usize> DropWithHeap for Bound<N> {
fn drop_with_heap<H: ContainsHeap>(self, heap: &mut H) {
for slot in self.slots {
slot.drop_with_heap(heap);
}
self.varargs.drop_with_heap(heap);
for (k, v) in self.varkwargs {
k.drop_with_heap(heap);
v.drop_with_heap(heap);
}
}
}
struct IterState {
pos: ArgPosIter,
kwargs: KwargsValuesIter,
}
impl DropWithHeap for IterState {
fn drop_with_heap<H: ContainsHeap>(self, heap: &mut H) {
self.pos.drop_with_heap(heap);
self.kwargs.drop_with_heap(heap);
}
}
fn find_param<'s>(spec: &'s ParamSpec, key: &EitherStr, interns: &Interns) -> Option<(usize, &'s Param)> {
spec.params
.iter()
.enumerate()
.find(|(_, p)| p.kwarg_id.is_some_and(|id| key.matches(id, interns)))
}
enum DuplicateOutcome {
Raise(RunError),
Defer(RunError),
}
#[cold]
fn duplicate_error(spec: &ParamSpec, idx: usize, param: &Param) -> DuplicateOutcome {
if matches!(param.kind, ParamKind::KwOnly) {
return DuplicateOutcome::Raise(ExcType::type_error_multiple_values(spec.func_name, param.name));
}
let named_conflict =
|| ExcType::type_error_positional_keyword_conflict(&format!("{}()", spec.func_name), param.name, idx + 1);
match spec.family {
ErrorFamily::C { .. } => DuplicateOutcome::Defer(ExcType::type_error_positional_keyword_conflict(
"function",
param.name,
idx + 1,
)),
ErrorFamily::CNamed => DuplicateOutcome::Defer(named_conflict()),
ErrorFamily::Clinic => DuplicateOutcome::Raise(named_conflict()),
ErrorFamily::Def | ErrorFamily::Unpack => {
DuplicateOutcome::Raise(ExcType::type_error_duplicate_arg(spec.func_name, param.name))
}
}
}
#[cold]
fn unpack_arity_error(spec: &ParamSpec, n_pos: usize) -> Option<RunError> {
let (min, max) = (spec.n_required_positional, spec.n_positional);
if n_pos < min {
Some(if min == max {
ExcType::type_error_expected_exact(spec.func_name, min, n_pos)
} else {
ExcType::type_error_at_least(spec.func_name, min, n_pos)
})
} else if n_pos > max {
Some(if min == max {
ExcType::type_error_expected_exact(spec.func_name, max, n_pos)
} else {
ExcType::type_error_at_most(spec.func_name, max, n_pos)
})
} else {
None
}
}
#[cold]
fn total_overflow_error(spec: &ParamSpec, total: usize) -> RunError {
match spec.family {
ErrorFamily::C {
positional_pivot: false,
} => ExcType::type_error_c_at_most(spec.n_positional, total),
ErrorFamily::C { positional_pivot: true } => ExcType::type_error_c_at_most_positional(spec.n_positional, total),
_ => ExcType::type_error_method_at_most(spec.func_name, spec.n_positional, total),
}
}
#[cold]
fn positional_overflow_error(spec: &ParamSpec, n_pos: usize, n_kw: usize) -> RunError {
let max = spec.n_positional;
match spec.family {
ErrorFamily::Def | ErrorFamily::Unpack => ExcType::type_error_at_most(spec.func_name, max, n_pos),
_ if spec.uses_c_method_arity() => ExcType::type_error_method_at_most(spec.func_name, max, n_pos + n_kw),
ErrorFamily::C {
positional_pivot: false,
} => ExcType::type_error_c_at_most(max, n_pos + n_kw),
ErrorFamily::C { positional_pivot: true } => {
ExcType::type_error_c_at_most_positional_or_total(max, spec.params.len(), n_pos + n_kw)
}
ErrorFamily::CNamed => ExcType::type_error_method_at_most(spec.func_name, max, n_pos + n_kw),
ErrorFamily::Clinic => ExcType::type_error_at_most(spec.func_name, max, n_pos + n_kw),
}
}
fn missing_names(spec: &ParamSpec, slots: &[Option<Value>], start: usize, end: usize) -> Vec<&'static str> {
spec.params[start..end]
.iter()
.zip(&slots[start..end])
.filter(|(p, slot)| p.required && slot.is_none())
.map(|(p, _)| p.name)
.collect()
}