Skip to content

Commit d9de16a

Browse files
committed
encoding
1 parent 884eba1 commit d9de16a

1 file changed

Lines changed: 39 additions & 45 deletions

File tree

src/pybase64/_pybase64.c

Lines changed: 39 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -599,13 +599,16 @@ static int decode_slow(const uint8_t *src, size_t srclen, uint8_t* out, size_t*
599599

600600
static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buffer, char const* alphabet, Py_ssize_t wrapcol, unsigned int flags)
601601
{
602+
size_t groups;
603+
size_t groups_remainder;
602604
size_t out_len;
603605
PyObject* out_object;
604606
#if PY_VERSION_HEX >= 0x030f0000
605607
PyBytesWriter* writer;
606608
#endif
607609
char* dst_start;
608610
char* dst;
611+
struct base64_state b64_state;
609612
pybase64_state *state = (pybase64_state*)PyModule_GetState(self);
610613
if (state == NULL) { /* GCOVR_EXCL_BR_WITHOUT_HIT: 1/2 */
611614
return NULL; /* GCOVR_EXCL_LINE */
@@ -624,7 +627,20 @@ static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buff
624627
return PyErr_NoMemory(); /* GCOVR_EXCL_LINE */
625628
}
626629

627-
out_len = (size_t)(((buffer->len + 2) / 3) * 4);
630+
groups = (size_t)(buffer->len / 3);
631+
groups_remainder = (size_t)buffer->len - groups * 3U;
632+
out_len = groups * 4;
633+
switch (groups_remainder)
634+
{
635+
case 1:
636+
out_len += (flags & PYBASE64_FLAGS_NO_PADDING) ? 2U : 4U;
637+
break;
638+
case 2:
639+
out_len += (flags & PYBASE64_FLAGS_NO_PADDING) ? 3U : 4U;
640+
break;
641+
default:
642+
break;
643+
}
628644
if (wrapcol > 0 && out_len > 0) {
629645
size_t newlines = (out_len - 1U) / (size_t)wrapcol;
630646
if (newlines > ((size_t)PY_SSIZE_T_MAX - out_len)) { /* GCOVR_EXCL_BR_WITHOUT_HIT: 1/2 */
@@ -679,7 +695,10 @@ static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buff
679695
/* not interacting with Python objects from here, release the GIL */
680696
Py_BEGIN_ALLOW_THREADS
681697

682-
int const libbase64_simd_flag = state->libbase64_simd_flag;
698+
int const b64_flags = state->libbase64_simd_flag |
699+
((flags & PYBASE64_FLAGS_NO_PADDING) ? BASE64_NO_PADDING : 0);
700+
701+
base64_stream_encode_init(&b64_state, b64_flags);
683702

684703
if (flags & PYBASE64_FLAGS_APPEND_NEW_LINE) {
685704
out_len--; /* only consider len without new line terminator */
@@ -690,13 +709,14 @@ static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buff
690709
const Py_ssize_t src_slice = (Py_ssize_t)((dst_slice / 4U) * 3U);
691710
Py_ssize_t len = buffer->len;
692711
const char* src = (const char*)buffer->buf;
693-
size_t remainder;
694712

695713
if (alphabet) {
714+
size_t remainder;
715+
696716
while (out_len > dst_slice) {
697717
size_t dst_len = (size_t)wrapcol;
698718

699-
base64_encode(src, src_slice, dst, &dst_len, libbase64_simd_flag);
719+
base64_stream_encode(&b64_state, src, src_slice, dst, &dst_len);
700720
translate_inplace(dst, dst_len, alphabet);
701721
dst[dst_len] = '\n';
702722

@@ -705,25 +725,28 @@ static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buff
705725
out_len -= dst_slice;
706726
dst += dst_slice;
707727
}
728+
base64_stream_encode(&b64_state, src, len, dst, &out_len);
708729
remainder = out_len;
709-
base64_encode(src, len, dst, &remainder, libbase64_simd_flag);
730+
base64_stream_encode_final(&b64_state, dst + out_len, &out_len);
731+
remainder += out_len;
710732
translate_inplace(dst, remainder, alphabet);
711733
dst += remainder;
712734
}
713735
else {
714736
while (out_len > dst_slice) {
715737
size_t dst_len = (size_t)wrapcol;
716-
base64_encode(src, src_slice, dst, &dst_len, libbase64_simd_flag);
738+
base64_stream_encode(&b64_state, src, src_slice, dst, &dst_len);
717739
dst[dst_len] = '\n';
718740

719741
len -= src_slice;
720742
src += src_slice;
721743
out_len -= dst_slice;
722744
dst += dst_slice;
723745
}
724-
remainder = out_len;
725-
base64_encode(src, len, dst, &remainder, libbase64_simd_flag);
726-
dst += remainder;
746+
base64_stream_encode(&b64_state, src, len, dst, &out_len);
747+
dst += out_len;
748+
base64_stream_encode_final(&b64_state, dst, &out_len);
749+
dst += out_len;
727750
}
728751
}
729752
else if (alphabet) {
@@ -737,32 +760,26 @@ static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buff
737760
while (out_len > dst_slice) {
738761
size_t dst_len = dst_slice;
739762

740-
base64_encode(src, src_slice, dst, &dst_len, libbase64_simd_flag);
763+
base64_stream_encode(&b64_state, src, src_slice, dst, &dst_len);
741764
translate_inplace(dst, dst_slice, alphabet);
742765

743766
len -= src_slice;
744767
src += src_slice;
745768
out_len -= dst_slice;
746769
dst += dst_slice;
747770
}
771+
base64_stream_encode(&b64_state, src, len, dst, &out_len);
748772
remainder = out_len;
749-
base64_encode(src, len, dst, &out_len, libbase64_simd_flag);
773+
base64_stream_encode_final(&b64_state, dst + out_len, &out_len);
774+
remainder += out_len;
750775
translate_inplace(dst, remainder, alphabet);
751776
dst += remainder;
752777
}
753778
else {
754-
base64_encode(buffer->buf, buffer->len, dst, &out_len, libbase64_simd_flag);
779+
base64_stream_encode(&b64_state, buffer->buf, buffer->len, dst, &out_len);
780+
dst += out_len;
781+
base64_stream_encode_final(&b64_state, dst, &out_len);
755782
dst += out_len;
756-
}
757-
if (flags & PYBASE64_FLAGS_NO_PADDING) {
758-
/* we have at least 4 bytes, at most 2 '=' */
759-
assert((dst - dst_start) >= 4);
760-
if (dst[-1] == '=') {
761-
dst -= 1;
762-
}
763-
if (dst[-1] == '=') {
764-
dst -= 1;
765-
}
766783
}
767784
if (flags & PYBASE64_FLAGS_APPEND_NEW_LINE) {
768785
*dst++ = '\n';
@@ -771,29 +788,6 @@ static PyObject* pybase64_encode_impl_core(PyObject* self, Py_buffer const* buff
771788
/* restore the GIL */
772789
Py_END_ALLOW_THREADS
773790

774-
if (flags & PYBASE64_FLAGS_NO_PADDING)
775-
{
776-
if (flags & PYBASE64_FLAGS_ENCODE_AS_STRING) {
777-
#if defined(PYPY_VERSION) || defined(GRAALVM_PYTHON)
778-
/* SystemError: PyUnicode_Resize called on already created string... */
779-
/* we'll be less efficient */
780-
PyObject* temp_object = PyUnicode_FromKindAndData(PyUnicode_1BYTE_KIND, dst_start, dst - dst_start);
781-
Py_DECREF(out_object);
782-
out_object = temp_object;
783-
#else
784-
if (PyUnicode_Resize(&out_object, dst - dst_start) != 0) { /* GCOVR_EXCL_BR_WITHOUT_HIT: 1/2 */
785-
Py_DECREF(out_object); /* GCOVR_EXCL_LINE */
786-
return NULL; /* GCOVR_EXCL_LINE */
787-
}
788-
#endif
789-
}
790-
else {
791-
#if PY_VERSION_HEX < 0x030f0000
792-
_PyBytes_Resize(&out_object, dst - dst_start);
793-
#endif
794-
}
795-
}
796-
797791
#if PY_VERSION_HEX >= 0x030f0000
798792
if (!(flags & PYBASE64_FLAGS_ENCODE_AS_STRING)) {
799793
out_object = PyBytesWriter_FinishWithPointer(writer, dst);

0 commit comments

Comments
 (0)