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
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@

package org.apache.doris.service.arrowflight;

import org.apache.doris.common.Config;
import org.apache.doris.common.util.DebugUtil;
import org.apache.doris.common.util.Util;
import org.apache.doris.mysql.MysqlCommand;
Expand All @@ -29,6 +30,7 @@
import org.apache.doris.service.arrowflight.results.FlightSqlResultCacheEntry;
import org.apache.doris.service.arrowflight.sessions.FlightSessionsManager;
import org.apache.doris.thrift.TUniqueId;
import org.apache.doris.tls.server.TlsProtocolSet;

import com.google.common.base.Preconditions;
import com.google.common.collect.Lists;
Expand Down Expand Up @@ -120,6 +122,14 @@ public DorisFlightSqlProducer(final Location location, FlightSessionsManager fli
.withSqlQuotedIdentifierCase(SqlSupportedCaseSensitivity.SQL_CASE_SENSITIVITY_CASE_INSENSITIVE);
}

public static Location createAdvertisedLocation(String host, int port) {
if (Config.enable_tls
&& TlsProtocolSet.isProtocolIncluded(TlsProtocolSet.Protocol.ARROWFLIGHT)) {
return Location.forGrpcTls(host, port);
}
return Location.forGrpcInsecure(host, port);
}

private static ByteBuffer serializeMetadata(final Schema schema) {
final ByteArrayOutputStream outputStream = new ByteArrayOutputStream();
try {
Expand Down Expand Up @@ -262,14 +272,14 @@ private FlightInfo executeQueryStatement(String peerIdentity, ConnectContext con
// If it is different from the Doris BE node randomly routed by nginx,
// data forwarding needs to be done inside the Doris BE node.
if (endpointLoc.getResultPublicAccessAddr().isSetPort()) {
location = Location.forGrpcInsecure(endpointLoc.getResultPublicAccessAddr().hostname,
location = createAdvertisedLocation(endpointLoc.getResultPublicAccessAddr().hostname,
endpointLoc.getResultPublicAccessAddr().port);
} else {
location = Location.forGrpcInsecure(endpointLoc.getResultPublicAccessAddr().hostname,
location = createAdvertisedLocation(endpointLoc.getResultPublicAccessAddr().hostname,
endpointLoc.getResultFlightServerAddr().port);
}
} else {
location = Location.forGrpcInsecure(endpointLoc.getResultFlightServerAddr().hostname,
location = createAdvertisedLocation(endpointLoc.getResultFlightServerAddr().hostname,
endpointLoc.getResultFlightServerAddr().port);
}
// By default, the query results of all BE nodes will be aggregated to one BE node.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@

package org.apache.doris.service.arrowflight;

import org.apache.doris.common.Config;
import org.apache.doris.common.FeConstants;
import org.apache.doris.qe.ConnectContext;
import org.apache.doris.qe.StmtExecutor;
Expand Down Expand Up @@ -45,18 +46,51 @@
public class DorisFlightSqlProducerTest {

private boolean prevRunningUnitTest;
private boolean prevEnableTls;
private String prevTlsExcludedProtocols;

@Before
public void setUp() {
// FlightSqlConnectContext.init() only reaches Env when this is false; keep it true so the
// context can be built without a running FE.
prevRunningUnitTest = FeConstants.runningUnitTest;
prevEnableTls = Config.enable_tls;
prevTlsExcludedProtocols = Config.tls_excluded_protocols;
FeConstants.runningUnitTest = true;
Config.enable_tls = false;
Config.tls_excluded_protocols = "";
}

@After
public void tearDown() {
FeConstants.runningUnitTest = prevRunningUnitTest;
Config.enable_tls = prevEnableTls;
Config.tls_excluded_protocols = prevTlsExcludedProtocols;
}

@Test
public void testAdvertisedLocationUsesTlsWhenArrowFlightTlsIsEnabled() {
Config.enable_tls = true;

Assert.assertEquals("grpc+tls",
DorisFlightSqlProducer.createAdvertisedLocation("be.example.com", 8050).getUri().getScheme());
}

@Test
public void testAdvertisedLocationUsesPlaintextWhenTlsIsDisabled() {
Config.enable_tls = false;

Assert.assertEquals("grpc+tcp",
DorisFlightSqlProducer.createAdvertisedLocation("be.example.com", 8050).getUri().getScheme());
}

@Test
public void testAdvertisedLocationUsesPlaintextWhenArrowFlightTlsIsExcluded() {
Config.enable_tls = true;
Config.tls_excluded_protocols = "arrowflight";

Assert.assertEquals("grpc+tcp",
DorisFlightSqlProducer.createAdvertisedLocation("be.example.com", 8050).getUri().getScheme());
}

/**
Expand Down
Loading