Skip to main content

rustpython_common/encodings/
utf7.rs

1//! UTF-7 encode / incremental decode.
2
3use super::*;
4use crate::wtf8::{CodePoint, Wtf8, Wtf8Buf};
5
6pub const ENCODING_NAME: &str = "utf-7";
7
8fn is_base64(b: u8) -> bool {
9    b.is_ascii_alphanumeric() || b == b'+' || b == b'/'
10}
11
12fn to_base64(n: u32) -> u8 {
13    b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"[(n & 0x3f) as usize]
14}
15
16fn from_base64(b: u8) -> u32 {
17    match b {
18        b'a'..=b'z' => (b - 71) as u32,
19        b'A'..=b'Z' => (b - 65) as u32,
20        b'0'..=b'9' => (b + 4) as u32,
21        b'+' => 62,
22        _ => 63,
23    }
24}
25
26fn decode_direct(b: u8) -> bool {
27    b <= 127 && b != b'+'
28}
29
30fn category(oc: u32) -> u8 {
31    if oc > 127 {
32        return 3;
33    }
34    let b = oc as u8;
35    if matches!(b, b'\t' | b'\n' | b'\r' | b' ') {
36        2
37    } else if b.is_ascii_alphanumeric() || b"'(),-./:?".contains(&b) {
38        0
39    } else if b"!\"#$%&*;<=>@[]^_`{|}".contains(&b) {
40        1
41    } else {
42        3
43    }
44}
45
46fn encode_direct(oc: u32) -> bool {
47    oc < 128 && oc > 0 && category(oc) != 3
48}
49
50fn encode_unit(out: &mut Vec<u8>, unit: u32, bits: &mut u32, buffer: &mut u32) {
51    *bits += 16;
52    *buffer = (*buffer << 16) | unit;
53    while *bits >= 6 {
54        out.push(to_base64(*buffer >> (*bits - 6)));
55        *bits -= 6;
56    }
57    *buffer &= (1 << *bits) - 1;
58}
59
60/// Encode `s` as UTF-7. Every code point is representable, so this cannot fail.
61#[must_use]
62pub fn encode_bytes(s: &Wtf8) -> Vec<u8> {
63    let mut out = Vec::new();
64    let mut in_shift = false;
65    let mut bits = 0;
66    let mut buffer = 0;
67    for cp in s.code_points() {
68        let oc = cp.to_u32();
69        if !in_shift {
70            if oc == b'+' as u32 {
71                out.extend_from_slice(b"+-");
72            } else if encode_direct(oc) {
73                out.push(oc as u8);
74            } else {
75                out.push(b'+');
76                in_shift = true;
77                emit_code(&mut out, oc, &mut bits, &mut buffer);
78            }
79        } else if encode_direct(oc) {
80            if bits != 0 {
81                out.push(to_base64(buffer << (6 - bits)));
82                buffer = 0;
83                bits = 0;
84            }
85            in_shift = false;
86            if is_base64(oc as u8) || oc == b'-' as u32 {
87                out.push(b'-');
88            }
89            out.push(oc as u8);
90        } else {
91            emit_code(&mut out, oc, &mut bits, &mut buffer);
92        }
93    }
94    if bits != 0 {
95        out.push(to_base64(buffer << (6 - bits)));
96    }
97    if in_shift {
98        out.push(b'-');
99    }
100    out
101}
102
103fn emit_code(out: &mut Vec<u8>, oc: u32, bits: &mut u32, buffer: &mut u32) {
104    if oc >= 0x10000 {
105        encode_unit(out, 0xd800 | ((oc - 0x10000) >> 10), bits, buffer);
106        encode_unit(out, 0xdc00 | ((oc - 0x10000) & 0x3ff), bits, buffer);
107    } else {
108        encode_unit(out, oc, bits, buffer);
109    }
110}
111
112pub fn encode<Ctx, E>(ctx: Ctx, _errors: &E) -> Result<Vec<u8>, Ctx::Error>
113where
114    Ctx: EncodeContext,
115    E: EncodeErrorHandler<Ctx>,
116{
117    Ok(encode_bytes(ctx.remaining_data()))
118}
119
120/// Incremental UTF-7 decode.
121pub fn decode<Ctx, E>(
122    mut ctx: Ctx,
123    errors: &E,
124    final_decode: bool,
125) -> Result<(Wtf8Buf, usize), Ctx::Error>
126where
127    Ctx: DecodeContext,
128    E: DecodeErrorHandler<Ctx>,
129{
130    let mut out = Wtf8Buf::new();
131    let mut in_shift = false;
132    let mut bits = 0u32;
133    let mut buffer = 0u32;
134    let mut surrogate = 0u32;
135    let mut shift_out_start = 0usize;
136    let mut startinpos = 0usize;
137    loop {
138        let rest = ctx.remaining_data();
139        if rest.is_empty() {
140            break;
141        }
142        let ch = rest[0];
143        if in_shift {
144            if is_base64(ch) {
145                buffer = (buffer << 6) | from_base64(ch);
146                bits += 6;
147                ctx.advance(1);
148                if bits >= 16 {
149                    let out_ch = buffer >> (bits - 16);
150                    bits -= 16;
151                    buffer &= (1 << bits) - 1;
152                    if surrogate != 0 {
153                        if (0xdc00..=0xdfff).contains(&out_ch) {
154                            let code = (((surrogate & 0x3ff) << 10) | (out_ch & 0x3ff)) + 0x10000;
155                            if let Some(cp) = CodePoint::from_u32(code) {
156                                out.push(cp);
157                            }
158                            surrogate = 0;
159                            continue;
160                        }
161                        if let Some(cp) = CodePoint::from_u32(surrogate) {
162                            out.push(cp);
163                        }
164                        surrogate = 0;
165                    }
166                    if (0xd800..=0xdbff).contains(&out_ch) {
167                        surrogate = out_ch;
168                    } else if let Some(cp) = CodePoint::from_u32(out_ch) {
169                        out.push(cp);
170                    }
171                }
172            } else {
173                in_shift = false;
174                if bits >= 6 {
175                    ctx.advance(1);
176                    let start = startinpos;
177                    let end = ctx.position();
178                    let replace = ctx.handle_error(
179                        errors,
180                        start..end,
181                        Some("partial character in shift sequence"),
182                    )?;
183                    out.push_wtf8(replace.as_ref());
184                    continue;
185                } else if bits > 0 && buffer != 0 {
186                    ctx.advance(1);
187                    let start = startinpos;
188                    let end = ctx.position();
189                    let replace = ctx.handle_error(
190                        errors,
191                        start..end,
192                        Some("non-zero padding bits in shift sequence"),
193                    )?;
194                    out.push_wtf8(replace.as_ref());
195                    continue;
196                }
197                if surrogate != 0
198                    && decode_direct(ch)
199                    && let Some(cp) = CodePoint::from_u32(surrogate)
200                {
201                    out.push(cp);
202                }
203                surrogate = 0;
204                if ch == b'-' {
205                    ctx.advance(1);
206                }
207            }
208        } else if ch == b'+' {
209            startinpos = ctx.position();
210            if rest.len() >= 2 && rest[1] == b'-' {
211                ctx.advance(2);
212                out.push_char('+');
213            } else if rest.len() >= 2 && !is_base64(rest[1]) {
214                let start = ctx.position();
215                let replace =
216                    ctx.handle_error(errors, start..start + 2, Some("ill-formed sequence"))?;
217                out.push_wtf8(replace.as_ref());
218            } else {
219                ctx.advance(1);
220                in_shift = true;
221                surrogate = 0;
222                shift_out_start = out.len();
223                bits = 0;
224                buffer = 0;
225            }
226        } else if decode_direct(ch) {
227            out.push_char(ch as char);
228            ctx.advance(1);
229        } else {
230            let start = ctx.position();
231            ctx.advance(1);
232            let replace = ctx.handle_error(
233                errors,
234                start..ctx.position(),
235                Some("unexpected special character"),
236            )?;
237            out.push_wtf8(replace.as_ref());
238        }
239    }
240    let mut consumed = ctx.position();
241    if in_shift && final_decode {
242        if surrogate != 0 || bits >= 6 || (bits > 0 && buffer != 0) {
243            let start = startinpos;
244            let end = ctx.position();
245            let replace =
246                ctx.handle_error(errors, start..end, Some("unterminated shift sequence"))?;
247            out.push_wtf8(replace.as_ref());
248        }
249    } else if in_shift {
250        consumed = startinpos;
251        out.truncate(shift_out_start);
252    }
253    Ok((out, consumed))
254}