Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 4 additions & 9 deletions ext/zstdruby/common.h
Original file line number Diff line number Diff line change
Expand Up @@ -10,15 +10,12 @@

extern VALUE rb_cCDict, rb_cDDict;

static int convert_compression_level(ZSTD_CCtx* ctx, VALUE compression_level_value)
static int convert_compression_level(VALUE compression_level_value)
{
if (NIL_P(compression_level_value)) {
return ZSTD_CLEVEL_DEFAULT;
}
if (!RB_INTEGER_TYPE_P(compression_level_value)) {
if (ctx) {
ZSTD_freeCCtx(ctx);
}
rb_raise(rb_eTypeError, "compression level must be an Integer");
}
return NUM2INT(compression_level_value);
Expand All @@ -27,7 +24,8 @@ static int convert_compression_level(ZSTD_CCtx* ctx, VALUE compression_level_val
/* Returns the Zstd::CDict given as `dict:`, or Qnil. ZSTD_CCtx_refCDict only
borrows the pointer, so a caller that keeps the ZSTD_CCtx alive beyond this
call has to keep the returned object reachable for just as long. A String
dictionary needs no such handling: ZSTD_CCtx_loadDictionary copies it. */
dictionary needs no such handling: ZSTD_CCtx_loadDictionary copies it.
Raises without freeing ctx: the caller owns it and has to release it. */
static VALUE set_compress_params(ZSTD_CCtx* const ctx, VALUE kwargs)
{
ID kwargs_keys[2];
Expand All @@ -38,7 +36,7 @@ static VALUE set_compress_params(ZSTD_CCtx* const ctx, VALUE kwargs)

int compression_level = ZSTD_CLEVEL_DEFAULT;
if (kwargs_values[0] != Qundef && kwargs_values[0] != Qnil) {
compression_level = convert_compression_level(ctx, kwargs_values[0]);
compression_level = convert_compression_level(kwargs_values[0]);
}
ZSTD_CCtx_setParameter(ctx, ZSTD_c_compressionLevel, compression_level);

Expand All @@ -47,7 +45,6 @@ static VALUE set_compress_params(ZSTD_CCtx* const ctx, VALUE kwargs)
ZSTD_CDict* cdict = DATA_PTR(kwargs_values[1]);
size_t ref_dict_ret = ZSTD_CCtx_refCDict(ctx, cdict);
if (ZSTD_isError(ref_dict_ret)) {
ZSTD_freeCCtx(ctx);
rb_raise(rb_eRuntimeError, "%s", "ZSTD_CCtx_refCDict failed");
}
return kwargs_values[1];
Expand All @@ -56,11 +53,9 @@ static VALUE set_compress_params(ZSTD_CCtx* const ctx, VALUE kwargs)
size_t dict_size = RSTRING_LEN(kwargs_values[1]);
size_t load_dict_ret = ZSTD_CCtx_loadDictionary(ctx, dict_buffer, dict_size);
if (ZSTD_isError(load_dict_ret)) {
ZSTD_freeCCtx(ctx);
rb_raise(rb_eRuntimeError, "%s", "ZSTD_CCtx_loadDictionary failed");
}
} else {
ZSTD_freeCCtx(ctx);
rb_raise(rb_eArgError, "`dict:` must be a Zstd::CDict or a String");
}
}
Expand Down
4 changes: 2 additions & 2 deletions ext/zstdruby/streaming_compress.c
Original file line number Diff line number Diff line change
Expand Up @@ -91,9 +91,9 @@ rb_streaming_compress_initialize(int argc, VALUE *argv, VALUE obj)
if (ctx == NULL) {
rb_raise(rb_eRuntimeError, "%s", "ZSTD_createCCtx error");
}
VALUE dict = set_compress_params(ctx, kwargs);

/* Before set_compress_params, which can raise: the free callback owns it. */
sc->ctx = ctx;
VALUE dict = set_compress_params(ctx, kwargs);
RB_OBJ_WRITE(obj, &sc->dict, dict);
RB_OBJ_WRITE(obj, &sc->buf, rb_str_new(NULL, buffOutSize));
sc->buf_size = buffOutSize;
Expand Down
54 changes: 38 additions & 16 deletions ext/zstdruby/zstdruby.c
Original file line number Diff line number Diff line change
Expand Up @@ -8,30 +8,25 @@ static VALUE zstdVersion(VALUE self)
return INT2NUM(version);
}

