diff --git a/clickhouse/client.cpp b/clickhouse/client.cpp index dae5cb7f..78e5d79b 100644 --- a/clickhouse/client.cpp +++ b/clickhouse/client.cpp @@ -1074,6 +1074,14 @@ void Client::Impl::SendQuery(const Query& query, bool finalize) { /// Per query settings if (server_info_.revision >= DBMS_MIN_REVISION_WITH_SETTINGS_SERIALIZED_AS_STRINGS) { + // The compression flag enables compression; this setting selects the response codec. + const auto method = options_.compression_method; + if ((method == CompressionMethod::LZ4 || method == CompressionMethod::ZSTD) && + query.GetQuerySettings().count("network_compression_method") == 0) { + WireFormat::WriteString(*output_, "network_compression_method"); + WireFormat::WriteVarint64(*output_, 0); + WireFormat::WriteString(*output_, method == CompressionMethod::ZSTD ? "ZSTD" : "LZ4"); + } for(const auto& [name, field] : query.GetQuerySettings()) { WireFormat::WriteString(*output_, name); WireFormat::WriteVarint64(*output_, field.flags); diff --git a/clickhouse/client.h b/clickhouse/client.h index dfe31536..ce029b4b 100644 --- a/clickhouse/client.h +++ b/clickhouse/client.h @@ -103,7 +103,8 @@ struct ClientOptions { /// Amount of time to wait before next retry. DECLARE_FIELD(retry_timeout, std::chrono::seconds, SetRetryTimeout, std::chrono::seconds(5)); - /// Compression method. + /// Compression method for outgoing and incoming blocks. + /// An explicit network_compression_method query setting overrides the incoming codec. DECLARE_FIELD(compression_method, CompressionMethod, SetCompressionMethod, CompressionMethod::None); /// TCP Keep alive options diff --git a/ut/BUILD.bazel b/ut/BUILD.bazel index a41e3e54..bd7541a4 100644 --- a/ut/BUILD.bazel +++ b/ut/BUILD.bazel @@ -42,6 +42,7 @@ cc_test( "bignum_string_ut.cpp", "bignum_ut.cpp", "block_ut.cpp", + "client_protocol_ut.cpp", "column_array_ut.cpp", "columns_ut.cpp", "itemview_ut.cpp", diff --git a/ut/client_protocol_ut.cpp b/ut/client_protocol_ut.cpp index 72cfdd5d..b0f978e1 100644 --- a/ut/client_protocol_ut.cpp +++ b/ut/client_protocol_ut.cpp @@ -2,6 +2,7 @@ #include #include #include +#include #include #include @@ -16,11 +17,13 @@ namespace { using namespace clickhouse; /// A socket that serves a pre-recorded byte script as the server response -/// and discards (but keeps) everything written by the client. +/// and records everything written by the client. An optional capture buffer +/// must outlive the socket. class ScriptedSocket : public SocketBase { public: - explicit ScriptedSocket(std::vector script) + explicit ScriptedSocket(std::vector script, std::vector* captured = nullptr) : script_(std::move(script)) + , captured_(captured) {} std::unique_ptr makeInputStream() const override { @@ -28,26 +31,29 @@ class ScriptedSocket : public SocketBase { } std::unique_ptr makeOutputStream() const override { - return std::make_unique(&written_); + return std::make_unique(captured_ ? captured_ : &written_); } private: const std::vector script_; + std::vector* captured_; mutable std::vector written_; }; class ScriptedSocketFactory : public SocketFactory { public: - explicit ScriptedSocketFactory(std::vector script) + explicit ScriptedSocketFactory(std::vector script, std::vector* captured = nullptr) : script_(std::move(script)) + , captured_(captured) {} std::unique_ptr connect(const ClientOptions&, const Endpoint&) override { - return std::make_unique(script_); + return std::make_unique(script_, captured_); } private: std::vector script_; + std::vector* captured_; }; ClientOptions ScriptedClientOptions() { @@ -238,3 +244,37 @@ TEST(ClientProtocol, WellFormedSelectResponseSucceeds) { })); EXPECT_EQ(blocks, 1u); } + +TEST(ClientProtocol, CompressionWithOldServerRevision) { + std::vector script = kServerHello; + script.push_back(0x05); // ServerCodes::EndOfStream + + for (const auto method : {CompressionMethod::LZ4, CompressionMethod::ZSTD}) { + std::vector written; + Client client(ScriptedClientOptions().SetCompressionMethod(method), + std::make_unique(script, &written)); + const auto query_offset = written.size(); // Skip the client handshake. + // Older servers cannot receive string-serialized query settings. + ASSERT_NO_THROW(client.Execute("SELECT 1")); + + ArrayInput input(written.data() + query_offset, written.size() - query_offset); + uint64_t value = 0; + std::string text; + ASSERT_TRUE(WireFormat::ReadUInt64(input, &value)); + ASSERT_EQ(1u, value); // ClientCodes::Query + ASSERT_TRUE(WireFormat::ReadString(input, &text)); + ASSERT_TRUE(text.empty()); // Query ID + ASSERT_TRUE(WireFormat::ReadString(input, &text)); + ASSERT_TRUE(text.empty()); // Settings terminator: no automatic setting + ASSERT_TRUE(WireFormat::ReadUInt64(input, &value)); + ASSERT_EQ(2u, value); // Stages::Complete + ASSERT_TRUE(WireFormat::ReadUInt64(input, &value)); + EXPECT_EQ(1u, value); // Compression remains enabled. + ASSERT_TRUE(WireFormat::ReadString(input, &text)); + EXPECT_EQ("SELECT 1", text); + + Query query("SELECT 1"); + query.SetSetting("network_compression_method", {"LZ4"}); + EXPECT_THROW(client.Execute(query), UnimplementedError); + } +} diff --git a/ut/roundtrip_tests.cpp b/ut/roundtrip_tests.cpp index 9476d2b1..e3526dd6 100644 --- a/ut/roundtrip_tests.cpp +++ b/ut/roundtrip_tests.cpp @@ -39,6 +39,55 @@ class RoundtripCase : public testing::TestWithParam { std::unique_ptr client_; }; +TEST_P(RoundtripCase, ResponseCompressionMethod) { + for (const std::string server_method : {"LZ4", "ZSTD"}) { + client_->Execute("SET network_compression_method = '" + server_method + "'"); + + const auto method = GetParam().compression_method; + const auto expected = method == CompressionMethod::None ? server_method + : method == CompressionMethod::LZ4 ? "LZ4" : "ZSTD"; + EXPECT_EQ(expected, GetSettingValue("network_compression_method")); + } +} + +TEST_P(RoundtripCase, ResponseCompressionMethodQuerySetting) { + for (const std::string method : {"LZ4", "ZSTD"}) { + Query query("SELECT value FROM system.settings WHERE name = 'network_compression_method'"); + query.SetSetting("network_compression_method", {method}); + + std::string result; + query.OnData([&result](const Block& block) { + if (block.GetRowCount() != 0) { + result = block[0]->AsStrict()->At(0); + } + }); + client_->Execute(query); + EXPECT_EQ(method, result); + } +} + +TEST_P(RoundtripCase, ResponseCompressionReadonly) { + const auto method = GetParam().compression_method; + const std::string server_method = method == CompressionMethod::ZSTD ? "LZ4" : "ZSTD"; + client_->Execute("SET network_compression_method = '" + server_method + "'"); + client_->Execute("SET readonly = 1"); + + Query query("SELECT 1"); + query.SetSetting("network_compression_method", {server_method}); + EXPECT_NO_THROW(client_->Execute(query)); + + if (method == CompressionMethod::None) { + EXPECT_NO_THROW(client_->Execute("SELECT 1")); + } else { + try { + client_->Execute("SELECT 1"); + FAIL() << "expected readonly to prevent changing the response codec"; + } catch (const ServerException& e) { + EXPECT_EQ(164, e.GetCode()); // READONLY + } + } +} + TEST_P(RoundtripCase, ArrayTUint64) { auto array = std::make_shared>(); array->Append({0, 1, 2});