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
8 changes: 8 additions & 0 deletions clickhouse/client.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
3 changes: 2 additions & 1 deletion clickhouse/client.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions ut/BUILD.bazel
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
50 changes: 45 additions & 5 deletions ut/client_protocol_ut.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
#include <clickhouse/base/input.h>
#include <clickhouse/base/output.h>
#include <clickhouse/base/socket.h>
#include <clickhouse/base/wire_format.h>
#include <clickhouse/exceptions.h>

#include <gtest/gtest.h>
Expand All @@ -16,38 +17,43 @@ 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<uint8_t> script)
explicit ScriptedSocket(std::vector<uint8_t> script, std::vector<uint8_t>* captured = nullptr)
: script_(std::move(script))
, captured_(captured)
{}

std::unique_ptr<InputStream> makeInputStream() const override {
return std::make_unique<ArrayInput>(script_.data(), script_.size());
}

std::unique_ptr<OutputStream> makeOutputStream() const override {
return std::make_unique<BufferOutput>(&written_);
return std::make_unique<BufferOutput>(captured_ ? captured_ : &written_);
}

private:
const std::vector<uint8_t> script_;
std::vector<uint8_t>* captured_;
mutable std::vector<uint8_t> written_;
};

class ScriptedSocketFactory : public SocketFactory {
public:
explicit ScriptedSocketFactory(std::vector<uint8_t> script)
explicit ScriptedSocketFactory(std::vector<uint8_t> script, std::vector<uint8_t>* captured = nullptr)
: script_(std::move(script))
, captured_(captured)
{}

std::unique_ptr<SocketBase> connect(const ClientOptions&, const Endpoint&) override {
return std::make_unique<ScriptedSocket>(script_);
return std::make_unique<ScriptedSocket>(script_, captured_);
}

private:
std::vector<uint8_t> script_;
std::vector<uint8_t>* captured_;
};

ClientOptions ScriptedClientOptions() {
Expand Down Expand Up @@ -238,3 +244,37 @@ TEST(ClientProtocol, WellFormedSelectResponseSucceeds) {
}));
EXPECT_EQ(blocks, 1u);
}

TEST(ClientProtocol, CompressionWithOldServerRevision) {
std::vector<uint8_t> script = kServerHello;
script.push_back(0x05); // ServerCodes::EndOfStream

for (const auto method : {CompressionMethod::LZ4, CompressionMethod::ZSTD}) {
std::vector<uint8_t> written;
Client client(ScriptedClientOptions().SetCompressionMethod(method),
std::make_unique<ScriptedSocketFactory>(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);
}
}
49 changes: 49 additions & 0 deletions ut/roundtrip_tests.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,55 @@ class RoundtripCase : public testing::TestWithParam<ClientOptions> {
std::unique_ptr<Client> 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<ColumnString>()->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<ColumnArrayT<ColumnUInt64>>();
array->Append({0, 1, 2});
Expand Down
Loading