You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 

145 lines
4.2 KiB

package generic
import (
"context"
"fmt"
"github.com/jhump/protoreflect/desc"
"github.com/jhump/protoreflect/dynamic"
"github.com/jhump/protoreflect/dynamic/grpcdynamic"
"github.com/jhump/protoreflect/grpcreflect"
"google.golang.org/grpc"
"google.golang.org/grpc/reflection/grpc_reflection_v1alpha"
"sonet/pkg/grpc/generic/desc_source"
"sonet/pkg/utils/logger"
"sync"
)
type GrpcGenericClient struct {
serviceName string
conn *grpc.ClientConn
descSource desc_source.DescriptorSource
serviceDesc *desc.ServiceDescriptor
callerCache *sync.Map
}
func NewGpcGenericClient(serviceName string, conn *grpc.ClientConn) *GrpcGenericClient {
return &GrpcGenericClient{
serviceName: serviceName,
conn: conn,
}
}
func (c *GrpcGenericClient) Init(ctx context.Context) (err error) {
// fetch service description
refClient := grpcreflect.NewClientV1Alpha(ctx, grpc_reflection_v1alpha.NewServerReflectionClient(c.conn))
desc_source.DescriptorSourceFromServer(ctx, refClient)
dsc, e1 := c.descSource.FindSymbol(c.serviceName)
if e1 != nil {
err = fmt.Errorf("service %s not found", c.serviceName)
if desc_source.IsNotFoundError(e1) {
logger.Errorf("target server not expose service %q in FindSymbol", c.serviceName)
return
}
logger.Errorf("failed to query for service descriptor %q: %v", c.serviceName, e1)
return
}
sd, ok := dsc.(*desc.ServiceDescriptor)
if !ok {
err = fmt.Errorf("service %s not found", c.serviceName)
logger.Errorf("target server not expose service %q", c.serviceName)
return
}
c.serviceDesc = sd
c.callerCache = &sync.Map{}
return
}
func (c *GrpcGenericClient) ServiceName() string {
return c.serviceName
}
type methodCaller struct {
Mtd *desc.MethodDescriptor
MsgFactory *dynamic.MessageFactory
Stub grpcdynamic.Stub
}
func (c *GrpcGenericClient) InvokeUnary(ctx context.Context, method string, reqBytes []byte, opts ...grpc.CallOption) (resp *dynamic.Message, err error) {
// cache method desc
caller, err := c.getMethodCaller(method)
if err != nil {
return
}
reqMessage := caller.MsgFactory.NewMessage(caller.Mtd.GetInputType())
if err = reqMessage.(*dynamic.Message).Unmarshal(reqBytes); err != nil {
err = fmt.Errorf("unmarshal req bytes error: %s", err.Error())
return
}
res, err := caller.Stub.InvokeRpc(ctx, caller.Mtd, reqMessage, opts...)
if err != nil {
return
}
resp = res.(*dynamic.Message)
return
}
// getMethodCaller load generic resource from cache
func (c *GrpcGenericClient) getMethodCaller(method string) (caller *methodCaller, err error) {
val, ok := c.callerCache.Load(method)
if ok {
caller = val.(*methodCaller)
} else {
// method desc
caller = &methodCaller{}
caller.Mtd = c.serviceDesc.FindMethodByName(method)
if caller.Mtd == nil {
logger.Errorf("service %q does not include a method named %q", c.serviceName, method)
err = fmt.Errorf("method %s not found", method)
return
}
// message factory
var ext dynamic.ExtensionRegistry
if err = c.fetchAllExtensions(&ext, caller.Mtd.GetInputType()); err != nil {
return
}
if err = c.fetchAllExtensions(&ext, caller.Mtd.GetOutputType()); err != nil {
return
}
caller.MsgFactory = dynamic.NewMessageFactoryWithExtensionRegistry(&ext)
// stub
caller.Stub = grpcdynamic.NewStubWithMessageFactory(c.conn, caller.MsgFactory)
c.callerCache.Store(method, caller)
}
return
}
func (c *GrpcGenericClient) fetchAllExtensions(ext *dynamic.ExtensionRegistry, md *desc.MessageDescriptor) (err error) {
msgTypeName := md.GetFullyQualifiedName()
if len(md.GetExtensionRanges()) > 0 {
fds, err := c.descSource.AllExtensionsForType(msgTypeName)
if err != nil {
return fmt.Errorf("failed to query for extensions of type %s: %v", msgTypeName, err)
}
for _, fd := range fds {
if err := ext.AddExtension(fd); err != nil {
return fmt.Errorf("could not register extension %s of type %s: %v", fd.GetFullyQualifiedName(), msgTypeName, err)
}
}
}
// recursively fetch extensions for the types of any message fields
for _, fd := range md.GetFields() {
if fd.GetMessageType() != nil {
err := c.fetchAllExtensions(ext, fd.GetMessageType())
if err != nil {
return err
}
}
}
return nil
}