Override domain in KT gRPC errors with chat domain

This commit is contained in:
moiseev-signal authored and GitHub committed 2026-09-16 09:56:41 -04:00
1 parent f55c3c22fe
commit 02b86d67ad
4 files changed
+96 -28

No files matched your search

@@ -0,0 +1,26 @@
/*
* Copyright 2026 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.whispersystems.textsecuregcm.grpc;
import com.google.protobuf.InvalidProtocolBufferException;
import com.google.rpc.ErrorInfo;
import java.io.UncheckedIOException;
import java.util.Optional;
public class ErrorUtil {
public static Optional<ErrorInfo> errorInfo(final com.google.rpc.Status statusProto) {
return statusProto.getDetailsList().stream()
.filter(any -> any.is(ErrorInfo.class))
.map(errorInfo -> {
try {
return errorInfo.unpack(ErrorInfo.class);
} catch (final InvalidProtocolBufferException e) {
throw new UncheckedIOException(e);
}
})
.findFirst();
}
}
@@ -6,22 +6,19 @@
package org.whispersystems.textsecuregcm.grpc;
import com.google.common.annotations.VisibleForTesting;
import io.grpc.Status;
import org.signal.keytransparency.client.AciMonitorRequest;
import org.signal.keytransparency.client.ConsistencyParameters;
import com.google.protobuf.Any;
import com.google.rpc.ErrorInfo;
import com.google.rpc.Status;
import io.grpc.StatusRuntimeException;
import io.grpc.protobuf.StatusProto;
import org.signal.keytransparency.client.DistinguishedRequest;
import org.signal.keytransparency.client.DistinguishedResponse;
import org.signal.keytransparency.client.E164MonitorRequest;
import org.signal.keytransparency.client.E164SearchRequest;
import org.signal.keytransparency.client.MonitorRequest;
import org.signal.keytransparency.client.MonitorResponse;
import org.signal.keytransparency.client.MonitorResponseV2;
import org.signal.keytransparency.client.SearchRequest;
import org.signal.keytransparency.client.SearchResponse;
import org.signal.keytransparency.client.SearchResponseV2;
import org.signal.keytransparency.client.SimpleKeyTransparencyQueryServiceGrpc;
import org.signal.keytransparency.client.UsernameHashMonitorRequest;
import org.whispersystems.textsecuregcm.controllers.AccountController;
import org.whispersystems.textsecuregcm.controllers.RateLimitExceededException;
import org.whispersystems.textsecuregcm.identity.AciServiceIdentifier;
import org.whispersystems.textsecuregcm.keytransparency.KeyTransparencyServiceClient;
@@ -88,4 +85,28 @@ public class KeyTransparencyGrpcService extends
return request;
}
@Override
public Throwable mapException(final Throwable throwable) {
// Reconstruct the exception so that we can override the backing service domain with chat's own domain.
if (throwable instanceof StatusRuntimeException s) {
final ErrorInfo.Builder errInfoBuilder = ErrorInfo.newBuilder()
.setDomain(GrpcExceptions.DOMAIN);
final Status statusProto = StatusProto.fromStatusAndTrailers(s.getStatus(), s.getTrailers());
ErrorUtil.errorInfo(statusProto)
.map(ErrorInfo::getReason)
.ifPresent(errInfoBuilder::setReason);
return StatusProto.toStatusRuntimeException(Status.newBuilder()
// See https://github.com/signalapp/key-transparency-server/blob/30b7a2a40604fbe04545d47f26e4c5a35f5d3161/cmd/kt-server/errors.go#L32
// for a list of error codes that the key transparency service can generate.
.setCode(statusProto.getCode())
.setMessage(statusProto.getMessage())
.addDetails(Any.pack(errInfoBuilder.build()))
.build());
}
return throwable;
}
}
@@ -127,7 +127,7 @@ public class MetricServerInterceptor implements ServerInterceptor {
@Override
public void close(final Status status, final Metadata responseHeaders) {
if (!status.isOk()) {
reason = errorInfo(StatusProto.fromStatusAndTrailers(status, responseHeaders))
reason = ErrorUtil.errorInfo(StatusProto.fromStatusAndTrailers(status, responseHeaders))
.map(ErrorInfo::getReason)
.orElse(DEFAULT_ERROR_REASON);
}
@@ -232,16 +232,4 @@ public class MetricServerInterceptor implements ServerInterceptor {
}
}
private static Optional<ErrorInfo> errorInfo(final com.google.rpc.Status statusProto) {
return statusProto.getDetailsList().stream()
.filter(any -> any.is(ErrorInfo.class))
.map(errorInfo -> {
try {
return errorInfo.unpack(ErrorInfo.class);
} catch (final InvalidProtocolBufferException e) {
throw new UncheckedIOException(e);
}
})
.findFirst();
}
}
@@ -5,10 +5,14 @@
package org.whispersystems.textsecuregcm.grpc;
import com.google.protobuf.Any;
import com.google.protobuf.ByteString;
import com.google.rpc.ErrorInfo;
import io.grpc.Channel;
import io.grpc.ServerInterceptor;
import io.grpc.Status;
import io.grpc.StatusRuntimeException;
import io.grpc.protobuf.StatusProto;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
@@ -32,7 +36,6 @@ import org.whispersystems.textsecuregcm.controllers.RateLimitExceededException;
import org.whispersystems.textsecuregcm.keytransparency.KeyTransparencyServiceClient;
import org.whispersystems.textsecuregcm.limits.RateLimiter;
import org.whispersystems.textsecuregcm.limits.RateLimiters;
import reactor.core.publisher.Mono;
import java.time.Duration;
import java.util.List;
@@ -40,6 +43,8 @@ import java.util.Optional;
import java.util.stream.Stream;
import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.Mockito.mock;
@@ -54,6 +59,7 @@ import static org.whispersystems.textsecuregcm.controllers.KeyTransparencyContro
import static org.whispersystems.textsecuregcm.controllers.KeyTransparencyControllerTest.USERNAME_HASH;
import static org.whispersystems.textsecuregcm.grpc.GrpcTestUtils.assertRateLimitExceeded;
import static org.whispersystems.textsecuregcm.grpc.GrpcTestUtils.assertStatusException;
import static org.whispersystems.textsecuregcm.grpc.GrpcTestUtils.extractErrorInfo;
import static org.whispersystems.textsecuregcm.grpc.KeyTransparencyGrpcService.COMMITMENT_INDEX_LENGTH;
@SuppressWarnings({"OptionalUsedAsFieldOrParameterType", "ResultOfMethodCallIgnored"})
@@ -63,6 +69,9 @@ public class KeyTransparencyGrpcServiceTest extends SimpleBaseGrpcTest<KeyTransp
@Mock
private RateLimiter rateLimiter;
// The key transparency service uses a different domain from the chat server.
private static final String KEY_TRANSPARENCY_DOMAIN = "key-transparency";
@Override
protected KeyTransparencyGrpcService createServiceBeforeEachTest() {
final RateLimiters rateLimiters = mock(RateLimiters.class);
@@ -79,9 +88,8 @@ public class KeyTransparencyGrpcServiceTest extends SimpleBaseGrpcTest<KeyTransp
}
@Test
void searchSuccess() throws RateLimitExceededException {
void searchSuccess() {
when(keyTransparencyServiceClient.search(any())).thenReturn(SearchResponseV2.getDefaultInstance());
Mockito.doNothing().when(rateLimiter).validate(any(String.class));
final SearchRequest request = SearchRequest.newBuilder()
.setAci(ByteString.copyFrom(ACI.toCompactByteArray()))
.setAciIdentityKey(ByteString.copyFrom(ACI_IDENTITY_KEY.serialize()))
@@ -156,11 +164,38 @@ public class KeyTransparencyGrpcServiceTest extends SimpleBaseGrpcTest<KeyTransp
verifyNoInteractions(keyTransparencyServiceClient);
}
@Test
void overrideBackingServiceErrorDomain() {
final StatusRuntimeException keyTransparencyException = StatusProto.toStatusRuntimeException(
com.google.rpc.Status.newBuilder()
.setCode(Status.Code.UNAVAILABLE.value())
.addDetails(Any.pack(ErrorInfo.newBuilder()
.setDomain(KEY_TRANSPARENCY_DOMAIN)
.setReason("UNAVAILABLE")
.build()))
.build());
when(keyTransparencyServiceClient.search(any())).thenThrow(keyTransparencyException);
final SearchRequest request = SearchRequest.newBuilder()
.setAci(ByteString.copyFrom(ACI.toCompactByteArray()))
.setAciIdentityKey(ByteString.copyFrom(ACI_IDENTITY_KEY.serialize()))
.setConsistency(ConsistencyParameters.newBuilder()
.setDistinguished(10)
.build())
.build();
final StatusRuntimeException exception =
assertThrows(StatusRuntimeException.class, () -> unauthenticatedServiceStub().searchV2(request));
assertEquals(GrpcExceptions.DOMAIN, extractErrorInfo(exception).getDomain(),
"errors from the key transparency service must report the chat server's error domain");
}
@Test
void monitorSuccess() {
when(keyTransparencyServiceClient.monitor(any())).thenReturn(MonitorResponseV2.getDefaultInstance());
when(rateLimiter.validateReactive(any(String.class)))
.thenReturn(Mono.empty());
final AciMonitorRequest aciMonitorRequest = AciMonitorRequest.newBuilder()
.setAci(ByteString.copyFrom(ACI.toCompactByteArray()))
.setCommitmentIndex(ByteString.copyFrom(new byte[COMMITMENT_INDEX_LENGTH]))
@@ -250,8 +285,6 @@ public class KeyTransparencyGrpcServiceTest extends SimpleBaseGrpcTest<KeyTransp
@Test
void distinguishedSuccess() {
when(keyTransparencyServiceClient.distinguished(any())).thenReturn(DistinguishedResponse.getDefaultInstance());
when(rateLimiter.validateReactive(any(String.class)))
.thenReturn(Mono.empty());
final DistinguishedRequest request = DistinguishedRequest.newBuilder().build();
assertDoesNotThrow(() -> unauthenticatedServiceStub().distinguishedV2(request));