Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
96 changes: 66 additions & 30 deletions fleetspeak/src/client/https/polling.go
Original file line number Diff line number Diff line change
Expand Up @@ -59,15 +59,19 @@ type Communicator struct {

certBytes []byte

wakeUp chan struct{}
pollComplete chan error
mu sync.Mutex
pollCancel context.CancelFunc
// Synchronization for Reset/Flush
wakeUp chan struct{}
mu sync.Mutex
pollDone chan struct{}
lastPollErr error
pollCancel context.CancelFunc
}

// Setup implements comms.Communicator.
func (c *Communicator) Setup(cl comms.Context) error {
c.cctx = cl
c.pollDone = make(chan struct{})
close(c.pollDone) // Start closed so Flush returns immediately if no messages.
return c.configure()
}

Expand Down Expand Up @@ -114,7 +118,6 @@ func (c *Communicator) configure() error {
c.ctx, c.done = context.WithCancel(context.Background())
c.clientCertificateHeader = si.ClientCertificateHeader
c.wakeUp = make(chan struct{}, 1)
c.pollComplete = make(chan error, 1)
c.certBytes = certBytes
return nil
}
Expand All @@ -141,13 +144,9 @@ func (c *Communicator) Reset() {
if c.pollCancel != nil {
c.pollCancel()
}
c.pollDone = make(chan struct{})
c.mu.Unlock()
c.hc.Transport.(*http.Transport).CloseIdleConnections()
// Drain pollComplete to ensure Flush waits for a new poll.
select {
case <-c.pollComplete:
default:
}
select {
case c.wakeUp <- struct{}{}:
default:
Expand All @@ -158,10 +157,23 @@ func (c *Communicator) Reset() {
func (c *Communicator) Flush(ctx context.Context) error {
log.InfoContextf(ctx, "Flush called")
for {
c.mu.Lock()
done := c.pollDone
c.mu.Unlock()

select {
case <-ctx.Done():
return ctx.Err()
case err := <-c.pollComplete:
case <-c.ctx.Done():
return c.ctx.Err()
case <-done:
c.mu.Lock()
err := c.lastPollErr
if err != nil && c.pollDone != done {
c.mu.Unlock()
continue
}
c.mu.Unlock()
return err
}
}
Expand Down Expand Up @@ -193,25 +205,48 @@ func (c *Communicator) processingLoop() {
// for the MinFailureDelay.
poll := func() {
var err error
c.mu.Lock()
select {
case <-c.pollDone:
c.pollDone = make(chan struct{})
default:
}
reqCtx, cancel := context.WithCancel(c.ctx)
c.pollCancel = cancel
myDone := c.pollDone
c.mu.Unlock()

defer func() {
select {
case c.pollComplete <- err:
default:
}
cancel()
c.mu.Lock()
c.pollCancel = nil
c.lastPollErr = err
close(myDone)
c.mu.Unlock()
}()

c.wd.Reset()
if c.cctx.CurrentID() != c.id {
c.configure()
}
var active bool
active, err = c.poll(toSend)
active, err = c.poll(reqCtx, toSend)
if err != nil {
c.mu.Lock()
wasReset := errors.Is(err, context.Canceled) && c.ctx.Err() == nil
c.mu.Unlock()

if wasReset {
return
}

log.Warningf("Failure during polling: %v", err)
for _, m := range toSend {
m.Nack()
}
toSend = nil
toSendSize = 0
oldestUnsent = time.Time{}

if (!lastPoll.IsZero()) && (time.Since(lastPoll) > time.Duration(c.conf.FailureSuicideTimeSeconds)*time.Second) {
// Die in the hopes that our replacement will be better configured, or otherwise have better luck.
Expand Down Expand Up @@ -310,6 +345,14 @@ func (c *Communicator) processingLoop() {
poll()
case m := <-c.cctx.Outbox():
t.Stop()
c.mu.Lock()
select {
case <-c.pollDone:
c.pollDone = make(chan struct{})
default:
}
c.mu.Unlock()

toSend = append(toSend, m)
toSendSize += 2 + proto.Size(m.M)
if toSendSize >= sendBytesThreshold ||
Expand All @@ -324,7 +367,7 @@ func (c *Communicator) processingLoop() {
}
}

func (c *Communicator) poll(toSend []comms.MessageInfo) (bool, error) {
func (c *Communicator) poll(ctx context.Context, toSend []comms.MessageInfo) (bool, error) {
var sent bool // records whether an interesting (non-LOW) priority message was sent.
msgs := make([]*fspb.Message, 0, len(toSend))
for _, m := range toSend {
Expand All @@ -347,9 +390,12 @@ func (c *Communicator) poll(toSend []comms.MessageInfo) (bool, error) {
}

for i, host := range c.hosts {
cd, err := c.pollHost(host, data)
cd, err := c.pollHost(ctx, host, data)
if err != nil {
log.Warningf("Error polling %q for ContactData: %v", host, err)
if ctx.Err() != nil {
return false, ctx.Err()
}
continue
}
if i != 0 {
Expand All @@ -367,7 +413,7 @@ func (c *Communicator) poll(toSend []comms.MessageInfo) (bool, error) {
return false, errors.New("unable to contact any server")
}

func (c *Communicator) pollHost(host string, data []byte) (*fspb.ContactData, error) {
func (c *Communicator) pollHost(ctx context.Context, host string, data []byte) (*fspb.ContactData, error) {
var sendErr, recvErr error
var sendSize, recvSize int
defer func() {
Expand Down Expand Up @@ -396,17 +442,7 @@ func (c *Communicator) pollHost(host string, data []byte) (*fspb.ContactData, er
if sendErr != nil {
return nil, sendErr
}
var reqCtx context.Context
c.mu.Lock()
reqCtx, c.pollCancel = context.WithCancel(c.ctx)
c.mu.Unlock()
defer func() {
c.mu.Lock()
c.pollCancel()
c.pollCancel = nil
c.mu.Unlock()
}()
req = req.WithContext(reqCtx)
req = req.WithContext(ctx)
SetContentEncoding(req.Header, c.conf.GetCompression())
if c.clientCertificateHeader != "" {
bc := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: c.certBytes})
Expand Down
Loading