package health import ( "context" "errors" "math" "strings" "time" controlplanev1 "proxy-pool/gen/controlplane/v1" "proxy-pool/internal/domain/activitypool" healthDomain "proxy-pool/internal/domain/health" proxyDomain "proxy-pool/internal/domain/proxy" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" "google.golang.org/protobuf/types/known/durationpb" "google.golang.org/protobuf/types/known/timestamppb" ) var ErrInvalidGRPCHandler = errors.New("invalid checker grpc handler") // CheckerIdentityAuthorizer proves that the calling mTLS identity owns the // checker ID in the request. It is deliberately separate from worker session // identity because Checkers do not own gateway snapshots. type CheckerIdentityAuthorizer interface { AuthorizeChecker(context.Context, string) error } type GRPCHandlerOptions struct { MaxObservationsPerBatch int MaxTasksPerClaim int TaskBroker TaskBroker Metrics healthDomain.TaskMetricsObserver Now func() time.Time } func DefaultGRPCHandlerOptions() GRPCHandlerOptions { return GRPCHandlerOptions{ MaxObservationsPerBatch: 1_000, MaxTasksPerClaim: defaultMaxTasksPerClaim, Now: time.Now, } } // GRPCHandler exposes the Controller's checker task and fact boundaries. The // Broker is optional while a deployment has no shared task store; reporting // remains available for its existing external fact integrations. type GRPCHandler struct { controlplanev1.UnimplementedCheckerControlPlaneServer reducer *Reducer identity CheckerIdentityAuthorizer options GRPCHandlerOptions } func NewGRPCHandler( reducer *Reducer, identity CheckerIdentityAuthorizer, options GRPCHandlerOptions, ) (*GRPCHandler, error) { if reducer == nil || nilInterface(identity) { return nil, ErrInvalidGRPCHandler } defaults := DefaultGRPCHandlerOptions() if options.MaxObservationsPerBatch == 0 { options.MaxObservationsPerBatch = defaults.MaxObservationsPerBatch } if options.MaxTasksPerClaim == 0 { options.MaxTasksPerClaim = defaults.MaxTasksPerClaim } if options.Now == nil { options.Now = defaults.Now } if options.MaxObservationsPerBatch < 0 || options.MaxTasksPerClaim < 0 { return nil, ErrInvalidGRPCHandler } return &GRPCHandler{reducer: reducer, identity: identity, options: options}, nil } func (handler *GRPCHandler) StreamCheckTasks( request *controlplanev1.StreamCheckTasksRequest, stream controlplanev1.CheckerControlPlane_StreamCheckTasksServer, ) error { if handler == nil || handler.reducer == nil || nilInterface(handler.identity) || request == nil || stream == nil || handler.options.Now == nil || request.GetCheckerId() == "" || request.GetInstanceId() == "" || request.GetMaxInFlight() == 0 || request.GetMaxInFlight() > uint32(handler.options.MaxTasksPerClaim) { return healthGRPCError(ErrInvalidGRPCHandler) } ctx := stream.Context() if ctx == nil { return healthGRPCError(ErrInvalidGRPCHandler) } if err := handler.identity.AuthorizeChecker(ctx, request.GetCheckerId()); err != nil { return status.Error(codes.PermissionDenied, "checker identity is not authorized") } if nilInterface(handler.options.TaskBroker) { return status.Error(codes.Unavailable, "checker task broker unavailable") } levels := make([]healthDomain.Level, 0, len(request.GetSupportedLevels())) for _, value := range request.GetSupportedLevels() { level, ok := decodeCheckLevel(value) if !ok { return healthGRPCError(ErrInvalidTaskClaim) } levels = append(levels, level) } tasks, err := handler.options.TaskBroker.Claim(ctx, TaskClaim{ CheckerID: request.GetCheckerId(), InstanceID: request.GetInstanceId(), MaxInFlight: int(request.GetMaxInFlight()), SupportedLevels: levels, }) if err != nil { return healthGRPCError(err) } now := handler.options.Now().UTC() if now.IsZero() { return healthGRPCError(ErrInvalidGRPCHandler) } for _, task := range tasks { wire, wireErr := wireCheckTask(task, now) if wireErr != nil { return healthGRPCError(wireErr) } if sendErr := stream.Send(wire); sendErr != nil { return sendErr } if handler.options.Metrics != nil { handler.options.Metrics.ObserveTaskDispatch(task.Level, 1) } } return nil } func (handler *GRPCHandler) ReportObservations( ctx context.Context, request *controlplanev1.ObservationBatch, ) (*controlplanev1.ReportObservationsResponse, error) { if ctx == nil || handler == nil || handler.reducer == nil || nilInterface(handler.identity) || request == nil || request.GetCheckerId() == "" || len(request.GetObservations()) == 0 || len(request.GetObservations()) > handler.options.MaxObservationsPerBatch { return nil, healthGRPCError(ErrInvalidGRPCHandler) } if err := handler.identity.AuthorizeChecker(ctx, request.GetCheckerId()); err != nil { return nil, status.Error(codes.PermissionDenied, "checker identity is not authorized") } response := &controlplanev1.ReportObservationsResponse{} for _, item := range request.GetObservations() { observation, err := decodeHealthObservation(item) if err == nil && !nilInterface(handler.options.TaskBroker) { now := handler.options.Now().UTC() if now.IsZero() { return nil, healthGRPCError(ErrInvalidGRPCHandler) } err = handler.options.TaskBroker.AuthorizeObservation(ctx, request.GetCheckerId(), item.GetLeaseToken(), observation, now) } if err == nil { _, err = handler.reducer.Apply(ctx, observation) } if err == nil && !nilInterface(handler.options.TaskBroker) { now := handler.options.Now().UTC() if now.IsZero() { return nil, healthGRPCError(ErrInvalidGRPCHandler) } err = handler.options.TaskBroker.CompleteObservation(ctx, request.GetCheckerId(), item.GetLeaseToken(), observation, now) } if err == nil { response.Accepted++ if handler.options.Metrics != nil { handler.options.Metrics.ObserveObservation(observation.Level, healthDomain.ObservationMetricAccepted) } continue } if rejectedObservationError(err) { response.Rejected++ if handler.options.Metrics != nil && observation.Level != "" { handler.options.Metrics.ObserveObservation(observation.Level, healthDomain.ObservationMetricRejected) } continue } return nil, healthGRPCError(err) } return response, nil } func decodeHealthObservation(item *controlplanev1.HealthObservation) (healthDomain.Observation, error) { if item == nil || item.GetLatency() == nil || item.GetLatency().CheckValid() != nil || item.GetObservedAt() == nil || item.GetObservedAt().CheckValid() != nil { return healthDomain.Observation{}, healthDomain.ErrInvalidObservation } level, ok := decodeCheckLevel(item.GetLevel()) if !ok { return healthDomain.Observation{}, healthDomain.ErrInvalidObservation } return healthDomain.Observation{ TaskID: item.GetTaskId(), ProxyID: item.GetProxyId(), Level: level, RoutingName: item.GetRoutingName(), TargetURL: item.GetTargetUrl(), Success: item.GetSuccess(), FailureClass: item.GetFailureClass(), Latency: item.GetLatency().AsDuration(), ObservedEgressIP: item.GetObservedEgressIp(), ObservedAt: item.GetObservedAt().AsTime(), }, nil } func decodeCheckLevel(level controlplanev1.CheckLevel) (healthDomain.Level, bool) { switch level { case controlplanev1.CheckLevel_CHECK_LEVEL_BASIC: return healthDomain.LevelBasic, true case controlplanev1.CheckLevel_CHECK_LEVEL_EGRESS: return healthDomain.LevelEgress, true case controlplanev1.CheckLevel_CHECK_LEVEL_TARGET: return healthDomain.LevelTarget, true default: return "", false } } func rejectedObservationError(err error) bool { return errors.Is(err, healthDomain.ErrInvalidObservation) || errors.Is(err, healthDomain.ErrInvalidFailureThreshold) || errors.Is(err, healthDomain.ErrStaleObservation) || errors.Is(err, healthDomain.ErrConflictingObservation) || errors.Is(err, activitypool.ErrInvalidHealthUpdate) || errors.Is(err, activitypool.ErrActivityNotFound) || errors.Is(err, ErrInvalidThreshold) || errors.Is(err, ErrUnconfiguredUpstream) || errors.Is(err, ErrTaskNotFound) || errors.Is(err, ErrTaskLeaseExpired) || errors.Is(err, ErrTaskLeaseNotOwned) || errors.Is(err, ErrTaskObservation) } func healthGRPCError(err error) error { switch { case errors.Is(err, context.Canceled): return status.Error(codes.Canceled, "checker control request canceled") case errors.Is(err, context.DeadlineExceeded): return status.Error(codes.DeadlineExceeded, "checker control request deadline exceeded") case errors.Is(err, ErrInvalidGRPCHandler), errors.Is(err, healthDomain.ErrInvalidObservation), errors.Is(err, healthDomain.ErrInvalidFailureThreshold), errors.Is(err, activitypool.ErrInvalidHealthUpdate), errors.Is(err, ErrInvalidThreshold), errors.Is(err, ErrInvalidTaskClaim), errors.Is(err, ErrInvalidLeasedTask): return status.Error(codes.InvalidArgument, "invalid checker observation batch") default: return status.Error(codes.Unavailable, "checker control plane unavailable") } } func wireCheckTask(task LeasedTask, now time.Time) (*controlplanev1.CheckTask, error) { if !validTaskIdentifier(task.TaskID) || !validTaskIdentifier(task.LeaseToken) || !validTaskIdentifier(task.ProxyID) || task.Host == "" || strings.TrimSpace(task.Host) != task.Host || task.Port == 0 || task.Attempts <= 0 || task.Attempts > math.MaxUint32 || !task.Deadline.After(now) || (task.SecretRef == "") != (task.CredentialVersion == "") { return nil, ErrInvalidLeasedTask } level, ok := wireCheckLevel(task.Level) if !ok { return nil, ErrInvalidLeasedTask } protocol, ok := wireTaskProtocol(task.Protocol) if !ok { return nil, ErrInvalidLeasedTask } switch task.Level { case healthDomain.LevelBasic: if task.RoutingName != "" || task.TargetURL != "" { return nil, ErrInvalidLeasedTask } case healthDomain.LevelEgress: targetURL, err := healthDomain.NormalizeEgressTarget(task.TargetURL) if task.RoutingName != "" || err != nil || targetURL != task.TargetURL { return nil, ErrInvalidLeasedTask } case healthDomain.LevelTarget: if _, err := healthDomain.NormalizeTargetProfile(healthDomain.TargetProfile{ RoutingName: task.RoutingName, TargetURL: task.TargetURL, }); err != nil { return nil, ErrInvalidLeasedTask } default: return nil, ErrInvalidLeasedTask } return &controlplanev1.CheckTask{ TaskId: task.TaskID, LeaseToken: task.LeaseToken, ProxyId: task.ProxyID, Protocol: protocol, Host: task.Host, Port: uint32(task.Port), SecretRef: task.SecretRef, CredentialVersion: task.CredentialVersion, Username: task.Username, Password: task.Password, Level: level, RoutingName: task.RoutingName, TargetUrl: task.TargetURL, Timeout: durationpb.New(task.Deadline.Sub(now)), Attempt: 1, MaxAttempts: uint32(task.Attempts), Deadline: timestamppb.New(task.Deadline.UTC()), }, nil } func wireCheckLevel(level healthDomain.Level) (controlplanev1.CheckLevel, bool) { switch level { case healthDomain.LevelBasic: return controlplanev1.CheckLevel_CHECK_LEVEL_BASIC, true case healthDomain.LevelEgress: return controlplanev1.CheckLevel_CHECK_LEVEL_EGRESS, true case healthDomain.LevelTarget: return controlplanev1.CheckLevel_CHECK_LEVEL_TARGET, true default: return controlplanev1.CheckLevel_CHECK_LEVEL_UNSPECIFIED, false } } func wireTaskProtocol(protocol proxyDomain.Scheme) (controlplanev1.ProxyProtocol, bool) { switch protocol { case proxyDomain.SchemeHTTP: return controlplanev1.ProxyProtocol_PROXY_PROTOCOL_HTTP, true case proxyDomain.SchemeHTTPS: return controlplanev1.ProxyProtocol_PROXY_PROTOCOL_HTTPS, true case proxyDomain.SchemeSOCKS5: return controlplanev1.ProxyProtocol_PROXY_PROTOCOL_SOCKS5, true default: return controlplanev1.ProxyProtocol_PROXY_PROTOCOL_UNSPECIFIED, false } }