| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295 |
- package grpcserver
- import (
- "context"
- "encoding/json"
- "io"
- "sync"
- "sync/atomic"
- "time"
- pbv1 "git3.techno-world.net/lrosales/broad-announce/gen/go/broadannounce/v1"
- "git3.techno-world.net/lrosales/broad-announce/internal/pipeline"
- )
- // Server implements broadannounce.v1.IngestServer.
- type Server struct {
- pbv1.UnimplementedIngestServer
- // Dep is the shared pipeline dependencies. All goroutines write to it
- // concurrently — it is safe for concurrent use.
- Dep pipeline.Deps
- MaxInflight int // max concurrent in-flight messages per stream (default 256)
- Logger interface{ Info(msg string, args ...any) }
- }
- // StreamAlerts is a bidirectional stream: client sends zero or more Alerts,
- // server sends exactly one Ack per Alert received.
- //
- // Bounded concurrency: each message runs in a goroutine controlled by a
- // semaphore channel of size MaxInflight. This is how we achieve 10k+/s on
- // a single stream — the bottleneck is the Redis dedupe + NATS publish (~2ms),
- // so 256-way concurrency gives us ~128k msg/s theoretical max.
- //
- // Backpressure: if the sem is full, we immediately send Error.RATE_LIMITED
- // without entering the pipeline. The client is expected to honour retry_after_ms.
- func (s *Server) StreamAlerts(stream pbv1.Ingest_StreamAlertsServer) error {
- ctx := stream.Context()
- // Authenticate the stream.
- src, err := authenticate(ctx, s.Dep.Sources)
- if err != nil {
- return err
- }
- sourceKey := src.CompanyID + ":" + src.SourceID
- // Per-stream semaphore: bounds concurrent processing to MaxInflight.
- if s.MaxInflight <= 0 {
- s.MaxInflight = 256
- }
- sem := make(chan struct{}, s.MaxInflight)
- // Per-stream ackCh: each goroutine sends its Ack here; the recv loop
- // forwards to the client. Buffer = MaxInflight so goroutines never block
- // on sending (only on the semaphore).
- ackCh := make(chan *pbv1.Ack, s.MaxInflight)
- // Track active goroutines for graceful shutdown.
- var wg sync.WaitGroup
- // streamErr holds the first non-nil error from goroutines or the recv loop.
- var streamErr atomic.Value // holds error
- setErr := func(err error) {
- if err != nil {
- streamErr.CompareAndSwap(nil, err)
- }
- }
- // Increment active-streams gauge.
- s.Dep.Metrics.StreamsActive.Add(1)
- defer s.Dep.Metrics.StreamsActive.Add(-1)
- // goroutine: receive loop — reads from client, dispatches to pipeline.
- go func() {
- for {
- alert, err := stream.Recv()
- if err == io.EOF {
- // Client called CloseSend. Drain in-flight, then exit.
- close(ackCh)
- wg.Wait()
- return
- }
- if err != nil {
- setErr(err)
- close(ackCh)
- wg.Wait()
- return
- }
- // Validate we have a non-empty Alert before spending goroutines.
- if alert.CompanyId == "" || alert.SourceId == "" {
- ackCh <- &pbv1.Ack{
- DedupeKey: alert.DedupeKey,
- Result: &pbv1.Ack_Error{
- Error: &pbv1.Error{
- Code: pbv1.Error_INVALID,
- Message: "company_id and source_id are required",
- },
- },
- }
- continue
- }
- // Per-stream rate limit (before entering the pipeline).
- // This is a fast rejection path that doesn't consume goroutines.
- rateLimitKey := "grpc_source:" + sourceKey
- allowed, retryAfter, _ := s.Dep.Limiter.Allow(ctx, rateLimitKey, src.RateLimitPerSec)
- if !allowed {
- s.Dep.Metrics.GRPCRateLimited.WithLabelValues(src.SourceID).Inc()
- ackCh <- &pbv1.Ack{
- AlertId: alert.DedupeKey,
- DedupeKey: alert.DedupeKey,
- Result: &pbv1.Ack_Error{
- Error: &pbv1.Error{
- Code: pbv1.Error_RATE_LIMITED,
- Message: "per-stream rate limit exceeded",
- RetryAfterMs: int32(retryAfter.Milliseconds()),
- },
- },
- }
- continue
- }
- // Acquire semaphore slot. If full, backpressure immediately.
- select {
- case sem <- struct{}{}:
- // Proceed.
- default:
- s.Dep.Metrics.GRPCRateLimited.WithLabelValues(src.SourceID).Inc()
- ackCh <- &pbv1.Ack{
- AlertId: alert.DedupeKey,
- DedupeKey: alert.DedupeKey,
- Result: &pbv1.Ack_Error{
- Error: &pbv1.Error{
- Code: pbv1.Error_RATE_LIMITED,
- Message: "in-flight capacity reached",
- RetryAfterMs: 5, // 5ms backoff, client should retry
- },
- },
- }
- continue
- }
- wg.Add(1)
- go func(alert *pbv1.Alert) {
- defer func() {
- <-sem
- wg.Done()
- }()
- start := time.Now()
- ack := s.processAlert(ctx, alert, src)
- // Record ack latency metric.
- s.Dep.Metrics.GRPCAckLatency.WithLabelValues(src.SourceID).Observe(
- time.Since(start).Seconds())
- // Send Ack to client. Non-blocking — if the channel is full the
- // recv loop has a backlog and we should not block the goroutine.
- select {
- case ackCh <- ack:
- default:
- // Channel full; log and drop.
- s.Dep.Logger.Info("grpc ack channel full, dropping", "alert_id", ack.AlertId)
- }
- }(alert)
- }
- }()
- // recv loop: forward Acks from goroutines to the client.
- // Also watches for goroutine errors.
- for ack := range ackCh {
- if err := stream.Send(ack); err != nil {
- return err
- }
- }
- // Check if the stream exited due to a goroutine error.
- if se := streamErr.Load().(error); se != nil {
- return se
- }
- return nil
- }
- // processAlert runs the shared pipeline on a single gRPC Alert and returns
- // the corresponding Ack. sig="" because gRPC auth is handled by authenticate()
- // before entering this function.
- func (s *Server) processAlert(ctx context.Context, alert *pbv1.Alert, src *pipeline.SourceConfig) *pbv1.Ack {
- body, err := alertToJSON(alert)
- if err != nil {
- return errorAck(alert.DedupeKey, pbv1.Error_INVALID, "marshal failed: "+err.Error(), 0)
- }
- // Override the Sources map for this call to use the authenticated source.
- // This ensures company_id/source_id from the API key match the alert body.
- sourceKey := src.CompanyID + ":" + src.SourceID
- origSources := s.Dep.Sources
- s.Dep.Sources = map[string]pipeline.SourceConfig{sourceKey: *src}
- res := s.Dep.Process(ctx, body, "" /* no HMAC for gRPC */)
- // Restore the original Sources map.
- s.Dep.Sources = origSources
- if !res.Accepted {
- code := errorCode(res.RejectReason)
- var retryAfter int32
- if res.RejectReason == "rate_limited_source" || res.RejectReason == "rate_limited_company" {
- if n, _ := parseRetryAfter(res.Detail); n > 0 {
- retryAfter = int32(n)
- }
- }
- return errorAck(alert.DedupeKey, code, res.Detail, retryAfter)
- }
- return &pbv1.Ack{
- AlertId: res.AlertID,
- DedupeKey: alert.DedupeKey,
- DedupeCount: res.DedupeCount,
- AcceptedAtMs: time.Now().UnixMilli(),
- Result: &pbv1.Ack_Ok{Ok: &pbv1.Ok{}},
- }
- }
- // alertToJSON converts a protobuf Alert to JSON bytes for the pipeline.
- func alertToJSON(a *pbv1.Alert) ([]byte, error) {
- m := map[string]any{
- "company_id": a.CompanyId,
- "source_id": a.SourceId,
- }
- if a.DedupeKey != "" {
- m["dedupe_key"] = a.DedupeKey
- }
- if a.Severity != "" {
- m["severity"] = a.Severity
- }
- if a.Category != "" {
- m["category"] = a.Category
- }
- if a.Title != "" {
- m["title"] = a.Title
- }
- if a.Body != "" {
- m["body"] = a.Body
- }
- if len(a.Data) > 0 {
- m["data"] = a.Data
- }
- if a.ClientTsMs > 0 {
- m["client_ts_ms"] = a.ClientTsMs
- }
- return json.Marshal(m)
- }
- func errorAck(dedupeKey string, code pbv1.Error_Code, message string, retryAfterMs int32) *pbv1.Ack {
- return &pbv1.Ack{
- DedupeKey: dedupeKey,
- Result: &pbv1.Ack_Error{
- Error: &pbv1.Error{
- Code: code,
- Message: message,
- RetryAfterMs: retryAfterMs,
- },
- },
- }
- }
- // errorCode maps pipeline reject reasons to gRPC Error codes.
- func errorCode(reason string) pbv1.Error_Code {
- switch reason {
- case "unknown_source", "bad_signature":
- return pbv1.Error_UNAUTHENTICATED
- case "rate_limited_source", "rate_limited_company":
- return pbv1.Error_RATE_LIMITED
- case "invalid", "invalid_json", "quarantined":
- return pbv1.Error_INVALID
- case "circuit_open", "broker_unavailable", "marshal_failed":
- return pbv1.Error_INTERNAL
- default:
- return pbv1.Error_UNKNOWN
- }
- }
- func parseRetryAfter(s string) (int, bool) {
- var n int
- var parsed bool
- for _, c := range s {
- if c >= '0' && c <= '9' {
- n = n*10 + int(c-'0')
- parsed = true
- } else {
- break
- }
- }
- return n, parsed
- }
|