Skip to main content

rustpython_vm/builtins/
interpolation.rs

1use super::{
2    PyStr, PyStrRef, PyTupleRef, PyType, PyTypeRef, genericalias::PyGenericAlias,
3    tuple::IntoPyTuple,
4};
5use crate::{
6    AsObject, Context, Py, PyObject, PyObjectRef, PyPayload, PyRef, PyResult, VirtualMachine,
7    class::PyClassImpl,
8    common::hash::PyHash,
9    convert::ToPyObject,
10    function::PyComparisonValue,
11    types::{Comparable, Constructor, Hashable, PyComparisonOp, Representable},
12};
13use itertools::Itertools;
14use rustpython_common::wtf8::Wtf8Buf;
15
16#[pyclass(module = "string.templatelib", name = "Interpolation")]
17#[derive(Debug, Clone)]
18pub struct PyInterpolation {
19    #[pymember(type = "object_ex")]
20    pub value: PyObjectRef,
21    #[pymember(type = "object_ex")]
22    pub expression: PyStrRef,
23    #[pymember(type = "object_ex")]
24    pub conversion: PyObjectRef, // None or 's', 'r', 'a'
25    #[pymember(type = "object_ex")]
26    pub format_spec: PyStrRef,
27}
28
29impl PyPayload for PyInterpolation {
30    #[inline]
31    fn class(ctx: &Context) -> &'static Py<PyType> {
32        ctx.types.interpolation_type
33    }
34}
35
36impl PyInterpolation {
37    pub fn new(
38        value: PyObjectRef,
39        expression: PyStrRef,
40        conversion: PyObjectRef,
41        format_spec: PyStrRef,
42        vm: &VirtualMachine,
43    ) -> PyResult<Self> {
44        // Validate conversion like _PyInterpolation_Build does
45        let is_valid = vm.is_none(&conversion)
46            || conversion
47                .downcast_ref::<PyStr>()
48                .is_some_and(|s| matches!(s.to_str(), Some("s" | "r" | "a")));
49        if !is_valid {
50            return Err(vm.new_system_error(
51                "Interpolation() argument 'conversion' must be one of 's', 'a' or 'r'",
52            ));
53        }
54        Ok(Self {
55            value,
56            expression,
57            conversion,
58            format_spec,
59        })
60    }
61}
62
63impl Constructor for PyInterpolation {
64    type Args = InterpolationArgs;
65
66    fn py_new(_cls: &Py<PyType>, args: Self::Args, vm: &VirtualMachine) -> PyResult<Self> {
67        let conversion: PyObjectRef = if let Some(s) = args.conversion {
68            let has_flag = s
69                .as_bytes()
70                .iter()
71                .exactly_one()
72                .is_ok_and(|s| matches!(*s, b's' | b'r' | b'a'));
73            if !has_flag {
74                return Err(vm.new_value_error(
75                    "Interpolation() argument 'conversion' must be one of 's', 'a' or 'r'",
76                ));
77            }
78            s.into()
79        } else {
80            vm.ctx.none()
81        };
82
83        let expression = args.expression;
84        let format_spec = args.format_spec;
85
86        Ok(Self {
87            value: args.value,
88            expression,
89            conversion,
90            format_spec,
91        })
92    }
93}
94
95#[derive(FromArgs)]
96pub struct InterpolationArgs {
97    #[pyarg(positional)]
98    value: PyObjectRef,
99    #[pyarg(any, default = "")]
100    expression: PyStrRef,
101    #[pyarg(
102        any,
103        optional,
104        error_msg = "Interpolation() argument 'conversion' must be str or None"
105    )]
106    conversion: Option<PyStrRef>,
107    #[pyarg(any, default = "")]
108    format_spec: PyStrRef,
109}
110
111#[pyclass(with(Constructor, Comparable, Hashable, Representable))]
112impl PyInterpolation {
113    #[pyattr]
114    fn __match_args__(ctx: &Context) -> PyTupleRef {
115        ctx.new_tuple(vec![
116            ctx.intern_str("value").to_owned().into(),
117            ctx.intern_str("expression").to_owned().into(),
118            ctx.intern_str("conversion").to_owned().into(),
119            ctx.intern_str("format_spec").to_owned().into(),
120        ])
121    }
122
123    #[pyclassmethod]
124    fn __class_getitem__(
125        cls: PyTypeRef,
126        args: PyObjectRef,
127        vm: &VirtualMachine,
128    ) -> PyResult<PyGenericAlias> {
129        PyGenericAlias::from_args(cls, args, vm)
130    }
131
132    #[pymethod]
133    fn __reduce__(zelf: PyRef<Self>, vm: &VirtualMachine) -> PyTupleRef {
134        let cls = zelf.class().to_owned();
135        let args = (
136            zelf.value.clone(),
137            zelf.expression.clone(),
138            zelf.conversion.clone(),
139            zelf.format_spec.clone(),
140        );
141        (cls, args.to_pyobject(vm)).into_pytuple(vm)
142    }
143}
144
145impl Comparable for PyInterpolation {
146    fn cmp(
147        zelf: &Py<Self>,
148        other: &PyObject,
149        op: PyComparisonOp,
150        vm: &VirtualMachine,
151    ) -> PyResult<PyComparisonValue> {
152        op.eq_only(|| {
153            let other = class_or_notimplemented!(Self, other);
154
155            let eq = vm.bool_eq(&zelf.value, &other.value)?
156                && vm.bool_eq(zelf.expression.as_object(), other.expression.as_object())?
157                && vm.bool_eq(&zelf.conversion, &other.conversion)?
158                && vm.bool_eq(zelf.format_spec.as_object(), other.format_spec.as_object())?;
159
160            Ok(eq.into())
161        })
162    }
163}
164
165impl Hashable for PyInterpolation {
166    fn hash(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<PyHash> {
167        // Hash based on (value, expression, conversion, format_spec)
168        let value_hash = zelf.value.hash(vm)?;
169        let expr_hash = zelf.expression.as_object().hash(vm)?;
170        let conv_hash = zelf.conversion.hash(vm)?;
171        let spec_hash = zelf.format_spec.as_object().hash(vm)?;
172
173        // Combine hashes
174        Ok(value_hash
175            .wrapping_add(expr_hash.wrapping_mul(3))
176            .wrapping_add(conv_hash.wrapping_mul(5))
177            .wrapping_add(spec_hash.wrapping_mul(7)))
178    }
179}
180
181impl Representable for PyInterpolation {
182    #[inline]
183    fn repr_wtf8(zelf: &Py<Self>, vm: &VirtualMachine) -> PyResult<Wtf8Buf> {
184        let value_repr = zelf.value.repr(vm)?;
185        let expr_repr = zelf.expression.repr(vm)?;
186        let spec_repr = zelf.format_spec.repr(vm)?;
187
188        let mut result = Wtf8Buf::from("Interpolation(");
189        result.push_wtf8(value_repr.as_wtf8());
190        result.push_str(", ");
191        result.push_str(&expr_repr);
192        result.push_str(", ");
193        if vm.is_none(&zelf.conversion) {
194            result.push_str("None");
195        } else {
196            result.push_wtf8(zelf.conversion.repr(vm)?.as_wtf8());
197        }
198        result.push_str(", ");
199        result.push_str(&spec_repr);
200        result.push_char(')');
201
202        Ok(result)
203    }
204}
205
206pub(crate) fn init(context: &'static Context) {
207    PyInterpolation::extend_class(context, context.types.interpolation_type);
208}