Merge pull request #5405 from daniel-j-born/custom_test_creds

Allow new credential types to be added to tests.
pull/5483/head
Dan Born 9 years ago
commit ef7e42660c
  1. 78
      test/cpp/util/test_credentials_provider.cc
  2. 19
      test/cpp/util/test_credentials_provider.h

@ -34,6 +34,8 @@
#include "test/cpp/util/test_credentials_provider.h" #include "test/cpp/util/test_credentials_provider.h"
#include <unordered_map>
#include <grpc/support/sync.h> #include <grpc/support/sync.h>
#include <grpc++/impl/sync.h> #include <grpc++/impl/sync.h>
@ -48,12 +50,36 @@ using grpc::InsecureServerCredentials;
using grpc::ServerCredentials; using grpc::ServerCredentials;
using grpc::SslCredentialsOptions; using grpc::SslCredentialsOptions;
using grpc::SslServerCredentialsOptions; using grpc::SslServerCredentialsOptions;
using grpc::testing::CredentialsProvider; using grpc::testing::CredentialTypeProvider;
// Provide test credentials. Thread-safe.
class CredentialsProvider {
public:
virtual ~CredentialsProvider() {}
virtual void AddSecureType(
const grpc::string& type,
std::unique_ptr<CredentialTypeProvider> type_provider) = 0;
virtual std::shared_ptr<ChannelCredentials> GetChannelCredentials(
const grpc::string& type, ChannelArguments* args) = 0;
virtual std::shared_ptr<ServerCredentials> GetServerCredentials(
const grpc::string& type) = 0;
virtual std::vector<grpc::string> GetSecureCredentialsTypeList() = 0;
};
class DefaultCredentialsProvider : public CredentialsProvider { class DefaultCredentialsProvider : public CredentialsProvider {
public: public:
~DefaultCredentialsProvider() override {} ~DefaultCredentialsProvider() override {}
void AddSecureType(
const grpc::string& type,
std::unique_ptr<CredentialTypeProvider> type_provider) override {
// This clobbers any existing entry for type, except the defaults, which
// can't be clobbered.
grpc::unique_lock<grpc::mutex> lock(mu_);
added_secure_types_[type] = std::move(type_provider);
}
std::shared_ptr<ChannelCredentials> GetChannelCredentials( std::shared_ptr<ChannelCredentials> GetChannelCredentials(
const grpc::string& type, ChannelArguments* args) override { const grpc::string& type, ChannelArguments* args) override {
if (type == grpc::testing::kInsecureCredentialsType) { if (type == grpc::testing::kInsecureCredentialsType) {
@ -63,10 +89,15 @@ class DefaultCredentialsProvider : public CredentialsProvider {
args->SetSslTargetNameOverride("foo.test.google.fr"); args->SetSslTargetNameOverride("foo.test.google.fr");
return SslCredentials(ssl_opts); return SslCredentials(ssl_opts);
} else { } else {
grpc::unique_lock<grpc::mutex> lock(mu_);
auto it(added_secure_types_.find(type));
if (it == added_secure_types_.end()) {
gpr_log(GPR_ERROR, "Unsupported credentials type %s.", type.c_str()); gpr_log(GPR_ERROR, "Unsupported credentials type %s.", type.c_str());
}
return nullptr; return nullptr;
} }
return it->second->GetChannelCredentials(args);
}
}
std::shared_ptr<ServerCredentials> GetServerCredentials( std::shared_ptr<ServerCredentials> GetServerCredentials(
const grpc::string& type) override { const grpc::string& type) override {
@ -80,33 +111,40 @@ class DefaultCredentialsProvider : public CredentialsProvider {
ssl_opts.pem_key_cert_pairs.push_back(pkcp); ssl_opts.pem_key_cert_pairs.push_back(pkcp);
return SslServerCredentials(ssl_opts); return SslServerCredentials(ssl_opts);
} else { } else {
grpc::unique_lock<grpc::mutex> lock(mu_);
auto it(added_secure_types_.find(type));
if (it == added_secure_types_.end()) {
gpr_log(GPR_ERROR, "Unsupported credentials type %s.", type.c_str()); gpr_log(GPR_ERROR, "Unsupported credentials type %s.", type.c_str());
}
return nullptr; return nullptr;
} }
return it->second->GetServerCredentials();
}
}
std::vector<grpc::string> GetSecureCredentialsTypeList() override { std::vector<grpc::string> GetSecureCredentialsTypeList() override {
std::vector<grpc::string> types; std::vector<grpc::string> types;
types.push_back(grpc::testing::kTlsCredentialsType); types.push_back(grpc::testing::kTlsCredentialsType);
grpc::unique_lock<grpc::mutex> lock(mu_);
for (const auto& type_pair : added_secure_types_) {
types.push_back(type_pair.first);
}
return types; return types;
} }
private:
grpc::mutex mu_;
std::unordered_map<grpc::string, std::unique_ptr<CredentialTypeProvider> >
added_secure_types_;
}; };
gpr_once g_once_init_provider_mu = GPR_ONCE_INIT; gpr_once g_once_init_provider = GPR_ONCE_INIT;
grpc::mutex* g_provider_mu = nullptr;
CredentialsProvider* g_provider = nullptr; CredentialsProvider* g_provider = nullptr;
void InitProviderMu() { g_provider_mu = new grpc::mutex; } void CreateDefaultProvider() {
g_provider = new DefaultCredentialsProvider;
grpc::mutex& GetMu() {
gpr_once_init(&g_once_init_provider_mu, &InitProviderMu);
return *g_provider_mu;
} }
CredentialsProvider* GetProvider() { CredentialsProvider* GetProvider() {
grpc::unique_lock<grpc::mutex> lock(GetMu()); gpr_once_init(&g_once_init_provider, &CreateDefaultProvider);
if (g_provider == nullptr) {
g_provider = new DefaultCredentialsProvider;
}
return g_provider; return g_provider;
} }
@ -115,15 +153,9 @@ CredentialsProvider* GetProvider() {
namespace grpc { namespace grpc {
namespace testing { namespace testing {
// Note that it is not thread-safe to set a provider while concurrently using void AddSecureType(const grpc::string& type,
// the previously set provider, as this deletes and replaces it. nullptr may be std::unique_ptr<CredentialTypeProvider> type_provider) {
// given to reset to the default. GetProvider()->AddSecureType(type, std::move(type_provider));
void SetTestCredentialsProvider(std::unique_ptr<CredentialsProvider> provider) {
grpc::unique_lock<grpc::mutex> lock(GetMu());
if (g_provider != nullptr) {
delete g_provider;
}
g_provider = provider.release();
} }
std::shared_ptr<ChannelCredentials> GetChannelCredentials( std::shared_ptr<ChannelCredentials> GetChannelCredentials(

@ -46,20 +46,21 @@ namespace testing {
const char kInsecureCredentialsType[] = "INSECURE_CREDENTIALS"; const char kInsecureCredentialsType[] = "INSECURE_CREDENTIALS";
const char kTlsCredentialsType[] = "TLS_CREDENTIALS"; const char kTlsCredentialsType[] = "TLS_CREDENTIALS";
class CredentialsProvider { // Provide test credentials of a particular type.
class CredentialTypeProvider {
public: public:
virtual ~CredentialsProvider() {} virtual ~CredentialTypeProvider() {}
virtual std::shared_ptr<ChannelCredentials> GetChannelCredentials( virtual std::shared_ptr<ChannelCredentials> GetChannelCredentials(
const grpc::string& type, ChannelArguments* args) = 0; ChannelArguments* args) = 0;
virtual std::shared_ptr<ServerCredentials> GetServerCredentials( virtual std::shared_ptr<ServerCredentials> GetServerCredentials() = 0;
const grpc::string& type) = 0;
virtual std::vector<grpc::string> GetSecureCredentialsTypeList() = 0;
}; };
// Set the CredentialsProvider used by the other functions in this file. If this // Add a secure type in addition to the defaults above
// is not set, a default provider will be used. // (kInsecureCredentialsType, kTlsCredentialsType) that can be returned from the
void SetTestCredentialsProvider(std::unique_ptr<CredentialsProvider> provider); // functions below.
void AddSecureType(const grpc::string& type,
std::unique_ptr<CredentialTypeProvider> type_provider);
// Provide channel credentials according to the given type. Alter the channel // Provide channel credentials according to the given type. Alter the channel
// arguments if needed. // arguments if needed.

Loading…
Cancel
Save