1use core::ops::{self, Range};
2
3use num_traits::ToPrimitive;
4
5use crate::str::StrKind;
6use crate::wtf8::{CodePoint, Wtf8, Wtf8Buf};
7
8#[cfg(feature = "cjk-codecs")]
9pub mod cjk;
10mod wide;
11pub use wide::ByteOrder;
12pub mod escape;
13pub mod raw_unicode_escape;
14pub mod unicode_escape;
15pub mod utf16;
16pub mod utf32;
17pub mod utf7;
18
19pub trait StrBuffer: AsRef<Wtf8> {
20 fn is_compatible_with(&self, kind: StrKind) -> bool {
21 let s = self.as_ref();
22 match kind {
23 StrKind::Ascii => s.is_ascii(),
24 StrKind::Utf8 => s.is_utf8(),
25 StrKind::Wtf8 => true,
26 }
27 }
28}
29
30pub trait CodecContext: Sized {
31 type Error;
32 type StrBuf: StrBuffer;
33 type BytesBuf: AsRef<[u8]>;
34
35 fn string(&self, s: Wtf8Buf) -> Self::StrBuf;
36 fn bytes(&self, b: Vec<u8>) -> Self::BytesBuf;
37}
38
39pub trait EncodeContext: CodecContext {
40 fn full_data(&self) -> &Wtf8;
41 fn data_len(&self) -> StrSize;
42
43 fn remaining_data(&self) -> &Wtf8;
44 fn position(&self) -> StrSize;
45
46 fn restart_from(&mut self, pos: StrSize) -> Result<(), Self::Error>;
47
48 fn error_encoding(&self, range: Range<StrSize>, reason: Option<&str>) -> Self::Error;
49
50 fn handle_error<E>(
51 &mut self,
52 errors: &E,
53 range: Range<StrSize>,
54 reason: Option<&str>,
55 ) -> Result<EncodeReplace<Self>, Self::Error>
56 where
57 E: EncodeErrorHandler<Self>,
58 {
59 let (replace, restart) = errors.handle_encode_error(self, range, reason)?;
60 self.restart_from(restart)?;
61 Ok(replace)
62 }
63}
64
65pub trait DecodeContext: CodecContext {
66 fn full_data(&self) -> &[u8];
67
68 fn remaining_data(&self) -> &[u8];
69 fn position(&self) -> usize;
70
71 fn advance(&mut self, by: usize);
72
73 fn restart_from(&mut self, pos: usize) -> Result<(), Self::Error>;
74
75 fn error_decoding(&self, byte_range: Range<usize>, reason: Option<&str>) -> Self::Error;
76
77 fn handle_error<E>(
78 &mut self,
79 errors: &E,
80 byte_range: Range<usize>,
81 reason: Option<&str>,
82 ) -> Result<Self::StrBuf, Self::Error>
83 where
84 E: DecodeErrorHandler<Self>,
85 {
86 let (replace, restart) = errors.handle_decode_error(self, byte_range, reason)?;
87 self.restart_from(restart)?;
88 Ok(replace)
89 }
90}
91
92pub trait EncodeErrorHandler<Ctx: EncodeContext> {
93 fn handle_encode_error(
94 &self,
95 ctx: &mut Ctx,
96 range: Range<StrSize>,
97 reason: Option<&str>,
98 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error>;
99}
100pub trait DecodeErrorHandler<Ctx: DecodeContext> {
101 fn handle_decode_error(
102 &self,
103 ctx: &mut Ctx,
104 byte_range: Range<usize>,
105 reason: Option<&str>,
106 ) -> Result<(Ctx::StrBuf, usize), Ctx::Error>;
107}
108
109pub enum EncodeReplace<Ctx: CodecContext> {
110 Str(Ctx::StrBuf),
111 Bytes(Ctx::BytesBuf),
112}
113
114#[derive(Copy, Clone, Default, Debug)]
115pub struct StrSize {
116 pub bytes: usize,
117 pub chars: usize,
118}
119
120fn iter_code_points(w: &Wtf8) -> impl Iterator<Item = (StrSize, CodePoint)> {
121 w.code_point_indices()
122 .enumerate()
123 .map(|(chars, (bytes, c))| (StrSize { bytes, chars }, c))
124}
125
126impl ops::Add for StrSize {
127 type Output = Self;
128 fn add(self, rhs: Self) -> Self::Output {
129 Self {
130 bytes: self.bytes + rhs.bytes,
131 chars: self.chars + rhs.chars,
132 }
133 }
134}
135
136impl ops::AddAssign for StrSize {
137 fn add_assign(&mut self, rhs: Self) {
138 self.bytes += rhs.bytes;
139 self.chars += rhs.chars;
140 }
141}
142
143struct DecodeError<'a> {
144 valid_prefix: &'a str,
145 rest: &'a [u8],
146 err_len: Option<usize>,
147}
148
149const unsafe fn make_decode_err(
152 v: &[u8],
153 valid_up_to: usize,
154 err_len: Option<usize>,
155) -> DecodeError<'_> {
156 let (valid_prefix, rest) = unsafe { v.split_at_unchecked(valid_up_to) };
157 let valid_prefix = unsafe { core::str::from_utf8_unchecked(valid_prefix) };
158 DecodeError {
159 valid_prefix,
160 rest,
161 err_len,
162 }
163}
164
165enum HandleResult<'a> {
166 Done,
167 Error {
168 err_len: Option<usize>,
169 reason: &'a str,
170 },
171}
172
173fn decode_utf8_compatible<Ctx, E, DecodeF, ErrF>(
174 mut ctx: Ctx,
175 errors: &E,
176 decode: DecodeF,
177 handle_error: ErrF,
178) -> Result<(Wtf8Buf, usize), Ctx::Error>
179where
180 Ctx: DecodeContext,
181 E: DecodeErrorHandler<Ctx>,
182 DecodeF: Fn(&[u8]) -> Result<&str, DecodeError<'_>>,
183 ErrF: Fn(&[u8], Option<usize>) -> HandleResult<'static>,
184{
185 if ctx.remaining_data().is_empty() {
186 return Ok((Wtf8Buf::new(), 0));
187 }
188 let mut out = Wtf8Buf::with_capacity(ctx.remaining_data().len());
189 loop {
190 match decode(ctx.remaining_data()) {
191 Ok(decoded) => {
192 out.push_str(decoded);
193 ctx.advance(decoded.len());
194 break;
195 }
196 Err(e) => {
197 out.push_str(e.valid_prefix);
198 match handle_error(e.rest, e.err_len) {
199 HandleResult::Done => {
200 ctx.advance(e.valid_prefix.len());
201 break;
202 }
203 HandleResult::Error { err_len, reason } => {
204 let err_start = ctx.position() + e.valid_prefix.len();
205 let err_end = match err_len {
206 Some(len) => err_start + len,
207 None => ctx.full_data().len(),
208 };
209 let err_range = err_start..err_end;
210 let replace = ctx.handle_error(errors, err_range, Some(reason))?;
211 out.push_wtf8(replace.as_ref());
212 continue;
213 }
214 }
215 }
216 }
217 }
218 Ok((out, ctx.position()))
219}
220
221#[inline]
222fn encode_utf8_compatible<Ctx, E>(
223 mut ctx: Ctx,
224 errors: &E,
225 err_reason: &str,
226 target_kind: StrKind,
227) -> Result<Vec<u8>, Ctx::Error>
228where
229 Ctx: EncodeContext,
230 E: EncodeErrorHandler<Ctx>,
231{
232 let mut out = Vec::<u8>::with_capacity(ctx.remaining_data().len());
235 loop {
236 let data = ctx.remaining_data();
237 let mut iter = iter_code_points(data);
238 let Some((i, _)) = iter.find(|(_, c)| !target_kind.can_encode(*c)) else {
239 break;
240 };
241
242 out.extend_from_slice(&ctx.remaining_data().as_bytes()[..i.bytes]);
243
244 let err_start = ctx.position() + i;
245 let err_end = match { iter }.find(|(_, c)| target_kind.can_encode(*c)) {
247 Some((i, _)) => ctx.position() + i,
248 None => ctx.data_len(),
249 };
250
251 let range = err_start..err_end;
252 let replace = ctx.handle_error(errors, range.clone(), Some(err_reason))?;
253 match replace {
254 EncodeReplace::Str(s) => {
255 if s.is_compatible_with(target_kind) {
256 out.extend_from_slice(s.as_ref().as_bytes());
257 } else {
258 return Err(ctx.error_encoding(range, Some(err_reason)));
259 }
260 }
261 EncodeReplace::Bytes(b) => {
262 out.extend_from_slice(b.as_ref());
263 }
264 }
265 }
266 out.extend_from_slice(ctx.remaining_data().as_bytes());
267 Ok(out)
268}
269
270pub mod errors {
271 use crate::str::UnicodeEscapeCodepoint;
272
273 use super::*;
274 use core::fmt::Write;
275
276 #[derive(Clone, Copy)]
277 pub struct Strict;
278
279 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for Strict {
280 fn handle_encode_error(
281 &self,
282 ctx: &mut Ctx,
283 range: Range<StrSize>,
284 reason: Option<&str>,
285 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
286 Err(ctx.error_encoding(range, reason))
287 }
288 }
289
290 impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for Strict {
291 fn handle_decode_error(
292 &self,
293 ctx: &mut Ctx,
294 byte_range: Range<usize>,
295 reason: Option<&str>,
296 ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
297 Err(ctx.error_decoding(byte_range, reason))
298 }
299 }
300
301 #[derive(Clone, Copy)]
302 pub struct Ignore;
303
304 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for Ignore {
305 fn handle_encode_error(
306 &self,
307 ctx: &mut Ctx,
308 range: Range<StrSize>,
309 _reason: Option<&str>,
310 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
311 Ok((EncodeReplace::Bytes(ctx.bytes(b"".into())), range.end))
312 }
313 }
314
315 impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for Ignore {
316 fn handle_decode_error(
317 &self,
318 ctx: &mut Ctx,
319 byte_range: Range<usize>,
320 _reason: Option<&str>,
321 ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
322 Ok((ctx.string("".into()), byte_range.end))
323 }
324 }
325
326 #[derive(Clone, Copy)]
327 pub struct Replace;
328
329 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for Replace {
330 fn handle_encode_error(
331 &self,
332 ctx: &mut Ctx,
333 range: Range<StrSize>,
334 _reason: Option<&str>,
335 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
336 let replace = "?".repeat(range.end.chars - range.start.chars);
337 Ok((EncodeReplace::Str(ctx.string(replace.into())), range.end))
338 }
339 }
340
341 impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for Replace {
342 fn handle_decode_error(
343 &self,
344 ctx: &mut Ctx,
345 byte_range: Range<usize>,
346 _reason: Option<&str>,
347 ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
348 Ok((
349 ctx.string(char::REPLACEMENT_CHARACTER.to_string().into()),
350 byte_range.end,
351 ))
352 }
353 }
354
355 #[derive(Clone, Copy)]
356 pub struct XmlCharRefReplace;
357
358 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for XmlCharRefReplace {
359 fn handle_encode_error(
360 &self,
361 ctx: &mut Ctx,
362 range: Range<StrSize>,
363 _reason: Option<&str>,
364 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
365 let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
366 let num_chars = range.end.chars - range.start.chars;
367 let mut out = String::with_capacity(num_chars * 6);
369 for c in err_str.code_points() {
370 write!(out, "&#{};", c.to_u32()).unwrap()
371 }
372 Ok((EncodeReplace::Str(ctx.string(out.into())), range.end))
373 }
374 }
375
376 #[derive(Clone, Copy)]
377 pub struct BackslashReplace;
378
379 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for BackslashReplace {
380 fn handle_encode_error(
381 &self,
382 ctx: &mut Ctx,
383 range: Range<StrSize>,
384 _reason: Option<&str>,
385 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
386 let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
387 let num_chars = range.end.chars - range.start.chars;
388 let mut out = String::with_capacity(num_chars * 4);
390 for c in err_str.code_points() {
391 write!(out, "{}", UnicodeEscapeCodepoint(c)).unwrap();
392 }
393 Ok((EncodeReplace::Str(ctx.string(out.into())), range.end))
394 }
395 }
396
397 impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for BackslashReplace {
398 fn handle_decode_error(
399 &self,
400 ctx: &mut Ctx,
401 byte_range: Range<usize>,
402 _reason: Option<&str>,
403 ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
404 let err_bytes = &ctx.full_data()[byte_range.clone()];
405 let mut replace = String::with_capacity(4 * err_bytes.len());
406 for &c in err_bytes {
407 write!(replace, "\\x{c:02x}").unwrap();
408 }
409 Ok((ctx.string(replace.into()), byte_range.end))
410 }
411 }
412
413 #[derive(Clone, Copy)]
414 pub struct NameReplace;
415
416 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for NameReplace {
417 fn handle_encode_error(
418 &self,
419 ctx: &mut Ctx,
420 range: Range<StrSize>,
421 _reason: Option<&str>,
422 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
423 let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
424 let num_chars = range.end.chars - range.start.chars;
425 let mut out = String::with_capacity(num_chars * 4);
426 for c in err_str.code_points() {
427 let c_u32 = c.to_u32();
428 if let Some(c_name) = c.to_char().and_then(rustpython_unicode::character_name) {
429 write!(out, "\\N{{{c_name}}}").unwrap();
430 } else if c_u32 >= 0x10000 {
431 write!(out, "\\U{c_u32:08x}").unwrap();
432 } else if c_u32 >= 0x100 {
433 write!(out, "\\u{c_u32:04x}").unwrap();
434 } else {
435 write!(out, "\\x{c_u32:02x}").unwrap();
436 }
437 }
438 Ok((EncodeReplace::Str(ctx.string(out.into())), range.end))
439 }
440 }
441
442 #[derive(Clone, Copy)]
443 pub struct SurrogateEscape;
444
445 impl<Ctx: EncodeContext> EncodeErrorHandler<Ctx> for SurrogateEscape {
446 fn handle_encode_error(
447 &self,
448 ctx: &mut Ctx,
449 range: Range<StrSize>,
450 reason: Option<&str>,
451 ) -> Result<(EncodeReplace<Ctx>, StrSize), Ctx::Error> {
452 let err_str = &ctx.full_data()[range.start.bytes..range.end.bytes];
453 let num_chars = range.end.chars - range.start.chars;
454 let mut out = Vec::with_capacity(num_chars);
455 let mut pos = range.start;
456 for ch in err_str.code_points() {
457 let ch_u32 = ch.to_u32();
458 if !(0xdc80..=0xdcff).contains(&ch_u32) {
459 if out.is_empty() {
460 return Err(ctx.error_encoding(range, reason));
462 }
463 return Ok((EncodeReplace::Bytes(ctx.bytes(out)), pos));
465 }
466 out.push((ch_u32 - 0xdc00) as u8);
467 pos += StrSize {
468 bytes: ch.len_wtf8(),
469 chars: 1,
470 };
471 }
472 Ok((EncodeReplace::Bytes(ctx.bytes(out)), range.end))
473 }
474 }
475
476 impl<Ctx: DecodeContext> DecodeErrorHandler<Ctx> for SurrogateEscape {
477 fn handle_decode_error(
478 &self,
479 ctx: &mut Ctx,
480 byte_range: Range<usize>,
481 reason: Option<&str>,
482 ) -> Result<(Ctx::StrBuf, usize), Ctx::Error> {
483 let err_bytes = &ctx.full_data()[byte_range.clone()];
484 let mut consumed = 0;
485 let mut replace = Wtf8Buf::with_capacity(4 * byte_range.len());
486 while consumed < 4 && consumed < byte_range.len() {
487 let c = err_bytes[consumed] as u16;
488 if c < 128 {
490 break;
491 }
492 replace.push(CodePoint::from(0xdc00 + c));
493 consumed += 1;
494 }
495 if consumed == 0 {
496 return Err(ctx.error_decoding(byte_range, reason));
497 }
498 Ok((ctx.string(replace), byte_range.start + consumed))
499 }
500 }
501}
502
503pub mod utf8 {
504 use super::*;
505
506 pub const ENCODING_NAME: &str = "utf-8";
507
508 #[inline]
509 pub fn encode<Ctx, E>(ctx: Ctx, errors: &E) -> Result<Vec<u8>, Ctx::Error>
510 where
511 Ctx: EncodeContext,
512 E: EncodeErrorHandler<Ctx>,
513 {
514 encode_utf8_compatible(ctx, errors, "surrogates not allowed", StrKind::Utf8)
515 }
516
517 pub fn decode<Ctx: DecodeContext, E: DecodeErrorHandler<Ctx>>(
518 ctx: Ctx,
519 errors: &E,
520 final_decode: bool,
521 ) -> Result<(Wtf8Buf, usize), Ctx::Error> {
522 decode_utf8_compatible(
523 ctx,
524 errors,
525 |v| {
526 core::str::from_utf8(v).map_err(|e| {
527 unsafe { make_decode_err(v, e.valid_up_to(), e.error_len()) }
530 })
531 },
532 |rest, err_len| {
533 let first_err = rest[0];
534 if matches!(first_err, 0x80..=0xc1 | 0xf5..=0xff) {
535 HandleResult::Error {
536 err_len: Some(1),
537 reason: "invalid start byte",
538 }
539 } else if err_len.is_none() {
540 if final_decode {
542 HandleResult::Error {
543 err_len,
544 reason: "unexpected end of data",
545 }
546 } else {
547 HandleResult::Done
548 }
549 } else if !final_decode && matches!(rest, [0xed, 0xa0..=0xbf]) {
550 HandleResult::Done
552 } else {
553 HandleResult::Error {
554 err_len,
555 reason: "invalid continuation byte",
556 }
557 }
558 },
559 )
560 }
561}
562
563pub mod latin_1 {
564 use super::*;
565
566 pub const ENCODING_NAME: &str = "latin-1";
567
568 const ERR_REASON: &str = "ordinal not in range(256)";
569
570 #[inline]
571 pub fn encode<Ctx, E>(mut ctx: Ctx, errors: &E) -> Result<Vec<u8>, Ctx::Error>
572 where
573 Ctx: EncodeContext,
574 E: EncodeErrorHandler<Ctx>,
575 {
576 let mut out = Vec::<u8>::new();
577 loop {
578 let data = ctx.remaining_data();
579 let mut iter = iter_code_points(ctx.remaining_data());
580 let Some((i, ch)) = iter.find(|(_, c)| !c.is_ascii()) else {
581 break;
582 };
583 out.extend_from_slice(&data.as_bytes()[..i.bytes]);
584 let err_start = ctx.position() + i;
585 if let Some(byte) = ch.to_u32().to_u8() {
586 drop(iter);
587 out.push(byte);
588 ctx.restart_from(err_start + StrSize { bytes: 2, chars: 1 })?;
590 } else {
591 let err_end = match { iter }.find(|(_, c)| c.to_u32() <= 255) {
593 Some((i, _)) => ctx.position() + i,
594 None => ctx.data_len(),
595 };
596 let err_range = err_start..err_end;
597 let replace = ctx.handle_error(errors, err_range.clone(), Some(ERR_REASON))?;
598 match replace {
599 EncodeReplace::Str(s) => {
600 if s.as_ref().code_points().any(|c| c.to_u32() > 255) {
601 return Err(ctx.error_encoding(err_range, Some(ERR_REASON)));
602 }
603 out.extend(s.as_ref().code_points().map(|c| c.to_u32() as u8));
604 }
605 EncodeReplace::Bytes(b) => {
606 out.extend_from_slice(b.as_ref());
607 }
608 }
609 }
610 }
611 out.extend_from_slice(ctx.remaining_data().as_bytes());
612 Ok(out)
613 }
614
615 pub fn decode<Ctx: DecodeContext, E: DecodeErrorHandler<Ctx>>(
616 ctx: Ctx,
617 _errors: &E,
618 ) -> Result<(Wtf8Buf, usize), Ctx::Error> {
619 let out: String = ctx.remaining_data().iter().map(|c| *c as char).collect();
620 let out_len = out.len();
621 Ok((out.into(), out_len))
622 }
623}
624
625pub mod ascii {
626 use super::*;
627 use ::ascii::AsciiStr;
628
629 pub const ENCODING_NAME: &str = "ascii";
630
631 const ERR_REASON: &str = "ordinal not in range(128)";
632
633 #[inline]
634 pub fn encode<Ctx, E>(ctx: Ctx, errors: &E) -> Result<Vec<u8>, Ctx::Error>
635 where
636 Ctx: EncodeContext,
637 E: EncodeErrorHandler<Ctx>,
638 {
639 encode_utf8_compatible(ctx, errors, ERR_REASON, StrKind::Ascii)
640 }
641
642 pub fn decode<Ctx: DecodeContext, E: DecodeErrorHandler<Ctx>>(
643 ctx: Ctx,
644 errors: &E,
645 ) -> Result<(Wtf8Buf, usize), Ctx::Error> {
646 decode_utf8_compatible(
647 ctx,
648 errors,
649 |v| {
650 AsciiStr::from_ascii(v).map(|s| s.as_str()).map_err(|e| {
651 unsafe { make_decode_err(v, e.valid_up_to(), Some(1)) }
654 })
655 },
656 |_rest, err_len| HandleResult::Error {
657 err_len,
658 reason: ERR_REASON,
659 },
660 )
661 }
662}