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 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 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 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 let mut new_strings: Vec<PyObjectRef> = Vec::new();
110 let mut new_interps: Vec<PyObjectRef> = Vec::new();
111
112 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 for interp in self.interpolations.as_slice() {
120 new_interps.push(interp.clone());
121 }
122
123 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 for i in 1..other.strings.as_slice().len() {
145 new_strings.push(other.strings.as_slice().get(i).unwrap().clone());
146 }
147
148 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 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 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 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}