grpcclient_test.go 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266
  1. package grpcclient
  2. import (
  3. "context"
  4. "io"
  5. "net"
  6. "sync/atomic"
  7. "testing"
  8. "time"
  9. pbv1 "git3.techno-world.net/lrosales/broad-announce/gen/go/broadannounce/v1"
  10. "google.golang.org/grpc"
  11. "google.golang.org/grpc/codes"
  12. "google.golang.org/grpc/credentials/insecure"
  13. "google.golang.org/grpc/metadata"
  14. "google.golang.org/grpc/test/bufconn"
  15. )
  16. const bufnetSize = 1 << 20 // 1 MB
  17. // mockIngestServer implements pbv1.IngestServer for testing.
  18. type mockIngestServer struct {
  19. pbv1.UnimplementedIngestServer
  20. alertsReceived atomic.Int64
  21. rateLimitAfter int64 // after N alerts, start rate-limiting
  22. authKey string
  23. authFail atomic.Bool
  24. }
  25. func (m *mockIngestServer) StreamAlerts(stream pbv1.Ingest_StreamAlertsServer) error {
  26. md, ok := metadata.FromIncomingContext(stream.Context())
  27. if !ok {
  28. return io.EOF
  29. }
  30. if m.authFail.Load() {
  31. return io.EOF
  32. }
  33. if len(m.authKey) > 0 {
  34. vals := md.Get("authorization")
  35. if len(vals) == 0 || vals[0] != "Bearer "+m.authKey {
  36. return io.EOF
  37. }
  38. }
  39. for {
  40. alert, err := stream.Recv()
  41. if err == io.EOF {
  42. return nil
  43. }
  44. if err != nil {
  45. return err
  46. }
  47. m.alertsReceived.Add(1)
  48. ack := &pbv1.Ack{
  49. AlertId: alert.DedupeKey,
  50. DedupeKey: alert.DedupeKey,
  51. AcceptedAtMs: time.Now().UnixMilli(),
  52. }
  53. if m.rateLimitAfter > 0 && m.alertsReceived.Load() > m.rateLimitAfter {
  54. ack.Result = &pbv1.Ack_Error{
  55. Error: &pbv1.Error{
  56. Code: pbv1.Error_RATE_LIMITED,
  57. Message: "rate limited by test server",
  58. RetryAfterMs: 10,
  59. },
  60. }
  61. } else {
  62. ack.Result = &pbv1.Ack_Ok{
  63. Ok: &pbv1.Ok{},
  64. }
  65. }
  66. if err := stream.Send(ack); err != nil {
  67. return err
  68. }
  69. }
  70. }
  71. func newTestServer(t *testing.T) (*grpc.Server, *bufconn.Listener) {
  72. lis := bufconn.Listen(bufnetSize)
  73. srv := grpc.NewServer()
  74. pbv1.RegisterIngestServer(srv, &mockIngestServer{authKey: "acme-001:prom:s3cret"})
  75. go srv.Serve(lis)
  76. return srv, lis
  77. }
  78. func dialBufconn(ctx context.Context, t *testing.T, lis *bufconn.Listener) *grpc.ClientConn {
  79. conn, err := grpc.NewClient(
  80. "passthrough://bufconn",
  81. grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) {
  82. return lis.Dial()
  83. }),
  84. grpc.WithTransportCredentials(insecure.NewCredentials()),
  85. )
  86. if err != nil {
  87. t.Fatalf("grpc.NewClient: %v", err)
  88. }
  89. return conn
  90. }
  91. func TestClient_New(t *testing.T) {
  92. c, err := New("localhost:9090",
  93. WithAPIKey("acme-001:prom:s3cret"),
  94. WithMaxRetries(5),
  95. WithInsecure(),
  96. )
  97. if err != nil {
  98. t.Fatalf("New() error = %v", err)
  99. }
  100. if c == nil {
  101. t.Fatal("New() returned nil client")
  102. }
  103. if c.maxRetries != 5 {
  104. t.Errorf("maxRetries = %d, want 5", c.maxRetries)
  105. }
  106. c.Close()
  107. }
  108. func TestClient_OptionDefaults(t *testing.T) {
  109. c, err := New("localhost:9090", WithInsecure())
  110. if err != nil {
  111. t.Fatalf("New() error = %v", err)
  112. }
  113. if c.maxRetries != 3 {
  114. t.Errorf("default maxRetries = %d, want 3", c.maxRetries)
  115. }
  116. c.Close()
  117. }
  118. func TestStream_SendRecv(t *testing.T) {
  119. ctx := context.Background()
  120. srv, lis := newTestServer(t)
  121. defer srv.Stop()
  122. conn := dialBufconn(ctx, t, lis)
  123. defer conn.Close()
  124. client := pbv1.NewIngestClient(conn)
  125. md := metadata.Pairs("authorization", "Bearer acme-001:prom:s3cret")
  126. sctx := metadata.NewOutgoingContext(ctx, md)
  127. stream, err := client.StreamAlerts(sctx)
  128. if err != nil {
  129. t.Fatalf("StreamAlerts() error = %v", err)
  130. }
  131. alert := &pbv1.Alert{
  132. CompanyId: "acme-001",
  133. SourceId: "prom",
  134. Severity: "critical",
  135. Title: "test alert",
  136. DedupeKey: "dk-001",
  137. ClientTsMs: time.Now().UnixMilli(),
  138. }
  139. if err := stream.Send(alert); err != nil {
  140. t.Fatalf("stream.Send() error = %v", err)
  141. }
  142. ack, err := stream.Recv()
  143. if err != nil {
  144. t.Fatalf("stream.Recv() error = %v", err)
  145. }
  146. if ack.DedupeKey != "dk-001" {
  147. t.Errorf("ack.DedupeKey = %q, want %q", ack.DedupeKey, "dk-001")
  148. }
  149. stream.CloseSend()
  150. }
  151. func TestStream_CloseSend(t *testing.T) {
  152. ctx := context.Background()
  153. srv, lis := newTestServer(t)
  154. defer srv.Stop()
  155. conn := dialBufconn(ctx, t, lis)
  156. defer conn.Close()
  157. client := pbv1.NewIngestClient(conn)
  158. md := metadata.Pairs("authorization", "Bearer acme-001:prom:s3cret")
  159. sctx := metadata.NewOutgoingContext(ctx, md)
  160. stream, err := client.StreamAlerts(sctx)
  161. if err != nil {
  162. t.Fatalf("StreamAlerts() error = %v", err)
  163. }
  164. if err := stream.CloseSend(); err != nil {
  165. t.Errorf("CloseSend() error = %v", err)
  166. }
  167. _, err = stream.Recv()
  168. if err != io.EOF {
  169. t.Errorf("Recv() after CloseSend = %v, want io.EOF", err)
  170. }
  171. }
  172. func TestRetryable(t *testing.T) {
  173. tests := []struct {
  174. code codes.Code
  175. want bool
  176. }{
  177. {codes.Unavailable, true},
  178. {codes.ResourceExhausted, true},
  179. {codes.Internal, true},
  180. {codes.OK, false},
  181. {codes.InvalidArgument, false},
  182. {codes.NotFound, false},
  183. {codes.Unauthenticated, false},
  184. }
  185. for _, tt := range tests {
  186. t.Run(tt.code.String(), func(t *testing.T) {
  187. if got := retryable(tt.code); got != tt.want {
  188. t.Errorf("retryable(%v) = %v, want %v", tt.code, got, tt.want)
  189. }
  190. })
  191. }
  192. }
  193. func TestStream_SendRecv_MultipleAlerts(t *testing.T) {
  194. ctx := context.Background()
  195. srv, lis := newTestServer(t)
  196. defer srv.Stop()
  197. conn := dialBufconn(ctx, t, lis)
  198. defer conn.Close()
  199. client := pbv1.NewIngestClient(conn)
  200. md := metadata.Pairs("authorization", "Bearer acme-001:prom:s3cret")
  201. sctx := metadata.NewOutgoingContext(ctx, md)
  202. stream, err := client.StreamAlerts(sctx)
  203. if err != nil {
  204. t.Fatalf("StreamAlerts() error = %v", err)
  205. }
  206. const n = 5
  207. for i := 0; i < n; i++ {
  208. alert := &pbv1.Alert{
  209. CompanyId: "acme-001",
  210. SourceId: "prom",
  211. Severity: "info",
  212. Title: "test alert",
  213. DedupeKey: "dk-multi-" + string(rune('0'+i)),
  214. ClientTsMs: time.Now().UnixMilli(),
  215. }
  216. if err := stream.Send(alert); err != nil {
  217. t.Fatalf("stream.Send() error = %v", err)
  218. }
  219. ack, err := stream.Recv()
  220. if err != nil {
  221. t.Fatalf("stream.Recv() error = %v", err)
  222. }
  223. if ack.DedupeKey != alert.DedupeKey {
  224. t.Errorf("ack.DedupeKey = %q, want %q", ack.DedupeKey, alert.DedupeKey)
  225. }
  226. }
  227. stream.CloseSend()
  228. }
  229. // Verify generated types implement the expected interfaces.
  230. func TestProtoInterfaces(t *testing.T) {
  231. var _ pbv1.IngestClient = nil
  232. var _ pbv1.IngestServer = nil
  233. }