From ccf667862e66fca1d9745ad4834507a86705ccde Mon Sep 17 00:00:00 2001 From: strange Date: Tue, 4 Nov 2025 18:29:59 +0800 Subject: [PATCH] grpc generic stub refresh --- cmd/exchange/main.go | 4 +- cmd/market/main.go | 4 +- cmd/sig-admin/main.go | 3 + cmd/trading/main.go | 4 +- config/exchange.toml | 4 +- pkg/grpc/discovery/consul/resolver.go | 1 + pkg/grpc/discovery/consul/watch.go | 112 +++++++++++++++++++++ pkg/grpc/discovery/consul_naming.go | 14 ++- pkg/grpc/generic/generic_client.go | 2 +- pkg/grpc/generic/generic_client_factory.go | 25 ++++- pkg/trade/risk_strategy.go | 7 +- pkg/trade/trade_strategy.go | 1 + pkg/utils/collect/sync_map.go | 8 +- 13 files changed, 178 insertions(+), 11 deletions(-) create mode 100644 pkg/grpc/discovery/consul/watch.go diff --git a/cmd/exchange/main.go b/cmd/exchange/main.go index de8a493..cadcecf 100644 --- a/cmd/exchange/main.go +++ b/cmd/exchange/main.go @@ -107,8 +107,10 @@ func main() { // consul 服务注册 register := exchangeConf.Register register.Name = pb.ExchangeService_ServiceDesc.ServiceName - if err := dis.Registry(grpcServer, register); err != nil { + if deregister, err := dis.Registry(grpcServer, register); err != nil { panic(err) + } else { + defer deregister() } // run grpc server diff --git a/cmd/market/main.go b/cmd/market/main.go index 76068bc..4af9cc7 100644 --- a/cmd/market/main.go +++ b/cmd/market/main.go @@ -68,8 +68,10 @@ func main() { register := marketConf.Register register.Name = pb.MarketService_ServiceDesc.ServiceName dis := discovery.NewConsulDiscovery(client) - if err := dis.Registry(grpcServer, register); err != nil { + if deregister, err := dis.Registry(grpcServer, register); err != nil { panic(err) + } else { + defer deregister() } // run grpc server diff --git a/cmd/sig-admin/main.go b/cmd/sig-admin/main.go index 6d2df84..620dc33 100644 --- a/cmd/sig-admin/main.go +++ b/cmd/sig-admin/main.go @@ -38,6 +38,9 @@ func main() { if err := gpcGenericClientFactory.Init(); err != nil { panic(err) } + if err := dis.WatchServices(gpcGenericClientFactory.RefreshService); err != nil { + panic(err) + } sigServer := sig.NewSigServer(gpcGenericClientFactory) if err := sigServer.Init(); err != nil { diff --git a/cmd/trading/main.go b/cmd/trading/main.go index a98122b..8fa4f46 100644 --- a/cmd/trading/main.go +++ b/cmd/trading/main.go @@ -99,8 +99,10 @@ func main() { // consul 服务注册 register := tradingConf.Register register.Name = pb.TradingService_ServiceDesc.ServiceName - if err := dis.Registry(grpcServer, register); err != nil { + if deregister, err := dis.Registry(grpcServer, register); err != nil { panic(err) + } else { + defer deregister() } // run grpc server diff --git a/config/exchange.toml b/config/exchange.toml index c0a42b7..4d567a5 100644 --- a/config/exchange.toml +++ b/config/exchange.toml @@ -18,8 +18,8 @@ receiveBuffer = 4096 marketSubscribeLimit = 16 consumeBatch = 1024 consumeLater = 2000 # 时间到达later或者数据累计到batch触发consume -httpProxy = "http://192.168.1.5:7890" -# httpProxy = "http://10.255.183.209:7890" +# httpProxy = "http://192.168.1.5:7890" +httpProxy = "http://10.255.183.209:7890" # 模拟盘API交易地址如下: # REST:https://www.okx.com diff --git a/pkg/grpc/discovery/consul/resolver.go b/pkg/grpc/discovery/consul/resolver.go index e549112..7bcebad 100644 --- a/pkg/grpc/discovery/consul/resolver.go +++ b/pkg/grpc/discovery/consul/resolver.go @@ -132,6 +132,7 @@ func (r *consulResolver) watch() { addr := resolver.Address{ Addr: fmt.Sprintf("%s:%d", s.Service.Address, s.Service.Port), Attributes: attributes.New("service_id", s.Service.ID), + ServerName: s.Service.Service, } addrs = append(addrs, addr) } diff --git a/pkg/grpc/discovery/consul/watch.go b/pkg/grpc/discovery/consul/watch.go new file mode 100644 index 0000000..c32f648 --- /dev/null +++ b/pkg/grpc/discovery/consul/watch.go @@ -0,0 +1,112 @@ +package consul + +import ( + "math" + "sig-pub/pkg/utils/collect" + "sig-pub/pkg/utils/retry" + "sig-pub/pkg/zlog" + "time" + + "github.com/hashicorp/consul/api" + "github.com/hashicorp/consul/api/watch" +) + +// 定义watcher +type Watcher struct { + client *api.Client + wp *watch.Plan // 总的Services变化对应的Plan + watchers *collect.SyncMap[string, *watch.Plan] // 对已经进行监控的service作个记录 + handler func(serviceName string, passing bool) +} + +func NewWatcher(client *api.Client, handler func(serviceName string, passing bool)) *Watcher { + return &Watcher{ + client: client, + watchers: collect.NewSyncMap[string, *watch.Plan](), + handler: handler, + } +} + +// WatchServices +// watch services doc: https://www.consul.io/docs/dynamic-app-config/watches#services +func (w *Watcher) WatchServices() (err error) { + w.wp, err = watch.Parse(map[string]any{ + "type": "services", + }) + if err != nil { + return + } + w.wp.Handler = func(i uint64, data any) { + switch d := data.(type) { + // services + case map[string][]string: + services := make([]string, 0, len(d)) + for i := range d { + if i == "consul" { + continue + } + services = append(services, i) + if _, loaded := w.watchers.LoadOrStore(i, nil); !loaded { + w.registerServiceWatcher(i) + } + } + zlog.Infof("consul services update: %v", services) + + // remove unknown services from watchers + // var dels []string + // w.watchers.Range(func(k string, wp *watch.Plan) bool { + // wp.Stop() + // dels = append(dels, k) + // return true + // }) + // for _, svc := range dels { + // w.watchers.Delete(svc) + // } + default: + zlog.Debugf("unknown consul watch type: %#v", &d) + } + } + + go func() { + if err := w.wp.RunWithClientAndHclog(w.client, nil); err != nil { + zlog.Error("watch services error: ", err) + retry.DoWithFixDelay(math.MaxUint32, 3*time.Second, func(retryTimes uint32) (_ struct{}, err error) { + err = w.WatchServices() + return + }) + } + }() + return +} + +// 将consul新增的service加入,并监控 +// doc: https://www.consul.io/docs/dynamic-app-config/watches#service +func (w *Watcher) registerServiceWatcher(serviceName string) error { + wp, err := watch.Parse(map[string]any{ + "type": "service", + "service": serviceName, + }) + if err != nil { + return err + } + + // 定义service变化后所执行的程序(函数)handler + wp.Handler = func(idx uint64, data any) { + switch d := data.(type) { + case []*api.ServiceEntry: + for _, i := range d { + if w.handler != nil { + w.handler(i.Service.Service, i.Checks.AggregatedStatus() == "passing") // status: passing / critical + } + // fmt.Printf("service %s 已变化", i.Service.Service) + // // 打印service的状态 + // fmt.Println("service status: ", i.Checks.AggregatedStatus()) + } + } + } + // 启动监控 + go wp.RunWithClientAndHclog(w.client, nil) + + w.watchers.Store(serviceName, wp) + return nil +} diff --git a/pkg/grpc/discovery/consul_naming.go b/pkg/grpc/discovery/consul_naming.go index 604fa50..1514b3c 100644 --- a/pkg/grpc/discovery/consul_naming.go +++ b/pkg/grpc/discovery/consul_naming.go @@ -23,7 +23,8 @@ func ConsulDialUrl(svrName string) string { } type ConsulDiscovery struct { - client *api.Client + client *api.Client + watcher *consul.Watcher } func NewConsulDiscovery(client *api.Client) *ConsulDiscovery { @@ -32,7 +33,7 @@ func NewConsulDiscovery(client *api.Client) *ConsulDiscovery { } } -func (r *ConsulDiscovery) Registry(grpcServer grpc.ServiceRegistrar, register Server) (err error) { +func (r *ConsulDiscovery) Registry(grpcServer grpc.ServiceRegistrar, register Server) (deregister func(), err error) { if strings.Trim(register.Name, " ") == "" { err = errors.New("registry name is empty") return @@ -66,6 +67,10 @@ func (r *ConsulDiscovery) Registry(grpcServer grpc.ServiceRegistrar, register Se zlog.Errorf("registry service %s error: %v", register.Name, err) return } + deregister = func() { + err := r.client.Agent().ServiceDeregister(reg.ID) + zlog.Infof("service id %s deregister to consul, err=%v", reg.ID, err) + } zlog.Infof("service %s registered to consul", register.Name) return } @@ -74,3 +79,8 @@ func (r *ConsulDiscovery) Resolver() (builder resolver.Builder) { builder = consul.NewConsulBuilder(r.client) return } + +func (r *ConsulDiscovery) WatchServices(handler func(service string, passing bool)) (err error) { + r.watcher = consul.NewWatcher(r.client, handler) + return r.watcher.WatchServices() +} diff --git a/pkg/grpc/generic/generic_client.go b/pkg/grpc/generic/generic_client.go index 9252488..a4bffdd 100644 --- a/pkg/grpc/generic/generic_client.go +++ b/pkg/grpc/generic/generic_client.go @@ -32,7 +32,7 @@ func NewGpcGenericClient(serviceName string, conn *grpc.ClientConn) *GrpcGeneric } } -func (c *GrpcGenericClient) Init(ctx context.Context) (err error) { +func (c *GrpcGenericClient) InitStub(ctx context.Context) (err error) { client := grpcreflect.NewClientV1(ctx, refv1.NewServerReflectionClient(c.conn)) marketServiceSymbol, err := client.FileContainingSymbol(protoreflect.FullName(c.serviceName)) if err != nil { diff --git a/pkg/grpc/generic/generic_client_factory.go b/pkg/grpc/generic/generic_client_factory.go index 30bc093..b5feadd 100644 --- a/pkg/grpc/generic/generic_client_factory.go +++ b/pkg/grpc/generic/generic_client_factory.go @@ -4,6 +4,9 @@ import ( "context" "fmt" "sig-pub/pkg/utils/collect" + "sig-pub/pkg/utils/retry" + "sig-pub/pkg/zlog" + "time" "google.golang.org/grpc" ) @@ -36,7 +39,7 @@ func (f *GrpcGenericClientFactory) NewClient(ctx context.Context, serviceName st return } client = NewGpcGenericClient(serviceName, conn) - err = client.Init(ctx) + err = client.InitStub(ctx) return } @@ -47,3 +50,23 @@ func (f *GrpcGenericClientFactory) GetClient(ctx context.Context, serviceName st }) return } + +// RefreshService 服务重新注册时可能有更改, 刷新旧的 grpc generic stub +func (f *GrpcGenericClientFactory) RefreshService(serviceName string, passing bool) { + zlog.Debugf("service status updated: %s, health=%v", serviceName, passing) + if passing { + service, ok := f.clientCache.Load(serviceName) + if ok { + go retry.DoWithFixDelay(10, 2*time.Second, func(_ uint32) (_ struct{}, err error) { + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + if err = service.InitStub(ctx); err != nil { + zlog.Errorf("refresh service stub error: %s, %v", serviceName, err) + } else { + zlog.Debugf("refresh service stub success: %s", serviceName) + } + return + }) + } + } +} diff --git a/pkg/trade/risk_strategy.go b/pkg/trade/risk_strategy.go index fb7a405..94077bf 100644 --- a/pkg/trade/risk_strategy.go +++ b/pkg/trade/risk_strategy.go @@ -5,7 +5,12 @@ import ( ) type IRickStrategy interface { - OnSingle(signalSide types.Side) + OnSignal(signalSide types.Side) (ok bool, cause Cause) +} + +type RiskStrategyParam struct { + SkipOnSideOpposite bool // 当前持有反方向单时 + SkipOnSideSame bool // 当前持有相同方向单时 } // RiskStrategy 风险管理策略 diff --git a/pkg/trade/trade_strategy.go b/pkg/trade/trade_strategy.go index c28377f..23a45ca 100644 --- a/pkg/trade/trade_strategy.go +++ b/pkg/trade/trade_strategy.go @@ -3,6 +3,7 @@ package trade // ITradeStrategy 下单策略 // 根据购买信号和账户信息生成下单参数 type ITradeStrategy interface { + OnSingal() } type TradeStrategyParam struct { diff --git a/pkg/utils/collect/sync_map.go b/pkg/utils/collect/sync_map.go index 1655093..48ce039 100644 --- a/pkg/utils/collect/sync_map.go +++ b/pkg/utils/collect/sync_map.go @@ -42,5 +42,11 @@ func (m *SyncMap[K, V]) CompareAndSwap(k K, old V, new V) (swapped bool) { } func (m *SyncMap[K, V]) LoadOrStore(k K, v V) (actual any, loaded bool) { - return m.m.LoadOrStore(k, v) + var cur any + cur, loaded = m.m.LoadOrStore(k, v) + if !loaded { + return + } + actual = cur.(V) + return }