Skip to content

Commit 1a85213

Browse files
authored
gh-158585: Adjust allocation in bytes_fromformat() (#158621)
Take in account preallocated bytes when calling PyBytesWriter_GrowAndUpdatePointer(). Use sprintf() return value instead of calling strlen().
1 parent b4197da commit 1a85213

2 files changed

Lines changed: 96 additions & 56 deletions

File tree

‎Lib/test/test_capi/test_bytes.py‎

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -607,15 +607,45 @@ def test_format_i(self):
607607

608608
writer = self.create_writer()
609609
writer.format_i(b'x=%i, ', 123)
610-
writer.format_i(b'y=%i', 456)
610+
writer.format_i(b'y=%d', 456)
611611
self.bytes_equal(writer.finish(), b'x=123, y=456')
612612

613+
# Make sure that %i allocates the right number of bytes
614+
n_str = '1234567'
615+
for i in range(1, len(n_str) + 1):
616+
n = int(n_str[:i])
617+
writer = self.create_writer(0)
618+
writer.format_i(b'n=%i', n)
619+
self.bytes_equal(writer.finish(), f'n={n}'.encode())
620+
621+
# Test invalid format: "%t" is invalid, stop the parser and copy
622+
# remaining bytes
623+
writer = self.create_writer()
624+
writer.write_bytes(b'prefix ', -1)
625+
writer.format_i(b'y=%._ t %i', 0)
626+
writer.write_bytes(b'.', 1)
627+
self.bytes_equal(writer.finish(), b'prefix y=%._ t %i.')
628+
613629
def test_format_s(self):
614630
# Test PyBytesWriter_Format()
615631
writer = self.create_writer()
616632
writer.format_s(b's=%s', b'Hello World')
617633
self.bytes_equal(writer.finish(), b's=Hello World')
618634

635+
# Test precision
636+
writer = self.create_writer()
637+
writer.format_s(b'%.6s', b'Monty Python')
638+
writer.format_s(b'%.20s', b'3.0')
639+
self.bytes_equal(writer.finish(), b'Monty 3.0')
640+
641+
# Make sure that %s allocates the right number of bytes
642+
text = b'1234567'
643+
for i in range(1, len(text) + 1):
644+
s = text[:i]
645+
writer = self.create_writer(0)
646+
writer.format_s(b's=%s', s)
647+
self.bytes_equal(writer.finish(), b's=' + s)
648+
619649
@support.nomemtest
620650
def test_format_s_memory_error(self):
621651
writer = self.create_writer()

‎Objects/bytesobject.c‎

Lines changed: 65 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -194,73 +194,77 @@ static char*
194194
bytes_fromformat(PyBytesWriter *writer, Py_ssize_t writer_pos,
195195
const char *format, va_list vargs)
196196
{
197-
const char *f;
198-
const char *p;
199-
Py_ssize_t prec;
200-
int longflag;
201-
int size_tflag;
202197
/* Longest 64-bit formatted numbers:
203198
- "18446744073709551615\0" (21 bytes)
204199
- "-9223372036854775808\0" (21 bytes)
205200
Decimal takes the most space (it isn't enough for octal.)
206201
207202
Longest 64-bit pointer representation:
208-
"0xffffffffffffffff\0" (19 bytes). */
203+
"0xffffffffffffffff\0" (19 bytes).
204+
205+
Longest 64-bit "%p" with "0x" prefix: len(hex(2**64-1)+'\0') = 19.
206+
*/
209207
char buffer[21];
210208

211209
char *s = (char*)PyBytesWriter_GetData(writer) + writer_pos;
212210

213-
#define WRITE_BYTES_LEN(str, len_expr) \
211+
#define WRITE_BYTES(str, len_expr) \
214212
do { \
215-
size_t len = (len_expr); \
216-
s = PyBytesWriter_GrowAndUpdatePointer(writer, len, s); \
217-
if (s == NULL) { \
218-
return NULL; \
213+
size_t _len = (len_expr); \
214+
size_t _prealloc = (f - p + 1); \
215+
if (_len > _prealloc) { \
216+
s = PyBytesWriter_GrowAndUpdatePointer(writer, _len - _prealloc, s); \
217+
if (s == NULL) { \
218+
return NULL; \
219+
} \
219220
} \
220-
memcpy(s, (str), len); \
221-
s += len; \
221+
memcpy(s, (str), _len); \
222+
s += _len; \
222223
} while (0)
223-
#define WRITE_BYTES(str) WRITE_BYTES_LEN(str, strlen(str))
224224

225-
for (f = format; *f; f++) {
225+
for (const char *f = format; *f; f++) {
226226
if (*f != '%') {
227227
*s++ = *f;
228228
continue;
229229
}
230230

231-
p = f++;
231+
const char *p = f++;
232232

233233
/* ignore the width (ex: 10 in "%10s") */
234234
while (Py_ISDIGIT(*f))
235235
f++;
236236

237237
/* parse the precision (ex: 10 in "%.10s") */
238-
prec = 0;
238+
Py_ssize_t prec = 0;
239239
if (*f == '.') {
240240
f++;
241241
for (; Py_ISDIGIT(*f); f++) {
242242
prec = (prec * 10) + (*f - '0');
243243
}
244244
}
245+
assert(prec >= 0);
245246

246-
while (*f && *f != '%' && !Py_ISALPHA(*f))
247+
while (*f && *f != '%' && !Py_ISALPHA(*f)) {
247248
f++;
249+
}
248250

249251
/* handle the long flag ('l'), but only for %ld and %lu.
250252
others can be added when necessary. */
251-
longflag = 0;
253+
int longflag = 0;
252254
if (*f == 'l' && (f[1] == 'd' || f[1] == 'u')) {
253255
longflag = 1;
254256
++f;
255257
}
256258

257259
/* handle the size_t flag ('z'). */
258-
size_tflag = 0;
260+
int size_tflag = 0;
259261
if (*f == 'z' && (f[1] == 'd' || f[1] == 'u')) {
260262
size_tflag = 1;
261263
++f;
262264
}
263265

266+
Py_ssize_t len;
267+
264268
switch (*f) {
265269
case 'c':
266270
{
@@ -277,91 +281,97 @@ bytes_fromformat(PyBytesWriter *writer, Py_ssize_t writer_pos,
277281

278282
case 'd':
279283
if (longflag) {
280-
sprintf(buffer, "%ld", va_arg(vargs, long));
284+
len = sprintf(buffer, "%ld", va_arg(vargs, long));
281285
}
282286
else if (size_tflag) {
283-
sprintf(buffer, "%zd", va_arg(vargs, Py_ssize_t));
287+
len = sprintf(buffer, "%zd", va_arg(vargs, Py_ssize_t));
284288
}
285289
else {
286-
sprintf(buffer, "%d", va_arg(vargs, int));
290+
len = sprintf(buffer, "%d", va_arg(vargs, int));
287291
}
288-
assert(strlen(buffer) < sizeof(buffer));
289-
WRITE_BYTES(buffer);
292+
assert(1 <= len && len < (Py_ssize_t)sizeof(buffer));
293+
WRITE_BYTES(buffer, len);
290294
break;
291295

292296
case 'u':
293297
if (longflag) {
294-
sprintf(buffer, "%lu", va_arg(vargs, unsigned long));
298+
len = sprintf(buffer, "%lu", va_arg(vargs, unsigned long));
295299
}
296300
else if (size_tflag) {
297-
sprintf(buffer, "%zu", va_arg(vargs, size_t));
301+
len = sprintf(buffer, "%zu", va_arg(vargs, size_t));
298302
}
299303
else {
300-
sprintf(buffer, "%u", va_arg(vargs, unsigned int));
304+
len = sprintf(buffer, "%u", va_arg(vargs, unsigned int));
301305
}
302-
assert(strlen(buffer) < sizeof(buffer));
303-
WRITE_BYTES(buffer);
306+
assert(1 <= len && len < (Py_ssize_t)sizeof(buffer));
307+
WRITE_BYTES(buffer, len);
304308
break;
305309

306310
case 'i':
307-
sprintf(buffer, "%i", va_arg(vargs, int));
308-
assert(strlen(buffer) < sizeof(buffer));
309-
WRITE_BYTES(buffer);
311+
len = sprintf(buffer, "%i", va_arg(vargs, int));
312+
assert(1 <= len && len < (Py_ssize_t)sizeof(buffer));
313+
WRITE_BYTES(buffer, len);
310314
break;
311315

312316
case 'x':
313-
sprintf(buffer, "%x", va_arg(vargs, int));
314-
assert(strlen(buffer) < sizeof(buffer));
315-
WRITE_BYTES(buffer);
317+
len = sprintf(buffer, "%x", va_arg(vargs, int));
318+
assert(1 <= len && len < (Py_ssize_t)sizeof(buffer));
319+
WRITE_BYTES(buffer, len);
316320
break;
317321

318322
case 's':
319323
{
320-
Py_ssize_t i;
321-
322-
p = va_arg(vargs, const char*);
323-
if (prec <= 0) {
324-
i = strlen(p);
324+
const char *str = va_arg(vargs, const char*);
325+
if (prec == 0) {
326+
len = strlen(str);
325327
}
326328
else {
327-
i = 0;
328-
while (i < prec && p[i]) {
329-
i++;
329+
const char *end = memchr(str, 0, prec);
330+
if (end != NULL) {
331+
len = (end - str);
332+
}
333+
else {
334+
len = (size_t)prec;
330335
}
331336
}
332-
WRITE_BYTES_LEN(p, i);
337+
WRITE_BYTES(str, len);
333338
break;
334339
}
335340

336341
case 'p':
337-
sprintf(buffer, "%p", va_arg(vargs, void*));
338-
assert(strlen(buffer) < sizeof(buffer));
342+
len = sprintf(buffer, "%p", va_arg(vargs, void*));
343+
assert(1 <= len && len < (Py_ssize_t)sizeof(buffer));
339344
/* %p is ill-defined: ensure leading 0x. */
340-
if (buffer[1] == 'X')
345+
if (buffer[1] == 'X') {
341346
buffer[1] = 'x';
347+
}
342348
else if (buffer[1] != 'x') {
343-
memmove(buffer+2, buffer, strlen(buffer)+1);
349+
memmove(buffer + 2, buffer, len + 1);
344350
buffer[0] = '0';
345351
buffer[1] = 'x';
352+
len += 2;
346353
}
347-
WRITE_BYTES(buffer);
354+
WRITE_BYTES(buffer, len);
348355
break;
349356

350357
case '%':
351358
*s++ = '%';
352359
break;
353360

354361
default:
355-
/* invalid format string: copy unformatted string and exit */
356-
WRITE_BYTES(p);
362+
// Invalid format string: copy unformatted string and exit.
363+
// No need to grow the writer buffer, we already preallocated
364+
// enough bytes.
365+
len = strlen(p);
366+
memcpy(s, p, len);
367+
s += len;
357368
return s;
358369
}
359370
}
360371

361-
#undef WRITE_BYTES
362-
#undef WRITE_BYTES_LEN
363-
364372
return s;
373+
374+
#undef WRITE_BYTES
365375
}
366376

367377

0 commit comments

Comments
 (0)