Skip to content

Commit a0491c8

Browse files
authored
gh-158439: Add _PyUnicodeWriter.recheck_maxchar (#158442)
If PyUnicodeWriter_WriteUTF8() or PyUnicodeWriter_Format() fails, PyUnicodeWriter_Finish() now checks if the string uses the most efficient storage. Previously, PyUnicodeWriter_Finish() could return a string with an inefficient storage. * Add _PyUnicodeWriter.recheck_maxchar. PyUnicodeWriter_WriteUTF8() and PyUnicodeWriter_Format() now set recheck_maxchar to 1 on error. * Move error handling from caller to _PyUnicode_DecodeUTF8Writer(). * PyUnicodeWriter_DecodeUTF8Stateful(): rename length parameter to size and rename string parameter to str.
1 parent 115c297 commit a0491c8

6 files changed

Lines changed: 121 additions & 53 deletions

File tree

‎Doc/c-api/unicode.rst‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1945,7 +1945,7 @@ object.
19451945
On success, return ``0``.
19461946
On error, set an exception, leave the writer unchanged, and return ``-1``.
19471947
1948-
.. c:function:: int PyUnicodeWriter_DecodeUTF8Stateful(PyUnicodeWriter *writer, const char *string, Py_ssize_t length, const char *errors, Py_ssize_t *consumed)
1948+
.. c:function:: int PyUnicodeWriter_DecodeUTF8Stateful(PyUnicodeWriter *writer, const char *str, Py_ssize_t size, const char *errors, Py_ssize_t *consumed)
19491949
19501950
Decode the string *str* from UTF-8 with *errors* error handler and write the
19511951
output into *writer*.

‎Include/cpython/unicodeobject.h‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -516,8 +516,8 @@ PyAPI_FUNC(int) PyUnicodeWriter_Format(
516516
...);
517517
PyAPI_FUNC(int) PyUnicodeWriter_DecodeUTF8Stateful(
518518
PyUnicodeWriter *writer,
519-
const char *string, /* UTF-8 encoded string */
520-
Py_ssize_t length, /* size of string */
519+
const char *str, /* UTF-8 encoded string */
520+
Py_ssize_t size, /* size of string */
521521
const char *errors, /* error handling */
522522
Py_ssize_t *consumed); /* bytes consumed */
523523

@@ -538,6 +538,9 @@ typedef struct {
538538
/* minimum character (default: 127, ASCII) */
539539
Py_UCS4 min_char;
540540

541+
// If non-zero, _PyUnicodeWriter_Finish() needs to check maxchar.
542+
int recheck_maxchar;
543+
541544
/* If non-zero, overallocate the buffer (default: 0). */
542545
unsigned char overallocate;
543546

‎Lib/test/test_capi/test_unicode.py‎

Lines changed: 33 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1976,7 +1976,11 @@ def test_invalid_utf8(self):
19761976
writer.write_utf8(b"invalid=\xFF", -1)
19771977

19781978
def test_recover_utf8_error(self):
1979-
# test recovering from PyUnicodeWriter_WriteUTF8() error
1979+
# Recover from PyUnicodeWriter_WriteUTF8() errors. A temporary write
1980+
# changes the buffer kind to UCS-2 before raising UnicodeDecodeError.
1981+
# Then, PyUnicodeWriter_Finish() has to change the buffer kind back to
1982+
# ASCII.
1983+
19801984
writer = self.create_writer(0)
19811985
writer.write_utf8(b"value=", -1)
19821986

@@ -1986,12 +1990,13 @@ def test_recover_utf8_error(self):
19861990
with self.assertRaises(UnicodeDecodeError):
19871991
s = "truncated\u20AC".encode()
19881992
writer.write_utf8(s, len(s) - 1)
1993+
with self.assertRaises(UnicodeDecodeError):
1994+
# Change buffer kind to UCS-2 then raise UnicodeDecodeError
1995+
s = "\u20AC\u20AC".encode()
1996+
writer.write_utf8(s, len(s) - 1)
19891997

1990-
# retry write with a valid string
19911998
writer.write_utf8(b"valid", -1)
1992-
1993-
self.assertEqual(writer.finish(),
1994-
"value=valid")
1999+
self.assertEqual(writer.finish(), "value=valid")
19952000

19962001
def test_decode_utf8(self):
19972002
# test PyUnicodeWriter_DecodeUTF8Stateful()
@@ -2299,6 +2304,29 @@ def test_recover_error(self):
22992304

23002305
self.assertEqual(writer.finish(), 'Hello World.')
23012306

2307+
def test_recheck_maxchar(self):
2308+
# PyUnicodeWriter_Format() changes buffer kind to UCS-2 before raising
2309+
# an exception. Then, PyUnicodeWriter_Finish() has to change the buffer
2310+
# kind back to ASCII.
2311+
from ctypes import py_object
2312+
2313+
class StrError:
2314+
def __str__(self):
2315+
raise RuntimeError("bug")
2316+
2317+
writer = self.create_writer(0)
2318+
# Allocate ASCII buffer
2319+
writer.write_str('ascii')
2320+
2321+
obj = StrError()
2322+
ucs2_utf8 = '\u20ac'.encode()
2323+
with self.assertRaises(RuntimeError):
2324+
# Change buffer kind to UCS-2, but then raise RuntimeError
2325+
self.writer_format(writer, b"%s%S", ucs2_utf8, py_object(obj))
2326+
2327+
writer.write_str('.')
2328+
self.assertEqual(writer.finish(), 'ascii.')
2329+
23022330
def test_readonly_optim(self):
23032331
# Read-only optimization: if the first and only write is a Python str
23042332
# object and no buffer was allocated yet, return the object unchanged
Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,4 @@
1+
If :c:func:`PyUnicodeWriter_WriteUTF8` or :c:func:`PyUnicodeWriter_Format`
2+
fails, :c:func:`PyUnicodeWriter_Finish` now checks if the string uses the most
3+
efficient storage. Previously, :c:func:`PyUnicodeWriter_Finish` could return a
4+
string with an inefficient storage. Patch by Victor Stinner.

‎Objects/unicode_writer.c‎

Lines changed: 43 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -160,7 +160,8 @@ PyUnicodeWriter_Create(Py_ssize_t length)
160160
}
161161

162162

163-
void PyUnicodeWriter_Discard(PyUnicodeWriter *writer)
163+
void
164+
PyUnicodeWriter_Discard(PyUnicodeWriter *writer)
164165
{
165166
if (writer == NULL) {
166167
return;
@@ -511,8 +512,7 @@ PyUnicodeWriter_WriteASCII(PyUnicodeWriter *writer,
511512
assert(writer != NULL);
512513
_Py_AssertHoldsTstate();
513514

514-
_PyUnicodeWriter *priv_writer = (_PyUnicodeWriter*)writer;
515-
return _PyUnicodeWriter_WriteASCIIString(priv_writer, str, size);
515+
return _PyUnicodeWriter_WriteASCIIString((_PyUnicodeWriter*)writer, str, size);
516516
}
517517

518518

@@ -525,40 +525,25 @@ PyUnicodeWriter_WriteUTF8(PyUnicodeWriter *writer,
525525
size = strlen(str);
526526
}
527527

528-
_PyUnicodeWriter *_writer = (_PyUnicodeWriter*)writer;
529-
Py_ssize_t old_pos = _writer->pos;
530-
int res = _PyUnicode_DecodeUTF8Writer(_writer, str, size,
531-
_Py_ERROR_STRICT, NULL, NULL);
532-
if (res < 0) {
533-
_writer->pos = old_pos;
534-
}
535-
return res;
528+
return _PyUnicode_DecodeUTF8Writer((_PyUnicodeWriter*)writer, str, size,
529+
_Py_ERROR_STRICT, NULL, NULL);
536530
}
537531

538532

539533
int
540534
PyUnicodeWriter_DecodeUTF8Stateful(PyUnicodeWriter *writer,
541-
const char *string,
542-
Py_ssize_t length,
535+
const char *str,
536+
Py_ssize_t size,
543537
const char *errors,
544538
Py_ssize_t *consumed)
545539
{
546-
if (length < 0) {
547-
length = strlen(string);
548-
}
549-
550-
_PyUnicodeWriter *_writer = (_PyUnicodeWriter*)writer;
551-
Py_ssize_t old_pos = _writer->pos;
552-
int res = _PyUnicode_DecodeUTF8Writer(_writer, string, length,
553-
_Py_ERROR_UNKNOWN, errors,
554-
consumed);
555-
if (res < 0) {
556-
_writer->pos = old_pos;
557-
if (consumed) {
558-
*consumed = 0;
559-
}
540+
if (size < 0) {
541+
size = strlen(str);
560542
}
561-
return res;
543+
544+
return _PyUnicode_DecodeUTF8Writer((_PyUnicodeWriter*)writer, str, size,
545+
_Py_ERROR_UNKNOWN, errors,
546+
consumed);
562547
}
563548

564549

@@ -600,30 +585,55 @@ _PyUnicodeWriter_Finish(_PyUnicodeWriter *writer)
600585

601586
Py_ssize_t final_size = writer->pos;
602587
if (final_size == 0) {
588+
// Get the empty string singleton
603589
PyObject *empty = _PyUnicode_GetEmpty();
604590
Py_XDECREF(str); // writer->buffer can be NULL if the position is 0
605591
return empty;
606592
}
607593

608-
Py_ssize_t length = PyUnicode_GET_LENGTH(str);
594+
if (writer->readonly) {
595+
assert(final_size == PyUnicode_GET_LENGTH(str));
596+
goto done;
597+
}
598+
assert(final_size <= PyUnicode_GET_LENGTH(str));
599+
609600
if (final_size == 1 && PyUnicode_KIND(str) == PyUnicode_1BYTE_KIND) {
610-
assert(length >= 1);
601+
// Get the single character singleton
602+
assert(PyUnicode_GET_LENGTH(str) >= 1);
611603
const Py_UCS1 *data = PyUnicode_1BYTE_DATA(str);
612604
Py_UCS1 ch = data[0];
613-
PyObject *latin1_char = _Py_LATIN1_CHR(ch);
614605
Py_DECREF(str);
615-
return latin1_char;
606+
str = _Py_LATIN1_CHR(ch);
607+
goto done;
608+
}
609+
610+
if (writer->recheck_maxchar) {
611+
Py_UCS4 maxchar = _PyUnicode_FindMaxChar(str, 0, final_size);
612+
if (maxchar != writer->maxchar) {
613+
// Adjust the string kind
614+
PyObject *str2 = PyUnicode_New(final_size, maxchar);
615+
if (str2 == NULL) {
616+
Py_DECREF(str);
617+
return NULL;
618+
}
619+
_PyUnicode_FastCopyCharacters(str2, 0, str, 0, final_size);
620+
Py_SETREF(str, str2);
621+
goto done;
622+
}
616623
}
617624

618-
if (!writer->readonly && length != final_size) {
625+
if (PyUnicode_GET_LENGTH(str) != final_size) {
626+
// Truncate the string
619627
PyObject *str2 = _PyUnicode_ResizeCompact(str, final_size);
620628
if (str2 == NULL) {
621629
Py_DECREF(str);
622630
return NULL;
623631
}
624632
str = str2;
633+
goto done;
625634
}
626635

636+
done:
627637
assert(_PyUnicode_CheckConsistency(str, 1));
628638
return str;
629639
}

‎Objects/unicodeobject.c‎

Lines changed: 35 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -3166,16 +3166,21 @@ PyUnicodeWriter_Format(PyUnicodeWriter *writer, const char *format, ...)
31663166
}
31673167

31683168
int
3169-
_PyUnicodeWriter_FormatV(PyUnicodeWriter *writer, const char *format,
3169+
_PyUnicodeWriter_FormatV(PyUnicodeWriter *pub_writer, const char *format,
31703170
va_list vargs)
31713171
{
3172-
_PyUnicodeWriter *_writer = (_PyUnicodeWriter*)writer;
3173-
Py_ssize_t old_pos = _writer->pos;
3172+
_PyUnicodeWriter *writer = (_PyUnicodeWriter*)pub_writer;
3173+
Py_ssize_t old_pos = writer->pos;
3174+
Py_UCS4 old_maxchar = writer->maxchar;
31743175

3175-
int res = unicode_from_format(_writer, format, vargs);
3176+
int res = unicode_from_format(writer, format, vargs);
31763177

31773178
if (res < 0) {
3178-
_writer->pos = old_pos;
3179+
writer->pos = old_pos;
3180+
if (writer->maxchar > old_maxchar) {
3181+
// _PyUnicodeWriter_Finish() will check maxchar
3182+
writer->recheck_maxchar = 1;
3183+
}
31793184
}
31803185
return res;
31813186
}
@@ -5364,7 +5369,7 @@ unicode_decode_utf8(const char *s, Py_ssize_t size,
53645369
}
53655370

53665371

5367-
// Used by PyUnicodeWriter_WriteUTF8() implementation
5372+
// Used by PyUnicodeWriter_WriteUTF8() and PyUnicodeWriter_DecodeUTF8Stateful()
53685373
int
53695374
_PyUnicode_DecodeUTF8Writer(_PyUnicodeWriter *writer,
53705375
const char *s, Py_ssize_t size,
@@ -5378,18 +5383,20 @@ _PyUnicode_DecodeUTF8Writer(_PyUnicodeWriter *writer,
53785383
return 0;
53795384
}
53805385

5386+
Py_ssize_t old_pos = writer->pos;
5387+
Py_UCS4 old_maxchar = writer->maxchar;
5388+
53815389
// fast path: try ASCII string.
53825390
if (_PyUnicodeWriter_Prepare(writer, size, 127) < 0) {
5383-
return -1;
5391+
goto error;
53845392
}
53855393
assert(_PyUnicodeWriter_CanWrite(writer));
53865394

53875395
const char *starts = s;
53885396
const char *end = s + size;
5389-
Py_ssize_t decoded = 0;
5390-
Py_UCS1 *dest = (Py_UCS1*)writer->data + writer->pos * writer->kind;
53915397
if (writer->kind == PyUnicode_1BYTE_KIND) {
5392-
decoded = ascii_decode(s, end, dest);
5398+
Py_UCS1 *dest = (Py_UCS1*)writer->data + writer->pos * writer->kind;
5399+
Py_ssize_t decoded = ascii_decode(s, end, dest);
53935400
writer->pos += decoded;
53945401

53955402
if (decoded == size) {
@@ -5401,8 +5408,24 @@ _PyUnicode_DecodeUTF8Writer(_PyUnicodeWriter *writer,
54015408
s += decoded;
54025409
}
54035410

5404-
return unicode_decode_utf8_impl(writer, starts, s, end,
5405-
error_handler, errors, consumed);
5411+
int res = unicode_decode_utf8_impl(writer, starts, s, end,
5412+
error_handler, errors, consumed);
5413+
if (res < 0) {
5414+
goto error;
5415+
}
5416+
return 0;
5417+
5418+
error:
5419+
// Restore the writer to its previous state
5420+
writer->pos = old_pos;
5421+
if (writer->maxchar > old_maxchar) {
5422+
// _PyUnicodeWriter_Finish() will check maxchar
5423+
writer->recheck_maxchar = 1;
5424+
}
5425+
if (consumed) {
5426+
*consumed = 0;
5427+
}
5428+
return -1;
54065429
}
54075430

54085431

0 commit comments

Comments
 (0)