1use crate::common::lock::LazyLock;
2use crate::{
3 AsObject, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine, atomic_func,
4 builtins::{
5 PyBaseExceptionRef, PyDict, PyStr, PyStrRef, PyTuple, PyTupleRef, PyType, PyTypeRef,
6 },
7 class::{PyClassImpl, StaticType, class_attr_item_doc},
8 function::{
9 Either, FuncArgs, KwArgs, NameChanges, OptionalArg, PyComparisonValue, PyMethodDef,
10 PyMethodFlags,
11 },
12 iter::PyExactSizeIterator,
13 protocol::{PyMappingMethods, PySequenceMethods},
14 sliceable::{SequenceIndex, SliceableSequenceOp},
15 types::PyComparisonOp,
16 vm::Context,
17};
18
19const DEFAULT_STRUCTSEQ_REDUCE: PyMethodDef = PyMethodDef::new_const(
20 "__reduce__",
21 |zelf: PyRef<PyTuple>, vm: &VirtualMachine| -> PyTupleRef {
22 vm.new_tuple((
23 zelf.class().to_owned(),
24 (vm.ctx.new_tuple(zelf.as_slice().to_vec()),),
25 ))
26 },
27 PyMethodFlags::METHOD,
28 crate::function::ItemDoc::static_text("__reduce__($self, /)\n--\n\n"),
29);
30
31pub const STRUCT_SEQUENCE_PARAMS: Option<&'static [crate::function::Param]> =
33 Some(&[crate::function::Param {
34 name: "iterable",
35 kind: crate::function::ParamKind::PositionalOnly,
36 default: Some(crate::function::DefaultRepr::Raw("()")),
37 }]);
38
39#[derive(FromArgs)]
41pub struct StructSequenceNewArgs {
42 #[pyarg(any)]
43 pub sequence: PyObjectRef,
44 #[pyarg(any, optional, py_default = "{}")]
45 pub dict: OptionalArg<PyObjectRef>,
46}
47
48pub fn struct_sequence_new(
58 cls: PyTypeRef,
59 args: StructSequenceNewArgs,
60 hidden_field_names: &[&str],
61 vm: &VirtualMachine,
62) -> PyResult {
63 let StructSequenceNewArgs {
65 sequence: seq,
66 dict,
67 } = args;
68
69 #[cold]
70 fn length_error(
71 tp_name: &str,
72 min_len: usize,
73 max_len: usize,
74 len: usize,
75 vm: &VirtualMachine,
76 ) -> PyBaseExceptionRef {
77 if min_len == max_len {
78 vm.new_type_error(format!(
79 "{tp_name}() takes a {min_len}-sequence ({len}-sequence given)"
80 ))
81 } else if len < min_len {
82 vm.new_type_error(format!(
83 "{tp_name}() takes an at least {min_len}-sequence ({len}-sequence given)"
84 ))
85 } else {
86 vm.new_type_error(format!(
87 "{tp_name}() takes an at most {max_len}-sequence ({len}-sequence given)"
88 ))
89 }
90 }
91
92 let min_len: usize = cls
93 .get_attr(identifier!(vm.ctx, n_sequence_fields))
94 .ok_or_else(|| vm.new_type_error("missing n_sequence_fields attribute"))?
95 .try_into_value(vm)?;
96 let max_len: usize = cls
97 .get_attr(identifier!(vm.ctx, n_fields))
98 .ok_or_else(|| vm.new_type_error("missing n_fields attribute"))?
99 .try_into_value(vm)?;
100
101 let dict = match dict {
102 OptionalArg::Missing => None,
103 OptionalArg::Present(dict) => Some(dict.downcast::<PyDict>().map_err(|_| {
104 vm.new_type_error(format!(
105 "{}() takes a dict as second arg, if any",
106 cls.slot_name()
107 ))
108 })?),
109 };
110
111 let seq: Vec<PyObjectRef> = seq.try_into_value(vm)?;
112 let len = seq.len();
113
114 if len < min_len || len > max_len {
115 return Err(length_error(&cls.slot_name(), min_len, max_len, len, vm));
116 }
117
118 let mut items = seq;
120 items.resize_with(max_len, || vm.ctx.none());
121
122 if let Some(dict) = dict.filter(|dict| !dict.is_empty()) {
126 let mut found = 0;
127 let names = hidden_field_names.get(len - min_len..).unwrap_or(&[]);
128 for (item, name) in items[len..].iter_mut().zip(names) {
129 if let Some(value) = dict.get_item_opt(*name, vm)? {
130 *item = value;
131 found += 1;
132 }
133 }
134 if found != dict.__len__() {
135 return Err(vm.new_type_error(format!(
136 "{}() got duplicate or unexpected field name(s)",
137 cls.slot_name()
138 )));
139 }
140 }
141
142 PyTuple::new_unchecked(items.into_boxed_slice())
143 .into_ref_with_type(vm, cls)
144 .map(Into::into)
145}
146
147fn get_visible_len(obj: &PyObject, vm: &VirtualMachine) -> PyResult<usize> {
148 obj.class()
149 .get_attr(identifier!(vm.ctx, n_sequence_fields))
150 .ok_or_else(|| vm.new_type_error("missing n_sequence_fields"))?
151 .try_into_value(vm)
152}
153
154static STRUCT_SEQUENCE_AS_SEQUENCE: LazyLock<PySequenceMethods> =
157 LazyLock::new(|| PySequenceMethods {
158 length: atomic_func!(|seq, vm| get_visible_len(seq.obj, vm)),
159 concat: atomic_func!(|seq, other, vm| {
160 let n_seq = get_visible_len(seq.obj, vm)?;
162 let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
163 let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
164 let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
165 visible_tuple
167 .as_object()
168 .sequence_unchecked()
169 .concat(other, vm)
170 }),
171 repeat: atomic_func!(|seq, n, vm| {
172 let n_seq = get_visible_len(seq.obj, vm)?;
174 let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
175 let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
176 let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
177 visible_tuple.as_object().sequence_unchecked().repeat(n, vm)
179 }),
180 item: atomic_func!(|seq, i, vm| {
181 let n_seq = get_visible_len(seq.obj, vm)?;
182 let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
183 let idx = if i < 0 {
184 let pos_i = n_seq as isize + i;
185 if pos_i < 0 {
186 return Err(vm.new_index_error("tuple index out of range"));
187 }
188 pos_i as usize
189 } else {
190 i as usize
191 };
192 if idx >= n_seq {
193 return Err(vm.new_index_error("tuple index out of range"));
194 }
195 Ok(tuple.as_slice()[idx].clone())
196 }),
197 contains: atomic_func!(|seq, needle, vm| {
198 let n_seq = get_visible_len(seq.obj, vm)?;
199 let tuple = seq.obj.downcast_ref::<PyTuple>().unwrap();
200 for item in tuple.as_slice().iter().take(n_seq) {
201 if item.rich_compare_bool(needle, PyComparisonOp::Eq, vm)? {
202 return Ok(true);
203 }
204 }
205 Ok(false)
206 }),
207 ..PySequenceMethods::NOT_IMPLEMENTED
208 });
209
210static STRUCT_SEQUENCE_AS_MAPPING: LazyLock<PyMappingMethods> =
213 LazyLock::new(|| PyMappingMethods {
214 length: atomic_func!(|mapping, vm| get_visible_len(mapping.obj, vm)),
215 subscript: atomic_func!(|mapping, needle, vm| {
216 let n_seq = get_visible_len(mapping.obj, vm)?;
217 let tuple = mapping.obj.downcast_ref::<PyTuple>().unwrap();
218 let visible_elements = &tuple.as_slice()[..n_seq];
219
220 match SequenceIndex::try_from_borrowed_object(vm, needle, "tuple")? {
221 SequenceIndex::Int(i) => visible_elements.getitem_by_index(vm, i),
222 SequenceIndex::Slice(slice) => visible_elements
223 .getitem_by_slice(vm, slice)
224 .map(|x| vm.ctx.new_tuple(x).into()),
225 }
226 }),
227 ..PyMappingMethods::NOT_IMPLEMENTED
228 });
229
230pub trait PyStructSequenceData: Sized {
235 const REQUIRED_FIELD_NAMES: &'static [&'static str];
237
238 const OPTIONAL_FIELD_NAMES: &'static [&'static str];
240
241 const UNNAMED_FIELDS_LEN: usize = 0;
243
244 fn into_tuple(self, vm: &VirtualMachine) -> PyTuple;
246
247 fn try_from_elements(_elements: Vec<PyObjectRef>, vm: &VirtualMachine) -> PyResult<Self> {
251 Err(vm.new_type_error("This struct sequence does not support construction from elements"))
252 }
253}
254
255#[pyclass]
260pub trait PyStructSequence: StaticType + PyClassImpl + Sized + 'static {
261 type Data: PyStructSequenceData;
263
264 #[pyslot]
265 fn slot_new(cls: PyTypeRef, args: FuncArgs, vm: &VirtualMachine) -> PyResult {
266 struct_sequence_new(
267 cls,
268 args.bind_for(vm, Self::NAME)?,
269 Self::Data::OPTIONAL_FIELD_NAMES,
270 vm,
271 )
272 }
273
274 fn from_data(data: Self::Data, vm: &VirtualMachine) -> PyTupleRef {
276 let tuple =
277 <Self::Data as ::rustpython_vm::types::PyStructSequenceData>::into_tuple(data, vm);
278 let typ = Self::static_type();
279 tuple
280 .into_ref_with_type(vm, typ.to_owned())
281 .expect("Every PyStructSequence must be a valid tuple. This is a RustPython bug.")
282 }
283
284 #[pyslot]
285 fn slot_repr(zelf: &PyObject, vm: &VirtualMachine) -> PyResult<PyStrRef> {
286 let zelf = zelf
287 .downcast_ref::<PyTuple>()
288 .ok_or_else(|| vm.new_type_error("unexpected payload for __repr__"))?;
289
290 let field_names = Self::Data::REQUIRED_FIELD_NAMES;
291 let format_field = |(value, name): (&PyObject, _)| {
292 let s = value.repr(vm)?;
293 Ok(format!("{name}={s}"))
294 };
295 let (body, suffix) =
296 if let Some(_guard) = rustpython_vm::recursion::ReprGuard::enter(vm, zelf.as_ref()) {
297 let fields: PyResult<Vec<_>> = zelf
298 .as_slice()
299 .iter()
300 .map(|value| value.as_ref())
301 .zip(field_names.iter().copied())
302 .map(format_field)
303 .collect();
304 (fields?.join(", "), "")
305 } else {
306 (String::new(), "...")
307 };
308 let type_name = if Self::MODULE_NAME.is_some() {
311 alloc::borrow::Cow::Borrowed(Self::TP_NAME)
312 } else {
313 let typ = zelf.class();
314 match typ.get_attr(identifier!(vm.ctx, __module__)) {
315 Some(module) if module.downcastable::<PyStr>() => {
316 let module_str = module.downcast_ref::<PyStr>().unwrap();
317 alloc::borrow::Cow::Owned(format!("{}.{}", module_str.as_wtf8(), Self::NAME))
318 }
319 _ => alloc::borrow::Cow::Borrowed(Self::TP_NAME),
320 }
321 };
322 let repr_str = format!("{type_name}({body}{suffix})");
323 Ok(vm.ctx.new_str(repr_str))
324 }
325
326 #[pymethod]
328 fn __replace__(
329 zelf: PyRef<PyTuple>,
330 changes: KwArgs<PyObjectRef, NameChanges>,
331 vm: &VirtualMachine,
332 ) -> PyResult {
333 if Self::Data::UNNAMED_FIELDS_LEN > 0 {
334 return Err(vm.new_type_error(format!(
335 "__replace__() is not supported for {} because it has unnamed field(s)",
336 zelf.class().slot_name()
337 )));
338 }
339
340 let n_fields =
341 Self::Data::REQUIRED_FIELD_NAMES.len() + Self::Data::OPTIONAL_FIELD_NAMES.len();
342 let mut items: Vec<PyObjectRef> = zelf.as_slice()[..n_fields].to_vec();
343
344 let mut kwargs = changes;
345
346 let all_field_names: Vec<&str> = Self::Data::REQUIRED_FIELD_NAMES
348 .iter()
349 .chain(Self::Data::OPTIONAL_FIELD_NAMES.iter())
350 .copied()
351 .collect();
352 for (i, &name) in all_field_names.iter().enumerate() {
353 if let Some(val) = kwargs.shift_remove(name) {
354 items[i] = val;
355 }
356 }
357
358 if !kwargs.is_empty() {
359 let names = vm.ctx.new_list(
360 kwargs
361 .keys()
362 .map(|k| vm.ctx.new_str(k.to_owned()).into())
363 .collect(),
364 );
365 let names_repr = names.as_object().repr(vm)?;
366 return Err(vm.new_type_error(format!("Got unexpected field name(s): {names_repr}")));
367 }
368
369 PyTuple::new_unchecked(items.into_boxed_slice())
370 .into_ref_with_type(vm, zelf.class().to_owned())
371 .map(Into::into)
372 }
373
374 #[pymethod]
375 fn __getitem__(zelf: PyRef<PyTuple>, needle: PyObjectRef, vm: &VirtualMachine) -> PyResult {
376 let n_seq = get_visible_len(zelf.as_ref(), vm)?;
377 let visible_elements = &zelf.as_slice()[..n_seq];
378
379 match SequenceIndex::try_from_borrowed_object(vm, &needle, "tuple")? {
380 SequenceIndex::Int(i) => visible_elements.getitem_by_index(vm, i),
381 SequenceIndex::Slice(slice) => visible_elements
382 .getitem_by_slice(vm, slice)
383 .map(|x| vm.ctx.new_tuple(x).into()),
384 }
385 }
386
387 #[extend_class]
388 fn extend_pyclass(ctx: &Context, class: &'static Py<PyType>) {
389 for (i, &name) in Self::Data::REQUIRED_FIELD_NAMES.iter().enumerate() {
391 class.set_attr(
392 ctx.intern_str(name),
393 ctx.new_readonly_tuple_member(name, class, i, class_attr_item_doc::<Self>(name))
394 .into(),
395 );
396 }
397
398 let visible_count = Self::Data::REQUIRED_FIELD_NAMES.len() + Self::Data::UNNAMED_FIELDS_LEN;
400 for (i, &name) in Self::Data::OPTIONAL_FIELD_NAMES.iter().enumerate() {
401 class.set_attr(
402 ctx.intern_str(name),
403 ctx.new_readonly_tuple_member(
404 name,
405 class,
406 visible_count + i,
407 class_attr_item_doc::<Self>(name),
408 )
409 .into(),
410 );
411 }
412
413 class.set_attr(
414 identifier!(ctx, __match_args__),
415 ctx.new_tuple(
416 Self::Data::REQUIRED_FIELD_NAMES
417 .iter()
418 .map(|&name| ctx.new_str(name).into())
419 .collect::<Vec<_>>(),
420 )
421 .into(),
422 );
423
424 let n_unnamed_fields = Self::Data::UNNAMED_FIELDS_LEN;
429 let n_sequence_fields = Self::Data::REQUIRED_FIELD_NAMES.len() + n_unnamed_fields;
430 let n_fields = n_sequence_fields + Self::Data::OPTIONAL_FIELD_NAMES.len();
431 class.set_attr(
432 identifier!(ctx, n_sequence_fields),
433 ctx.new_int(n_sequence_fields).into(),
434 );
435 class.set_attr(identifier!(ctx, n_fields), ctx.new_int(n_fields).into());
436 class.set_attr(
437 identifier!(ctx, n_unnamed_fields),
438 ctx.new_int(n_unnamed_fields).into(),
439 );
440
441 class
443 .slots
444 .as_sequence
445 .copy_from(&STRUCT_SEQUENCE_AS_SEQUENCE);
446 class
447 .slots
448 .as_mapping
449 .copy_from(&STRUCT_SEQUENCE_AS_MAPPING);
450
451 class.slots.iter.store(Some(struct_sequence_iter));
453
454 class.slots.hash.store(Some(struct_sequence_hash));
456
457 class
459 .slots
460 .richcompare
461 .store(Some(struct_sequence_richcompare));
462
463 if !class.attributes.contains(ctx.intern_str("__reduce__")) {
467 class.set_attr(
468 ctx.intern_str("__reduce__"),
469 DEFAULT_STRUCTSEQ_REDUCE.to_proper_method(class, ctx),
470 );
471 }
472 }
473}
474
475fn struct_sequence_iter(zelf: PyObjectRef, vm: &VirtualMachine) -> PyResult {
477 let tuple = zelf
478 .downcast_ref::<PyTuple>()
479 .ok_or_else(|| vm.new_type_error("expected tuple"))?;
480 let n_seq = get_visible_len(&zelf, vm)?;
481 let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
482 let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
483 visible_tuple
484 .as_object()
485 .to_owned()
486 .get_iter(vm)
487 .map(Into::into)
488}
489
490fn struct_sequence_hash(
492 zelf: &PyObject,
493 vm: &VirtualMachine,
494) -> PyResult<crate::common::hash::PyHash> {
495 let tuple = zelf
496 .downcast_ref::<PyTuple>()
497 .ok_or_else(|| vm.new_type_error("expected tuple"))?;
498 let n_seq = get_visible_len(zelf, vm)?;
499 let visible: Vec<_> = tuple.as_slice().iter().take(n_seq).cloned().collect();
501 let visible_tuple = PyTuple::new_ref(visible, &vm.ctx);
502 visible_tuple.as_object().hash(vm)
503}
504
505fn struct_sequence_richcompare(
507 zelf: &PyObject,
508 other: &PyObject,
509 op: PyComparisonOp,
510 vm: &VirtualMachine,
511) -> PyResult<Either<PyObjectRef, PyComparisonValue>> {
512 let zelf_tuple = zelf
513 .downcast_ref::<PyTuple>()
514 .ok_or_else(|| vm.new_type_error("expected tuple"))?;
515
516 let Some(other_tuple) = other.downcast_ref::<PyTuple>() else {
518 return Ok(Either::B(PyComparisonValue::NotImplemented));
519 };
520
521 let zelf_len = get_visible_len(zelf, vm)?;
522 let other_len = get_visible_len(other, vm).unwrap_or(other_tuple.as_slice().len());
524
525 let zelf_visible = &zelf_tuple.as_slice()[..zelf_len];
526 let other_visible = &other_tuple.as_slice()[..other_len];
527
528 zelf_visible
530 .iter()
531 .map(|o| &**o)
532 .richcompare(other_visible.iter().map(|o| &**o), op, vm)
533 .map(|v| Either::B(PyComparisonValue::Implemented(v)))
534}