package executor import ( "context" "errors" "fmt" "maps" "sync" "github.com/google/uuid" "tangled.org/core/spindle/quota" millproto "tangled.org/core/spindle/mill/proto" millv1 "tangled.org/core/spindle/mill/proto/gen" ) type QuotaClient struct { mu sync.Mutex pending map[string]chan *millv1.QuotaResponse sendFn func(*millproto.Message) error } func NewQuotaClient() *QuotaClient { return &QuotaClient{ pending: make(map[string]chan *millv1.QuotaResponse), } } func (c *QuotaClient) SetSendFn(sendFn func(*millproto.Message) error) { c.mu.Lock() defer c.mu.Unlock() c.sendFn = sendFn } func (c *QuotaClient) OnDisconnect() { c.mu.Lock() oldPending := c.pending c.pending = make(map[string]chan *millv1.QuotaResponse) c.sendFn = nil c.mu.Unlock() for _, ch := range oldPending { select { case ch <- nil: default: } } } func (c *QuotaClient) HandleResponse(reqID string, resp *millv1.QuotaResponse) { c.mu.Lock() ch, ok := c.pending[reqID] c.mu.Unlock() if ok && resp != nil { select { case ch <- resp: default: } } } func (c *QuotaClient) ForLease(leaseID string) quota.ReservationStore { return &leaseQuotaStore{client: c, leaseID: leaseID} } type leaseQuotaStore struct { client *QuotaClient leaseID string } var _ quota.ReservationStore = (*leaseQuotaStore)(nil) func (s *leaseQuotaStore) Reserve(ctx context.Context, req quota.ReserveRequest) (quota.Reservation, error) { if err := quota.Validate(req); err != nil { return quota.Reservation{}, err } if s.leaseID == "" { return quota.Reservation{}, errors.New("empty workflow lease") } resp, err := s.client.roundTrip(ctx, &millv1.QuotaRequest{ Operation: millv1.QuotaOperation_QUOTA_OPERATION_RESERVE, LeaseId: s.leaseID, Kind: string(req.Kind), Key: req.Key, Resources: maps.Clone(req.Resources), }) if err != nil { return quota.Reservation{}, err } if resp.GetAllowed() && resp.GetTemporary() { return quota.Reservation{}, errors.New("mill allowed a reservation and deferred it at once") } return quota.Reservation{ ID: resp.GetReservationId(), Allowed: resp.GetAllowed(), Temporary: resp.GetTemporary(), Reason: resp.GetReason(), Resource: resp.GetResource(), }, nil } func (s *leaseQuotaStore) BeginCommit(ctx context.Context, reservationID string) error { return s.transition(ctx, millv1.QuotaOperation_QUOTA_OPERATION_BEGIN_COMMIT, reservationID) } func (s *leaseQuotaStore) Commit(ctx context.Context, reservationID string) error { return s.transition(ctx, millv1.QuotaOperation_QUOTA_OPERATION_COMMIT, reservationID) } func (s *leaseQuotaStore) Release(ctx context.Context, reservationID string) error { return s.transition(ctx, millv1.QuotaOperation_QUOTA_OPERATION_RELEASE, reservationID) } func (s *leaseQuotaStore) transition(ctx context.Context, operation millv1.QuotaOperation, reservationID string) error { if reservationID == "" { return nil } if s.leaseID == "" { return fmt.Errorf("reservation %q is not associated with a live lease", reservationID) } _, err := s.client.roundTrip(ctx, &millv1.QuotaRequest{ Operation: operation, LeaseId: s.leaseID, ReservationId: reservationID, }) return err } func (c *QuotaClient) roundTrip(ctx context.Context, req *millv1.QuotaRequest) (*millv1.QuotaResponse, error) { reqID := uuid.NewString() req.RequestId = reqID ch := make(chan *millv1.QuotaResponse, 1) c.mu.Lock() send := c.sendFn if send == nil { c.mu.Unlock() return nil, errors.New("quota client not connected") } c.pending[reqID] = ch c.mu.Unlock() defer func() { c.mu.Lock() delete(c.pending, reqID) c.mu.Unlock() }() if err := send(&millproto.Message{QuotaReq: req}); err != nil { return nil, fmt.Errorf("send quota request: %w", err) } select { case <-ctx.Done(): return nil, ctx.Err() case resp := <-ch: if resp == nil { return nil, errors.New("quota client disconnected") } if resp.GetError() != "" { return nil, errors.New(resp.GetError()) } return resp, nil } }