Skip to main content

rustpython_vm/builtins/
template.rs

1use super::{
2    PyStr, PyTupleRef, PyType, PyTypeRef, genericalias::PyGenericAlias,
3    interpolation::PyInterpolation,
4};
5use crate::{
6    AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
7    atomic_func,
8    class::{PyClassDef, PyClassImpl},
9    common::lock::LazyLock,
10    function::{FuncArgs, PyComparisonValue},
11    protocol::{PyIterReturn, PySequenceMethods},
12    types::{
13        AsSequence, Comparable, Constructor, IterNext, Iterable, PyComparisonOp, Representable,
14        SelfIter,
15    },
16};
17use rustpython_common::wtf8::{Wtf8Buf, wtf8_concat};
18
19#[pyclass(module = "string.templatelib", name = "Template")]
20#[derive(Debug, Clone)]
21pub struct PyTemplate {
22    #[pymember(type = "object_ex")]
23    pub strings: PyTupleRef,
24    #[pymember(type = "object_ex")]
25    pub interpolations: PyTupleRef,
26}
27
28impl PyPayload for PyTemplate {
29    #[inline]
30    fn class(ctx: &Context) -> &'static Py<PyType> {
31        ctx.types.template_type
32    }
33}
34
35impl PyTemplate {
36    #[must_use]
37    pub fn new(strings: PyTupleRef, interpolations: PyTupleRef) -> Self {
38        Self {
39            strings,
40            interpolations,
41        }
42    }
43}
44
45impl Constructor for PyTemplate {
46    type Args = FuncArgs;
47
48    fn py_new(_cls: &Py<PyType>, args: Self::Args, vm: &VirtualMachine) -> PyResult<Self> {
49        if !args.kwargs.is_empty() {
50            return Err(vm.new_type_error("Template.__new__ only accepts *args arguments"));
51        }
52
53        let mut strings: Vec<PyObjectRef> = Vec::new();
54        let mut interpolations: Vec<PyObjectRef> = Vec::new();
55        let mut last_was_str = false;
56
57        for item in &args.args {
58            if let Ok(s) = item.clone().downcast::<PyStr>() {
59                if last_was_str {
60                    // Concatenate adjacent strings
61                    if let Some(last) = strings.last_mut() {
62                        let last_str = last.downcast_ref::<PyStr>().unwrap();
63                        let mut buf = last_str.as_wtf8().to_owned();
64                        buf.push_wtf8(s.as_wtf8());
65                        *last = vm.ctx.new_str(buf).into();
66                    }
67                } else {
68                    strings.push(s.into());
69                }
70                last_was_str = true;
71            } else if item.class().is(vm.ctx.types.interpolation_type) {
72                if !last_was_str {
73                    // Add empty string before interpolation
74                    strings.push(vm.ctx.empty_str.to_owned().into());
75                }
76                interpolations.push(item.clone());
77                last_was_str = false;
78            } else {
79                return Err(vm.new_type_error(format!(
80                    "Template.__new__ *args need to be of type 'str' or 'Interpolation', got {}",
81                    item.class().name()
82                )));
83            }
84        }
85
86        if !last_was_str {
87            // Add trailing empty string
88            strings.push(vm.ctx.empty_str.to_owned().into());
89        }
90
91        Ok(Self {
92            strings: vm.ctx.new_tuple(strings),
93            interpolations: vm.ctx.new_tuple(interpolations),
94        })
95    }
96}
97
98impl PyTemplate {
99    fn concat(&self, other: &PyObject, vm: &VirtualMachine) -> PyResult<PyRef<Self>> {
100        let other = other.downcast_ref::<Self>().ok_or_else(|| {
101            vm.new_type_error(format!(
102                r#"can only concatenate {name} (not "{}") to {name}"#,
103                other.class().slot_name(),
104                name = Self::TP_NAME,
105            ))
106        })?;
107
108        // Concatenate the two templates
109        let mut new_strings: Vec<PyObjectRef> = Vec::new();
110        let mut new_interps: Vec<PyObjectRef> = Vec::new();
111
112        // Add all strings from self except the last one
113        let self_strings_len = self.strings.as_slice().len();
114        for i in 0..self_strings_len.saturating_sub(1) {
115            new_strings.push(self.strings.as_slice().get(i).unwrap().clone());
116        }
117
118        // Add all interpolations from self
119        for interp in self.interpolations.as_slice() {
120            new_interps.push(interp.clone());
121        }
122
123        // Concatenate last string of self with first string of other
124        let mut buf = Wtf8Buf::new();
125        if let Some(s) = self
126            .strings
127            .as_slice()
128            .get(self_strings_len.saturating_sub(1))
129            .and_then(|s| s.downcast_ref::<PyStr>())
130        {
131            buf.push_wtf8(s.as_wtf8());
132        }
133        if let Some(s) = other
134            .strings
135            .as_slice()
136            .first()
137            .and_then(|s| s.downcast_ref::<PyStr>())
138        {
139            buf.push_wtf8(s.as_wtf8());
140        }
141        new_strings.push(vm.ctx.new_str(buf).into());
142
143        // Add remaining strings from other (skip first)
144        for i in 1..other.strings.as_slice().len() {
145            new_strings.push(other.strings.as_slice().get(i).unwrap().clone());
146        }
147
148        // Add all interpolations from other
149        for interp in other.interpolations.as_slice() {
150            new_interps.push(interp.clone());
151        }
152
153        let template = Self {
154            strings: vm.ctx.new_tuple(new_strings),
155            interpolations: vm.ctx.new_tuple(new_interps),
156        };
157
158        Ok(template.into_ref(&vm.ctx))
159    }
160
161    fn __add__(&self, other: &PyObject, vm: &VirtualMachine) -> PyResult<PyRef<Self>> {
162        self.concat(other, vm)
163    }
164}
165
166#[pyclass(with(Constructor, Comparable, Iterable, Representable, AsSequence))]
167impl Py<PyTemplate> {
168    #[pygetset]
169    fn values(&self, vm: &VirtualMachine) -> PyTupleRef {
170        let values: Vec<PyObjectRef> = self
171            .interpolations
172            .as_slice()
173            .iter()
174            .map(|interp| {
175                interp
176                    .downcast_ref::<PyInterpolation>()
177                    .map_or_else(|| interp.clone(), |i| i.value.clone())
178            })
179            .collect();
180        vm.ctx.new_tuple(values)
181    }
182
183    #[pyclassmethod]
184    fn __class_getitem__(
185        cls: PyTypeRef,
186        args: PyObjectRef,
187        vm: &VirtualMachine,
188    ) -> PyResult<PyGenericAlias> {
189        PyGenericAlias::from_args(cls, args, vm)
190    }
191
192    #[pymethod]
193    fn __reduce__(&self, vm: &VirtualMachine) -> PyResult<PyTupleRef> {
194        // Import string.templatelib._template_unpickle
195        // We need to import string first, then get templatelib from it,
196        // because import("string.templatelib", 0) with empty from_list returns the top-level module
197        let string_mod = vm.import("string.templatelib", 0)?;
198        let templatelib = string_mod.get_attr("templatelib", vm)?;
199        let unpickle_func = templatelib.get_attr("_template_unpickle", vm)?;
200
201        // Return (func, (strings, interpolations))
202        let args = vm.ctx.new_tuple(vec![
203            self.strings.clone().into(),
204            self.interpolations.clone().into(),
205        ]);
206        Ok(vm.ctx.new_tuple(vec![unpickle_func, args.into()]))
207    }
208}
209
210impl AsSequence for PyTemplate {
211    fn as_sequence() -> &'static PySequenceMethods {
212        static AS_SEQUENCE: LazyLock<PySequenceMethods> = LazyLock::new(|| PySequenceMethods {
213            concat: atomic_func!(|seq, other, vm| {
214                let zelf = PyTemplate::sequence_downcast(seq);
215                zelf.concat(other, vm).map(|t| t.into())
216            }),
217            ..PySequenceMethods::NOT_IMPLEMENTED
218        });
219        &AS_SEQUENCE
220    }
221}
222
223impl Comparable for PyTemplate {
224    fn cmp(
225        zelf: &Py<Self>,
226        other: &PyObject,
227        op: PyComparisonOp,
228        vm: &VirtualMachine,
229    ) -> PyResult<PyComparisonValue> {
230        op.eq_only(|| {
231            let other = class_or_notimplemented!(Self, other);
232
233            let eq = vm.bool_eq(zelf.strings.as_object(), other.strings.as_object())?
234                && vm.bool_eq(
235                    zelf.interpolations.as_object(),
236                    other.interpolations.as_object(),
237                )?;
238
239            Ok(eq.into())
240        })
241    }
242}
243
244impl Iterable for PyTemplate {
245    fn iter(zelf: PyRef<Self>, vm: &VirtualMachine) -> PyResult {
246        Ok(PyTemplateIter::new(zelf).into_pyobject(vm))
247    }
248}
249
250impl Representable for PyTemplate {
251    #[inline]
252    fn repr_wtf8(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<Wtf8Buf> {
253        let strings_repr = zelf.strings.as_object().repr(vm)?;
254        let interp_repr = zelf.interpolations.as_object().repr(vm)?;
255        Ok(wtf8_concat!(
256            "Template(strings=",
257            strings_repr.as_wtf8(),
258            ", interpolations=",
259            interp_repr.as_wtf8(),
260            ')',
261        ))
262    }
263}
264
265#[pyclass(module = "string.templatelib", name = "TemplateIter")]
266#[derive(Debug)]
267pub struct PyTemplateIter {
268    template: PyRef<PyTemplate>,
269    index: core::sync::atomic::AtomicUsize,
270    from_strings: core::sync::atomic::AtomicBool,
271}
272
273impl PyPayload for PyTemplateIter {
274    #[inline]
275    fn class(ctx: &Context) -> &'static Py<PyType> {
276        ctx.types.template_iter_type
277    }
278}
279
280impl PyTemplateIter {
281    fn new(template: PyRef<PyTemplate>) -> Self {
282        Self {
283            template,
284            index: core::sync::atomic::AtomicUsize::new(0),
285            from_strings: core::sync::atomic::AtomicBool::new(true),
286        }
287    }
288}
289
290#[pyclass(with(IterNext, Iterable))]
291impl PyTemplateIter {}
292
293impl SelfIter for PyTemplateIter {}
294
295impl IterNext for PyTemplateIter {
296    fn next(zelf: &Py<Self>, _vm: &VirtualMachine) -> PyResult<PyIterReturn> {
297        use core::sync::atomic::Ordering;
298
299        loop {
300            let from_strings = zelf.from_strings.load(Ordering::SeqCst);
301            let index = zelf.index.load(Ordering::SeqCst);
302
303            if from_strings {
304                if index < zelf.template.strings.as_slice().len() {
305                    let item = zelf.template.strings.as_slice().get(index).unwrap();
306                    zelf.from_strings.store(false, Ordering::SeqCst);
307
308                    // Skip empty strings
309                    if let Some(s) = item.downcast_ref::<PyStr>()
310                        && s.as_wtf8().is_empty()
311                    {
312                        continue;
313                    }
314                    return Ok(PyIterReturn::Return(item.clone()));
315                }
316                return Ok(PyIterReturn::StopIteration(None));
317            } else if index < zelf.template.interpolations.as_slice().len() {
318                let item = zelf.template.interpolations.as_slice().get(index).unwrap();
319                zelf.index.fetch_add(1, Ordering::SeqCst);
320                zelf.from_strings.store(true, Ordering::SeqCst);
321                return Ok(PyIterReturn::Return(item.clone()));
322            }
323            return Ok(PyIterReturn::StopIteration(None));
324        }
325    }
326}
327
328pub(crate) fn init(context: &'static Context) {
329    PyTemplate::extend_class(context, context.types.template_type);
330    PyTemplateIter::extend_class(context, context.types.template_iter_type);
331}