static VALUE rb_compress(int argc, VALUE *argv, VALUE self)
{
struct compress_args {
ZSTD_CCtx* ctx;
VALUE input_value;
VALUE kwargs;
rb_scan_args(argc, argv, "10:", &input_value, &kwargs);

StringValue(input_value);

ZSTD_CCtx* const ctx = ZSTD_createCCtx();
if (ctx == NULL) {
rb_raise(rb_eRuntimeError, "%s", "ZSTD_createCCtx error");
}
};

set_compress_params(ctx, kwargs);
static VALUE compress_body(VALUE arg)
{
struct compress_args* args = (struct compress_args*)arg;
set_compress_params(args->ctx, args->kwargs);

char* input_data = RSTRING_PTR(input_value);
size_t input_size = RSTRING_LEN(input_value);
char* input_data = RSTRING_PTR(args->input_value);
size_t input_size = RSTRING_LEN(args->input_value);

size_t max_compressed_size = ZSTD_compressBound(input_size);
VALUE output = rb_str_new(NULL, max_compressed_size);
char* output_data = RSTRING_PTR(output);

size_t const ret = zstd_compress(ctx, output_data, max_compressed_size, input_data, input_size, false);
ZSTD_freeCCtx(ctx);
size_t const ret = zstd_compress(args->ctx, output_data, max_compressed_size, input_data, input_size, false);
if (ZSTD_isError(ret)) {
rb_raise(rb_eRuntimeError, "compress error error code: %s", ZSTD_getErrorName(ret));
}
Expand All @@ -40,6 +35,33 @@ static VALUE rb_compress(int argc, VALUE *argv, VALUE self)
return output;
}

static VALUE compress_ensure(VALUE arg)
{
ZSTD_freeCCtx(((struct compress_args*)arg)->ctx);
return Qnil;
}

static VALUE rb_compress(int argc, VALUE *argv, VALUE self)
{
VALUE input_value;
VALUE kwargs;
rb_scan_args(argc, argv, "10:", &input_value, &kwargs);

StringValue(input_value);

ZSTD_CCtx* const ctx = ZSTD_createCCtx();
if (ctx == NULL) {
rb_raise(rb_eRuntimeError, "%s", "ZSTD_createCCtx error");
}

/* The body can raise -- a bad keyword, or an interrupt delivered when
zstd_compress reacquires the GVL -- so the context is freed under ensure. */
struct compress_args args = { ctx, input_value, kwargs };
VALUE output = rb_ensure(compress_body, (VALUE)&args, compress_ensure, (VALUE)&args);
RB_GC_GUARD(input_value);
return output;
}

static VALUE decode_one_frame(ZSTD_DCtx* dctx, const unsigned char* src, size_t size, VALUE kwargs, size_t* consumed) {
VALUE out = rb_str_buf_new(0);
size_t cap = ZSTD_DStreamOutSize();
Expand Down Expand Up @@ -196,7 +218,7 @@ static VALUE rb_cdict_initialize(int argc, VALUE *argv, VALUE self)
VALUE dict;
VALUE compression_level_value;
rb_scan_args(argc, argv, "11", &dict, &compression_level_value);
int compression_level = convert_compression_level(NULL, compression_level_value);
int compression_level = convert_compression_level(compression_level_value);

StringValue(dict);
char* dict_buffer = RSTRING_PTR(dict);
Expand Down
12 changes: 12 additions & 0 deletions spec/zstd-ruby-streaming-compress_spec.rb
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,18 @@
end
end

# initialize raises after creating the ZSTD_CCtx; the object's free callback
# has to own it by then. `rake spec:valgrind` catches a leak or a double free.
describe 'initialize raising after creating the compression context' do
it 'raises ArgumentError for an unknown keyword' do
expect { Zstd::StreamingCompress.new(unknown: 1) }.to raise_error(ArgumentError)
end

it 'raises ArgumentError for a dict: that is neither a CDict nor a String' do
expect { Zstd::StreamingCompress.new(dict: 123) }.to raise_error(ArgumentError)
end
end

if Gem::Version.new(RUBY_VERSION) >= Gem::Version.new('3.0.0')
describe 'Ractor' do
it 'should be supported' do
Expand Down
27 changes: 27 additions & 0 deletions spec/zstd-ruby_spec.rb
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,33 @@ def to_str
decompressed = Zstd.decompress(compressed)
expect(decompressed).to eq('abc')
end

# Each of these raises after the ZSTD_CCtx has been created. The assertions
# pin the behaviour; `rake spec:valgrind` is what catches a leaked context.
context 'when it raises after creating the compression context' do
it 'raises ArgumentError for an unknown keyword' do
expect { Zstd.compress('abc', unknown: 1) }.to raise_error(ArgumentError)
end

it 'raises RangeError for a level that does not fit in an int' do
expect { Zstd.compress('abc', level: 2**40) }.to raise_error(RangeError)
end

it 'raises ArgumentError for a dict: that is neither a CDict nor a String' do
expect { Zstd.compress('abc', dict: 123) }.to raise_error(ArgumentError)
end

it 'can be interrupted by Thread#raise while compressing' do
interrupt = Class.new(StandardError)
input = Random.new(42).bytes(16 * 1024 * 1024)
thread = Thread.new { Zstd.compress(input, level: 19) }
thread.report_on_exception = false
sleep 0.1
thread.raise(interrupt)

expect { thread.join }.to raise_error(interrupt)
end
end
end

describe 'decompress' do
Expand Down
Loading