/src/duckdb/extension/parquet/zstd_file_system.cpp
Line | Count | Source |
1 | | #include "zstd_file_system.hpp" |
2 | | |
3 | | #include <exception> |
4 | | #include <utility> |
5 | | |
6 | | #include "zstd.h" |
7 | | #include "duckdb/common/assert.hpp" |
8 | | #include "duckdb/common/enums/file_compression_type.hpp" |
9 | | #include "duckdb/common/error_data.hpp" |
10 | | #include "duckdb/common/exception.hpp" |
11 | | #include "duckdb/common/helper.hpp" |
12 | | #include "duckdb/common/numeric_utils.hpp" |
13 | | #include "duckdb/common/shared_ptr_ipp.hpp" |
14 | | #include "duckdb/common/string.hpp" |
15 | | #include "duckdb/logging/logger.hpp" |
16 | | |
17 | | namespace duckdb { |
18 | | |
19 | | namespace { |
20 | | |
21 | | struct ZstdStreamWrapper : public StreamWrapper { |
22 | | ~ZstdStreamWrapper() override; |
23 | | |
24 | | CompressedFile *file = nullptr; |
25 | | duckdb_zstd::ZSTD_DStream *zstd_stream_ptr = nullptr; |
26 | | duckdb_zstd::ZSTD_CStream *zstd_compress_ptr = nullptr; |
27 | | bool writing = false; |
28 | | |
29 | | public: |
30 | | void Initialize(QueryContext context, CompressedFile &file, bool write) override; |
31 | | bool Read(StreamData &stream_data) override; |
32 | | void Write(CompressedFile &file, StreamData &stream_data, data_ptr_t buffer, int64_t nr_bytes) override; |
33 | | |
34 | | void Close() override; |
35 | | |
36 | | void FlushStream(); |
37 | | }; |
38 | | |
39 | 0 | ZstdStreamWrapper::~ZstdStreamWrapper() { |
40 | 0 | if (Exception::UncaughtException()) { |
41 | 0 | return; |
42 | 0 | } |
43 | 0 | try { |
44 | 0 | Close(); |
45 | 0 | } catch (std::exception &ex) { |
46 | 0 | if (file && file->child_handle) { |
47 | | // FIXME: Make any log context available here. |
48 | 0 | ErrorData data(ex); |
49 | 0 | try { |
50 | 0 | const auto logger = file->child_handle->logger; |
51 | 0 | if (logger) { |
52 | 0 | DUCKDB_LOG_ERROR(logger, "ZstdStreamWrapper::~ZstdStreamWrapper()\t\t" + data.Message()); |
53 | 0 | } |
54 | 0 | } catch (...) { // NOLINT |
55 | 0 | } |
56 | 0 | } |
57 | 0 | } catch (...) { // NOLINT |
58 | 0 | } |
59 | 0 | } |
60 | | |
61 | 0 | void ZstdStreamWrapper::Initialize(QueryContext context, CompressedFile &file, bool write) { |
62 | 0 | D_ASSERT(!zstd_stream_ptr && !zstd_compress_ptr); |
63 | 0 | this->file = &file; |
64 | 0 | this->writing = write; |
65 | 0 | if (write) { |
66 | 0 | zstd_compress_ptr = duckdb_zstd::ZSTD_createCStream(); |
67 | 0 | } else { |
68 | 0 | zstd_stream_ptr = duckdb_zstd::ZSTD_createDStream(); |
69 | 0 | } |
70 | 0 | } |
71 | | |
72 | 0 | bool ZstdStreamWrapper::Read(StreamData &sd) { |
73 | 0 | D_ASSERT(!writing); |
74 | |
|
75 | 0 | duckdb_zstd::ZSTD_inBuffer in_buffer; |
76 | 0 | duckdb_zstd::ZSTD_outBuffer out_buffer; |
77 | |
|
78 | 0 | in_buffer.src = sd.in_buff_start; |
79 | 0 | in_buffer.size = sd.in_buff_end - sd.in_buff_start; |
80 | 0 | in_buffer.pos = 0; |
81 | |
|
82 | 0 | out_buffer.dst = sd.out_buff_start; |
83 | 0 | out_buffer.size = sd.out_buf_size; |
84 | 0 | out_buffer.pos = 0; |
85 | |
|
86 | 0 | auto res = duckdb_zstd::ZSTD_decompressStream(zstd_stream_ptr, &out_buffer, &in_buffer); |
87 | 0 | if (duckdb_zstd::ZSTD_isError(res)) { |
88 | 0 | throw IOException(duckdb_zstd::ZSTD_getErrorName(res)); |
89 | 0 | } |
90 | | |
91 | 0 | sd.in_buff_start = (data_ptr_t)in_buffer.src + in_buffer.pos; // NOLINT |
92 | 0 | sd.in_buff_end = (data_ptr_t)in_buffer.src + in_buffer.size; // NOLINT |
93 | 0 | sd.out_buff_end = (data_ptr_t)out_buffer.dst + out_buffer.pos; // NOLINT |
94 | 0 | return false; |
95 | 0 | } |
96 | | |
97 | | void ZstdStreamWrapper::Write(CompressedFile &file, StreamData &sd, data_ptr_t uncompressed_data, |
98 | 0 | int64_t uncompressed_size) { |
99 | 0 | D_ASSERT(writing); |
100 | |
|
101 | 0 | auto remaining = uncompressed_size; |
102 | 0 | while (remaining > 0) { |
103 | 0 | D_ASSERT(sd.out_buff.get() + sd.out_buf_size > sd.out_buff_start); |
104 | 0 | idx_t output_remaining = (sd.out_buff.get() + sd.out_buf_size) - sd.out_buff_start; |
105 | |
|
106 | 0 | duckdb_zstd::ZSTD_inBuffer in_buffer; |
107 | 0 | duckdb_zstd::ZSTD_outBuffer out_buffer; |
108 | |
|
109 | 0 | in_buffer.src = uncompressed_data; |
110 | 0 | in_buffer.size = remaining; |
111 | 0 | in_buffer.pos = 0; |
112 | |
|
113 | 0 | out_buffer.dst = sd.out_buff_start; |
114 | 0 | out_buffer.size = output_remaining; |
115 | 0 | out_buffer.pos = 0; |
116 | 0 | auto res = |
117 | 0 | duckdb_zstd::ZSTD_compressStream2(zstd_compress_ptr, &out_buffer, &in_buffer, duckdb_zstd::ZSTD_e_continue); |
118 | 0 | if (duckdb_zstd::ZSTD_isError(res)) { |
119 | 0 | throw IOException(duckdb_zstd::ZSTD_getErrorName(res)); |
120 | 0 | } |
121 | 0 | idx_t input_consumed = in_buffer.pos; |
122 | 0 | idx_t written_to_output = out_buffer.pos; |
123 | 0 | sd.out_buff_start += written_to_output; |
124 | 0 | if (sd.out_buff_start == sd.out_buff.get() + sd.out_buf_size) { |
125 | | // no more output buffer available: flush |
126 | 0 | file.child_handle->Write(sd.out_buff.get(), sd.out_buff_start - sd.out_buff.get()); |
127 | 0 | sd.out_buff_start = sd.out_buff.get(); |
128 | 0 | } |
129 | 0 | uncompressed_data += input_consumed; |
130 | 0 | remaining -= UnsafeNumericCast<int64_t>(input_consumed); |
131 | 0 | } |
132 | 0 | } |
133 | | |
134 | 0 | void ZstdStreamWrapper::FlushStream() { |
135 | 0 | auto &sd = file->stream_data; |
136 | 0 | duckdb_zstd::ZSTD_inBuffer in_buffer; |
137 | 0 | duckdb_zstd::ZSTD_outBuffer out_buffer; |
138 | |
|
139 | 0 | in_buffer.src = nullptr; |
140 | 0 | in_buffer.size = 0; |
141 | 0 | in_buffer.pos = 0; |
142 | 0 | while (true) { |
143 | 0 | idx_t output_remaining = (sd.out_buff.get() + sd.out_buf_size) - sd.out_buff_start; |
144 | |
|
145 | 0 | out_buffer.dst = sd.out_buff_start; |
146 | 0 | out_buffer.size = output_remaining; |
147 | 0 | out_buffer.pos = 0; |
148 | |
|
149 | 0 | auto res = |
150 | 0 | duckdb_zstd::ZSTD_compressStream2(zstd_compress_ptr, &out_buffer, &in_buffer, duckdb_zstd::ZSTD_e_end); |
151 | 0 | if (duckdb_zstd::ZSTD_isError(res)) { |
152 | 0 | throw IOException(duckdb_zstd::ZSTD_getErrorName(res)); |
153 | 0 | } |
154 | 0 | idx_t written_to_output = out_buffer.pos; |
155 | 0 | sd.out_buff_start += written_to_output; |
156 | 0 | if (sd.out_buff_start > sd.out_buff.get()) { |
157 | 0 | file->child_handle->Write(sd.out_buff.get(), sd.out_buff_start - sd.out_buff.get()); |
158 | 0 | sd.out_buff_start = sd.out_buff.get(); |
159 | 0 | } |
160 | 0 | if (res == 0) { |
161 | 0 | break; |
162 | 0 | } |
163 | 0 | } |
164 | 0 | } |
165 | | |
166 | 0 | void ZstdStreamWrapper::Close() { |
167 | 0 | if (!zstd_stream_ptr && !zstd_compress_ptr) { |
168 | 0 | return; |
169 | 0 | } |
170 | 0 | if (writing) { |
171 | 0 | FlushStream(); |
172 | 0 | } |
173 | 0 | if (zstd_stream_ptr) { |
174 | 0 | duckdb_zstd::ZSTD_freeDStream(zstd_stream_ptr); |
175 | 0 | } |
176 | 0 | if (zstd_compress_ptr) { |
177 | 0 | duckdb_zstd::ZSTD_freeCStream(zstd_compress_ptr); |
178 | 0 | } |
179 | 0 | zstd_stream_ptr = nullptr; |
180 | 0 | zstd_compress_ptr = nullptr; |
181 | 0 | } |
182 | | |
183 | | struct ZStdFileSystemHolder { |
184 | | ZStdFileSystem zstd_fs; |
185 | | }; |
186 | | |
187 | | class ZStdFile : private ZStdFileSystemHolder, public CompressedFile { |
188 | | public: |
189 | | ZStdFile(QueryContext context, unique_ptr<FileHandle> child_handle_p, const string &path, bool write) |
190 | 0 | : CompressedFile(zstd_fs, std::move(child_handle_p), path) { |
191 | 0 | Initialize(context, write); |
192 | 0 | } |
193 | | |
194 | 0 | FileCompressionType GetFileCompressionType() override { |
195 | 0 | return FileCompressionType::ZSTD; |
196 | 0 | } |
197 | | }; |
198 | | |
199 | | } // namespace |
200 | | |
201 | | unique_ptr<FileHandle> ZStdFileSystem::OpenCompressedFile(QueryContext context, unique_ptr<FileHandle> handle, |
202 | 0 | bool write) { |
203 | 0 | auto path = handle->path; |
204 | 0 | return make_uniq<ZStdFile>(context, std::move(handle), path, write); |
205 | 0 | } |
206 | | |
207 | 0 | unique_ptr<StreamWrapper> ZStdFileSystem::CreateStream() { |
208 | 0 | return make_uniq<ZstdStreamWrapper>(); |
209 | 0 | } |
210 | | |
211 | 0 | idx_t ZStdFileSystem::InBufferSize() { |
212 | 0 | return duckdb_zstd::ZSTD_DStreamInSize(); |
213 | 0 | } |
214 | | |
215 | 0 | idx_t ZStdFileSystem::OutBufferSize() { |
216 | 0 | return duckdb_zstd::ZSTD_DStreamOutSize(); |
217 | 0 | } |
218 | | |
219 | 0 | int64_t ZStdFileSystem::DefaultCompressionLevel() { |
220 | 0 | return duckdb_zstd::ZSTD_defaultCLevel(); |
221 | 0 | } |
222 | | |
223 | 0 | int64_t ZStdFileSystem::MinimumCompressionLevel() { |
224 | 0 | return duckdb_zstd::ZSTD_minCLevel(); |
225 | 0 | } |
226 | | |
227 | 0 | int64_t ZStdFileSystem::MaximumCompressionLevel() { |
228 | 0 | return duckdb_zstd::ZSTD_maxCLevel(); |
229 | 0 | } |
230 | | |
231 | | } // namespace duckdb |