Simplify API of streaming base64 decoder further

This commit is contained in:
Kovid Goyal
2024-07-29 21:24:45 +05:30
parent 212d7accfc
commit b52275e0b5
4 changed files with 26 additions and 85 deletions

View File

@@ -107,100 +107,51 @@ pybase64_decode(PyObject UNUSED *self, PyObject *args) {
typedef struct StreamingBase64Decoder {
PyObject_HEAD
PyObject *output;
size_t output_sz, output_capacity, initial_capacity;
struct base64_state state;
} StreamingBase64Decoder;
static int
StreamingBase64Decoder_init(PyObject *s, PyObject *args, PyObject *kwds) {
static char *kwlist[] = {"initial_capacity", NULL};
unsigned long initial_capacity = 8 * 1024;
if (!PyArg_ParseTupleAndKeywords(args, kwds, "|k", kwlist, &initial_capacity)) return -1;
StreamingBase64Decoder_init(PyObject *s, PyObject *args UNUSED, PyObject *kwds UNUSED) {
if (PyTuple_GET_SIZE(args)) { PyErr_SetString(PyExc_TypeError, "constructor takes no arguments"); return -1; }
StreamingBase64Decoder *self = (StreamingBase64Decoder*)s;
self->initial_capacity = initial_capacity;
base64_stream_decode_init(&self->state, 0);
return 0;
}
static void
StreamingBase64Decoder_dealloc(PyObject *self) {
StreamingBase64Decoder *h = (StreamingBase64Decoder*)self;
Py_CLEAR(h->output);
Py_TYPE(self)->tp_free(self);
}
static bool
write_base64_data(StreamingBase64Decoder *self, const void *data, const size_t len) {
if (!len) return true;
size_t sz = required_buffer_size_for_base64_decode(len);
if ((self->output_sz + sz) > self->output_capacity) {
size_t cap = MAX(self->output_capacity * 2, MAX(self->output_sz + sz, self->initial_capacity));
if (self->output) { if (_PyBytes_Resize(&self->output, cap) != 0) return false; }
else { self->output = PyBytes_FromStringAndSize(NULL, cap); if (!self->output) return false; }
self->output_capacity = cap;
}
bool ok = base64_stream_decode(&self->state, data, len, PyBytes_AS_STRING(self->output) + self->output_sz, &sz);
if (!ok) {
PyErr_SetString(PyExc_ValueError, "Invalid base64 input data");
return false;
}
self->output_sz += sz;
return true;
}
static PyObject*
StreamingBase64Decoder_add(StreamingBase64Decoder *self, PyObject *a) {
StreamingBase64Decoder_decode(StreamingBase64Decoder *self, PyObject *a) {
RAII_PY_BUFFER(data);
if (PyObject_GetBuffer(a, &data, PyBUF_SIMPLE) != 0) return NULL;
if (!data.buf || !data.len) return PyLong_FromLong(0);
size_t before = self->output_sz;
if (!write_base64_data(self, data.buf, data.len)) return NULL;
return PyLong_FromSize_t(self->output_sz - before);
if (!data.buf || !data.len) return PyBytes_FromStringAndSize(NULL, 0);
size_t sz = required_buffer_size_for_base64_decode(data.len);
RAII_PyObject(ans, PyBytes_FromStringAndSize(NULL, sz));
if (!ans) return NULL;
if (!base64_stream_decode(&self->state, data.buf, data.len, PyBytes_AS_STRING(ans), &sz)) {
PyErr_SetString(PyExc_ValueError, "Invalid base64 input data");
return NULL;
}
if (_PyBytes_Resize(&ans, sz) != 0) return NULL;
return Py_NewRef(ans);
}
static Py_ssize_t
StreamingBase64Decoder_len(PyObject *s) { return ((StreamingBase64Decoder*)s)->output_sz; }
static PyObject*
StreamingBase64Decoder_reinitialize(StreamingBase64Decoder *self, PyObject *args UNUSED) {
StreamingBase64Decoder_reset(StreamingBase64Decoder *self, PyObject *args UNUSED) {
base64_stream_decode_init(&self->state, 0);
Py_RETURN_NONE;
}
static PyObject*
StreamingBase64Decoder_copy_output(StreamingBase64Decoder *self, PyObject *args UNUSED) {
return PyBytes_FromStringAndSize(PyBytes_AS_STRING(self->output), self->output_sz);
}
static PyObject*
StreamingBase64Decoder_take_output(StreamingBase64Decoder *self, PyObject *args UNUSED) {
if (!self->output_sz) return PyBytes_FromStringAndSize(NULL, 0);
if (_PyBytes_Resize(&self->output, self->output_sz) != 0) return NULL;
PyObject *ans = self->output;
self->output = NULL; self->output_sz = 0; self->output_capacity = 0;
return ans;
}
static PyTypeObject StreamingBase64Decoder_Type = {
PyVarObject_HEAD_INIT(NULL, 0)
.tp_name = "kitty.fast_data_types.StreamingBase64Decoder",
.tp_basicsize = sizeof(StreamingBase64Decoder),
.tp_dealloc = StreamingBase64Decoder_dealloc,
.tp_flags = Py_TPFLAGS_DEFAULT,
.tp_doc = "StreamingBase64Decoder",
.tp_methods = (PyMethodDef[]){
{"add", (PyCFunction)StreamingBase64Decoder_add, METH_O, ""},
{"reinitialize", (PyCFunction)StreamingBase64Decoder_reinitialize, METH_NOARGS, ""},
{"take_output", (PyCFunction)StreamingBase64Decoder_take_output, METH_NOARGS, ""},
{"copy_output", (PyCFunction)StreamingBase64Decoder_copy_output, METH_NOARGS, ""},
{"decode", (PyCFunction)StreamingBase64Decoder_decode, METH_O, ""},
{"reset", (PyCFunction)StreamingBase64Decoder_reset, METH_NOARGS, ""},
{NULL, NULL, 0, NULL},
},
.tp_new = PyType_GenericNew,
.tp_init = StreamingBase64Decoder_init,
.tp_as_sequence = &(PySequenceMethods){
.sq_length = StreamingBase64Decoder_len,
},
};
static PyObject*