mirror of
https://github.com/zarzet/SpotiFLAC-Mobile.git
synced 2026-09-02 16:20:57 +02:00
perf(extensions): coalesce signed session lifecycle
This commit is contained in:
@@ -476,12 +476,16 @@ func preflightExtensionDownloadSession(extensionID string) (bool, error) {
|
||||
if _, err := ext.lockReadyVM(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer ext.VMMu.Unlock()
|
||||
if ext.runtime == nil {
|
||||
runtime := ext.runtime
|
||||
ext.VMMu.Unlock()
|
||||
if runtime == nil {
|
||||
return false, fmt.Errorf("extension '%s' runtime is unavailable", extensionID)
|
||||
}
|
||||
|
||||
return ext.runtime.preflightSignedSession()
|
||||
// Preflight touches only the runtime's thread-safe HTTP/session state. Do
|
||||
// not hold the Goja VM lock across bootstrap network I/O: metadata/status
|
||||
// calls on the same extension must remain responsive while auth is slow.
|
||||
return runtime.preflightSignedSession()
|
||||
}
|
||||
|
||||
func DownloadWithExtensionsJSON(requestJSON string) (string, error) {
|
||||
|
||||
@@ -122,8 +122,14 @@ type signedSessionCoordinator struct {
|
||||
// exchangeInFlight serializes grant exchanges without keeping mu held over
|
||||
// HTTP or Retry-After backoff. Waiters observe the completion channel and
|
||||
// retry their state check after the owner commits or fails.
|
||||
exchangeInFlight bool
|
||||
exchangeDone chan struct{}
|
||||
exchangeInFlight bool
|
||||
exchangeDone chan struct{}
|
||||
bootstrapInFlight bool
|
||||
bootstrapDone chan struct{}
|
||||
bootstrapErr error
|
||||
refreshInFlight bool
|
||||
refreshDone chan struct{}
|
||||
refreshErr error
|
||||
}
|
||||
|
||||
func (c *signedSessionCoordinator) beginExchange(ctx context.Context) (func(), error) {
|
||||
@@ -899,6 +905,38 @@ func (r *extensionRuntime) signedSessionFetch(call goja.FunctionCall) goja.Value
|
||||
}
|
||||
coordinator.mu.Unlock()
|
||||
|
||||
// Refresh can involve a slow HTTP request. Coalesce it across parallel
|
||||
// extension runtimes without keeping the shared coordinator mutex locked,
|
||||
// then reload the committed generation before signing the request.
|
||||
if signedSessionRefreshDue(config, record) {
|
||||
if _, refreshErr := r.refreshSignedSessionCoalesced(config, coordinator); refreshErr != nil {
|
||||
LogWarn("SignedSession", "Session refresh failed for extension %s: %v", r.extensionID, refreshErr)
|
||||
}
|
||||
coordinator.mu.Lock()
|
||||
latest, loadErr := r.loadSignedSession(config)
|
||||
if loadErr != nil {
|
||||
coordinator.mu.Unlock()
|
||||
return r.vm.ToValue(map[string]any{"ok": false, "error": loadErr.Error()})
|
||||
}
|
||||
if !signedSessionRecordIsUsable(latest) || coordinator.generationIsBlocked(latest) {
|
||||
authURL, verificationErr := r.startSignedSessionVerificationLocked(
|
||||
config,
|
||||
coordinator,
|
||||
"signed-fetch-refresh",
|
||||
)
|
||||
coordinator.mu.Unlock()
|
||||
if authURL != "" {
|
||||
return r.signedSessionVerificationRequiredValue(authURL)
|
||||
}
|
||||
if verificationErr != nil {
|
||||
return r.vm.ToValue(map[string]any{"ok": false, "error": verificationErr.Error()})
|
||||
}
|
||||
return r.vm.ToValue(map[string]any{"ok": false, "error": "signed session is not authenticated"})
|
||||
}
|
||||
record = latest
|
||||
coordinator.mu.Unlock()
|
||||
}
|
||||
|
||||
// A request that loses a race with a successful grant exchange may return a
|
||||
// canonical SESSION_INVALID response for the old secret. Reload and retry
|
||||
// with the newer shared session; never let that stale response erase its
|
||||
@@ -1104,32 +1142,44 @@ func (r *extensionRuntime) ensureSignedSession(config SignedSessionConfig) (*sig
|
||||
_ = r.saveSignedSession(config, record)
|
||||
return nil, fmt.Errorf("signed session expired")
|
||||
}
|
||||
if config.Endpoints.Refresh != "" && time.Until(expiresAt) <= signedSessionRefreshSkew {
|
||||
_ = r.refreshSignedSession(config, record)
|
||||
}
|
||||
}
|
||||
return record, nil
|
||||
}
|
||||
|
||||
func (r *extensionRuntime) refreshSignedSession(config SignedSessionConfig, record *signedSessionRecord) error {
|
||||
func signedSessionRefreshDue(config SignedSessionConfig, record *signedSessionRecord) bool {
|
||||
if config.Endpoints.Refresh == "" || !signedSessionRecordIsUsable(record) {
|
||||
return false
|
||||
}
|
||||
expiresAt, ok := parseSignedSessionTime(record.ExpiresAt)
|
||||
return ok && time.Now().Before(expiresAt) && time.Until(expiresAt) <= signedSessionRefreshSkew
|
||||
}
|
||||
|
||||
func (r *extensionRuntime) fetchSignedSessionRefresh(
|
||||
config SignedSessionConfig,
|
||||
record *signedSessionRecord,
|
||||
) (signedSessionExchangeResponse, error) {
|
||||
var refreshed signedSessionExchangeResponse
|
||||
body, _ := json.Marshal(map[string]string{"install_id": record.InstallID})
|
||||
resp, respBody, _, err := r.doSignedSessionRequest(config, record, http.MethodPost, config.Endpoints.Refresh, body, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
return refreshed, err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("session refresh failed: HTTP %d", resp.StatusCode)
|
||||
return refreshed, fmt.Errorf("session refresh failed: HTTP %d", resp.StatusCode)
|
||||
}
|
||||
var refreshed signedSessionExchangeResponse
|
||||
if err := json.Unmarshal(respBody, &refreshed); err != nil {
|
||||
return err
|
||||
return refreshed, err
|
||||
}
|
||||
return refreshed, nil
|
||||
}
|
||||
|
||||
func applySignedSessionRefresh(record *signedSessionRecord, refreshed signedSessionExchangeResponse) bool {
|
||||
changed := false
|
||||
if refreshed.SessionID != "" {
|
||||
if refreshed.SessionID != "" && refreshed.SessionID != record.SessionID {
|
||||
record.SessionID = refreshed.SessionID
|
||||
changed = true
|
||||
}
|
||||
if refreshed.SessionSecret != "" {
|
||||
if refreshed.SessionSecret != "" && refreshed.SessionSecret != record.SessionSecret {
|
||||
record.SessionSecret = refreshed.SessionSecret
|
||||
changed = true
|
||||
}
|
||||
@@ -1137,7 +1187,85 @@ func (r *extensionRuntime) refreshSignedSession(config SignedSessionConfig, reco
|
||||
record.ExpiresAt = refreshed.ExpiresAt
|
||||
changed = true
|
||||
}
|
||||
if changed {
|
||||
return changed
|
||||
}
|
||||
|
||||
func (r *extensionRuntime) refreshSignedSessionCoalesced(
|
||||
config SignedSessionConfig,
|
||||
coordinator *signedSessionCoordinator,
|
||||
) (*signedSessionRecord, error) {
|
||||
ctx := r.activeOperationContext(context.Background())
|
||||
for {
|
||||
coordinator.mu.Lock()
|
||||
latest, err := r.loadSignedSession(config)
|
||||
if err != nil {
|
||||
coordinator.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
if !signedSessionRefreshDue(config, latest) {
|
||||
coordinator.mu.Unlock()
|
||||
return latest, nil
|
||||
}
|
||||
if coordinator.refreshInFlight {
|
||||
done := coordinator.refreshDone
|
||||
coordinator.mu.Unlock()
|
||||
select {
|
||||
case <-done:
|
||||
coordinator.mu.Lock()
|
||||
sharedErr := coordinator.refreshErr
|
||||
sameOperation := coordinator.refreshDone == done
|
||||
coordinator.mu.Unlock()
|
||||
if sameOperation && sharedErr != nil {
|
||||
return nil, sharedErr
|
||||
}
|
||||
continue
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
coordinator.refreshInFlight = true
|
||||
coordinator.refreshDone = make(chan struct{})
|
||||
coordinator.refreshErr = nil
|
||||
done := coordinator.refreshDone
|
||||
clearGeneration := coordinator.clearGeneration
|
||||
refreshGeneration := *latest
|
||||
coordinator.mu.Unlock()
|
||||
|
||||
refreshed, refreshErr := r.fetchSignedSessionRefresh(config, &refreshGeneration)
|
||||
|
||||
coordinator.mu.Lock()
|
||||
var current *signedSessionRecord
|
||||
finalErr := refreshErr
|
||||
if finalErr == nil && coordinator.clearGeneration != clearGeneration {
|
||||
finalErr = fmt.Errorf("signed-session refresh was superseded by session clear")
|
||||
}
|
||||
if finalErr == nil {
|
||||
var loadErr error
|
||||
current, loadErr = r.loadSignedSession(config)
|
||||
finalErr = loadErr
|
||||
}
|
||||
if finalErr == nil && sameSignedSession(current, &refreshGeneration) {
|
||||
if applySignedSessionRefresh(current, refreshed) {
|
||||
finalErr = r.saveSignedSession(config, current)
|
||||
}
|
||||
}
|
||||
if coordinator.refreshInFlight && coordinator.refreshDone == done {
|
||||
coordinator.refreshInFlight = false
|
||||
coordinator.refreshErr = finalErr
|
||||
close(done)
|
||||
}
|
||||
coordinator.mu.Unlock()
|
||||
return current, finalErr
|
||||
}
|
||||
}
|
||||
|
||||
func (r *extensionRuntime) refreshSignedSession(config SignedSessionConfig, record *signedSessionRecord) error {
|
||||
refreshed, err := r.fetchSignedSessionRefresh(config, record)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if applySignedSessionRefresh(record, refreshed) {
|
||||
return r.saveSignedSession(config, record)
|
||||
}
|
||||
return nil
|
||||
@@ -1192,27 +1320,127 @@ func (r *extensionRuntime) startSignedSessionVerificationLocked(
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("load signed-session bootstrap state: %w", err)
|
||||
}
|
||||
if signedSessionRecordIsUsable(record) &&
|
||||
!coordinator.generationIsBlocked(record) {
|
||||
return "", nil
|
||||
}
|
||||
if coordinator.bootstrapInFlight {
|
||||
done := coordinator.bootstrapDone
|
||||
ctx := r.activeOperationContext(context.Background())
|
||||
coordinator.mu.Unlock()
|
||||
select {
|
||||
case <-done:
|
||||
coordinator.mu.Lock()
|
||||
sharedErr := coordinator.bootstrapErr
|
||||
sameOperation := coordinator.bootstrapDone == done
|
||||
if sameOperation && sharedErr != nil {
|
||||
return "", sharedErr
|
||||
}
|
||||
return r.startSignedSessionVerificationLocked(config, coordinator, reason)
|
||||
case <-ctx.Done():
|
||||
coordinator.mu.Lock()
|
||||
return "", ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
coordinator.bootstrapInFlight = true
|
||||
coordinator.bootstrapDone = make(chan struct{})
|
||||
coordinator.bootstrapErr = nil
|
||||
bootstrapDone := coordinator.bootstrapDone
|
||||
clearGeneration := coordinator.clearGeneration
|
||||
ctx := r.activeOperationContext(context.Background())
|
||||
coordinator.mu.Unlock()
|
||||
bootstrap, bootstrapErr := r.performSignedSessionBootstrap(ctx, config, record, reason)
|
||||
coordinator.mu.Lock()
|
||||
finalErr := bootstrapErr
|
||||
authURL := ""
|
||||
if finalErr == nil && coordinator.clearGeneration != clearGeneration {
|
||||
finalErr = fmt.Errorf("signed-session bootstrap was superseded by session clear")
|
||||
}
|
||||
|
||||
var latest *signedSessionRecord
|
||||
if finalErr == nil {
|
||||
latest, finalErr = r.loadSignedSession(config)
|
||||
}
|
||||
if finalErr == nil && signedSessionRecordIsUsable(latest) &&
|
||||
!sameSignedSession(latest, record) &&
|
||||
!coordinator.generationIsBlocked(latest) {
|
||||
coordinator.clearChallenge()
|
||||
} else if finalErr == nil && bootstrap.SessionID != "" {
|
||||
latest.SessionID = bootstrap.SessionID
|
||||
latest.SessionSecret = bootstrap.SessionSecret
|
||||
latest.ExpiresAt = bootstrap.ExpiresAt
|
||||
if saveErr := r.saveSignedSession(config, latest); saveErr != nil {
|
||||
finalErr = fmt.Errorf("save bootstrapped signed session: %w", saveErr)
|
||||
} else {
|
||||
coordinator.clearBlockedGeneration()
|
||||
coordinator.clearChallenge()
|
||||
}
|
||||
} else if finalErr == nil {
|
||||
request := &PendingAuthRequest{
|
||||
ExtensionID: r.extensionID,
|
||||
AuthURL: bootstrap.AuthURL,
|
||||
CallbackURL: bootstrap.CallbackURL,
|
||||
State: bootstrap.CallbackState,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if registerErr := registerPendingAuthRequest(request); registerErr != nil {
|
||||
finalErr = registerErr
|
||||
} else {
|
||||
coordinator.rememberChallenge(
|
||||
r.extensionID,
|
||||
bootstrap.AuthURL,
|
||||
bootstrap.CallbackURL,
|
||||
bootstrap.CallbackState,
|
||||
)
|
||||
authURL = bootstrap.AuthURL
|
||||
}
|
||||
}
|
||||
if coordinator.bootstrapInFlight && coordinator.bootstrapDone == bootstrapDone {
|
||||
coordinator.bootstrapInFlight = false
|
||||
coordinator.bootstrapErr = finalErr
|
||||
close(bootstrapDone)
|
||||
}
|
||||
return authURL, finalErr
|
||||
}
|
||||
|
||||
type signedSessionBootstrapResult struct {
|
||||
SessionID string
|
||||
SessionSecret string
|
||||
ExpiresAt string
|
||||
AuthURL string
|
||||
CallbackURL string
|
||||
CallbackState string
|
||||
}
|
||||
|
||||
func (r *extensionRuntime) performSignedSessionBootstrap(
|
||||
ctx context.Context,
|
||||
config SignedSessionConfig,
|
||||
record *signedSessionRecord,
|
||||
reason string,
|
||||
) (signedSessionBootstrapResult, error) {
|
||||
var result signedSessionBootstrapResult
|
||||
bootstrapURL, err := signedSessionURL(config, config.Endpoints.Bootstrap)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("build signed-session bootstrap URL: %w", err)
|
||||
return result, fmt.Errorf("build signed-session bootstrap URL: %w", err)
|
||||
}
|
||||
parsed, err := url.Parse(bootstrapURL)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("parse signed-session bootstrap URL: %w", err)
|
||||
return result, fmt.Errorf("parse signed-session bootstrap URL: %w", err)
|
||||
}
|
||||
query := parsed.Query()
|
||||
query.Set("app_version", config.AppVersion)
|
||||
query.Set("install_id", record.InstallID)
|
||||
parsed.RawQuery = query.Encode()
|
||||
if r.httpClient == nil {
|
||||
return "", fmt.Errorf("signed-session bootstrap HTTP client is unavailable")
|
||||
return result, fmt.Errorf("signed-session bootstrap HTTP client is unavailable")
|
||||
}
|
||||
|
||||
var resp *http.Response
|
||||
for attempt := 0; attempt < 2; attempt++ {
|
||||
req, requestErr := http.NewRequest(http.MethodGet, parsed.String(), nil)
|
||||
req, requestErr := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
|
||||
if requestErr != nil {
|
||||
return "", fmt.Errorf("build signed-session bootstrap request: %w", requestErr)
|
||||
return result, fmt.Errorf("build signed-session bootstrap request: %w", requestErr)
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
req.Header.Set("User-Agent", "SpotiFLAC-Mobile/"+config.AppVersion)
|
||||
@@ -1241,7 +1469,7 @@ func (r *extensionRuntime) startSignedSessionVerificationLocked(
|
||||
err,
|
||||
)
|
||||
LogWarn("SignedSession", "Bootstrap failed for extension %s (%s): %v", r.extensionID, reason, bootstrapErr)
|
||||
return "", bootstrapErr
|
||||
return result, bootstrapErr
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
@@ -1257,26 +1485,21 @@ func (r *extensionRuntime) startSignedSessionVerificationLocked(
|
||||
}
|
||||
bootstrapErr := errors.New(message)
|
||||
LogWarn("SignedSession", "Bootstrap failed for extension %s (%s): %v", r.extensionID, reason, bootstrapErr)
|
||||
return "", bootstrapErr
|
||||
return result, bootstrapErr
|
||||
}
|
||||
body, err := readExtensionHTTPResponseBody(resp)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("read signed-session bootstrap response: %w", err)
|
||||
return result, fmt.Errorf("read signed-session bootstrap response: %w", err)
|
||||
}
|
||||
var boot signedSessionExchangeResponse
|
||||
if err := json.Unmarshal(body, &boot); err != nil {
|
||||
return "", fmt.Errorf("decode signed-session bootstrap response: %w", err)
|
||||
return result, fmt.Errorf("decode signed-session bootstrap response: %w", err)
|
||||
}
|
||||
if boot.SessionID != "" && boot.SessionSecret != "" && boot.ExpiresAt != "" {
|
||||
record.SessionID = boot.SessionID
|
||||
record.SessionSecret = boot.SessionSecret
|
||||
record.ExpiresAt = boot.ExpiresAt
|
||||
if err := r.saveSignedSession(config, record); err != nil {
|
||||
return "", fmt.Errorf("save bootstrapped signed session: %w", err)
|
||||
}
|
||||
coordinator.clearBlockedGeneration()
|
||||
coordinator.clearChallenge()
|
||||
return "", nil
|
||||
result.SessionID = boot.SessionID
|
||||
result.SessionSecret = boot.SessionSecret
|
||||
result.ExpiresAt = boot.ExpiresAt
|
||||
return result, nil
|
||||
}
|
||||
authURL := boot.AuthURL
|
||||
if authURL == "" && boot.ChallengeURL != "" {
|
||||
@@ -1286,7 +1509,7 @@ func (r *extensionRuntime) startSignedSessionVerificationLocked(
|
||||
// host-generated nonce so the callback is bound to this one challenge.
|
||||
callbackState, err := newExtensionCallbackState()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("prepare signed-session callback state: %w", err)
|
||||
return result, fmt.Errorf("prepare signed-session callback state: %w", err)
|
||||
}
|
||||
if parsedAuthURL, parseErr := url.Parse(authURL); parseErr == nil {
|
||||
if serverState := strings.TrimSpace(parsedAuthURL.Query().Get("state")); serverState != "" {
|
||||
@@ -1294,32 +1517,24 @@ func (r *extensionRuntime) startSignedSessionVerificationLocked(
|
||||
} else if authURL != "" {
|
||||
authURL, err = setOAuthState(authURL, callbackState)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("prepare signed-session verification URL: %w", err)
|
||||
return result, fmt.Errorf("prepare signed-session verification URL: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
callbackURL, err := setOAuthState(config.CallbackURL, callbackState)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("prepare signed-session callback: %w", err)
|
||||
return result, fmt.Errorf("prepare signed-session callback: %w", err)
|
||||
}
|
||||
if authURL == "" && boot.ChallengeID != "" {
|
||||
authURL = r.buildSignedSessionChallengeURL(config, boot.ChallengeID, callbackState)
|
||||
}
|
||||
if authURL == "" {
|
||||
return "", fmt.Errorf("signed-session bootstrap did not return a session or verification challenge")
|
||||
return result, fmt.Errorf("signed-session bootstrap did not return a session or verification challenge")
|
||||
}
|
||||
request := &PendingAuthRequest{
|
||||
ExtensionID: r.extensionID,
|
||||
AuthURL: authURL,
|
||||
CallbackURL: callbackURL,
|
||||
State: callbackState,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := registerPendingAuthRequest(request); err != nil {
|
||||
return "", err
|
||||
}
|
||||
coordinator.rememberChallenge(r.extensionID, authURL, callbackURL, callbackState)
|
||||
return authURL, nil
|
||||
result.AuthURL = authURL
|
||||
result.CallbackURL = callbackURL
|
||||
result.CallbackState = callbackState
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *extensionRuntime) buildSignedSessionChallengeURL(config SignedSessionConfig, challengeID, callbackState string) string {
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"path/filepath"
|
||||
goruntime "runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -438,6 +439,68 @@ func TestParallelSignedSessionPreflightSharesOneBootstrap(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestParallelSignedSessionPreflightSharesBootstrapFailure(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
requestStarted := make(chan struct{})
|
||||
releaseRequest := make(chan struct{})
|
||||
transport := roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if calls.Add(1) == 1 {
|
||||
close(requestStarted)
|
||||
}
|
||||
<-releaseRequest
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusServiceUnavailable,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{"error":"busy"}`)),
|
||||
Request: req,
|
||||
}, nil
|
||||
})
|
||||
|
||||
root := t.TempDir()
|
||||
config := &SignedSessionConfig{
|
||||
Namespace: "shared-failed-session",
|
||||
BaseURL: "https://auth.example.com",
|
||||
}
|
||||
newRuntime := func(extensionID string) *extensionRuntime {
|
||||
return &extensionRuntime{
|
||||
extensionID: extensionID,
|
||||
manifest: &ExtensionManifest{
|
||||
Name: extensionID,
|
||||
SignedSession: config,
|
||||
},
|
||||
dataDir: filepath.Join(root, extensionID),
|
||||
vm: goja.New(),
|
||||
httpClient: &http.Client{Transport: transport},
|
||||
}
|
||||
}
|
||||
|
||||
const workers = 12
|
||||
errors := make(chan error, workers)
|
||||
for worker := range workers {
|
||||
go func(worker int) {
|
||||
_, err := newRuntime(fmt.Sprintf("failed-provider-%d", worker)).preflightSignedSession()
|
||||
errors <- err
|
||||
}(worker)
|
||||
}
|
||||
select {
|
||||
case <-requestStarted:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("bootstrap request did not start")
|
||||
}
|
||||
// Give every parallel caller time to join the in-flight generation. The
|
||||
// request remains blocked, so no caller can observe a completed operation.
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
close(releaseRequest)
|
||||
for range workers {
|
||||
if err := <-errors; err == nil || !strings.Contains(err.Error(), "HTTP 503") {
|
||||
t.Fatalf("unexpected coalesced bootstrap error: %v", err)
|
||||
}
|
||||
}
|
||||
if got := calls.Load(); got != 1 {
|
||||
t.Fatalf("parallel failed bootstrap calls = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadWithExtensionsStopsAfterFailedSignedSessionPreflight(t *testing.T) {
|
||||
extensionID := "preflight-network-failure"
|
||||
itemID := "preflight-network-item"
|
||||
@@ -1950,6 +2013,160 @@ func TestRefreshSignedSession(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestRefreshSignedSessionCoalescesWithoutHoldingCoordinatorMutex(t *testing.T) {
|
||||
var calls int32
|
||||
requestStarted := make(chan struct{})
|
||||
releaseRequest := make(chan struct{})
|
||||
transport := roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if atomic.AddInt32(&calls, 1) == 1 {
|
||||
close(requestStarted)
|
||||
}
|
||||
<-releaseRequest
|
||||
payload := signedSessionExchangeResponse{
|
||||
SessionSecret: "rotated-secret",
|
||||
ExpiresAt: time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339),
|
||||
}
|
||||
body, _ := json.Marshal(payload)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(string(body))),
|
||||
Request: req,
|
||||
}, nil
|
||||
})
|
||||
runtime := newSignedSessionTestRuntime(t, "refresh-coalesced", transport)
|
||||
config := signedSessionConfigWithDefaults(&SignedSessionConfig{
|
||||
Namespace: "refresh-coalesced",
|
||||
BaseURL: "https://auth.example.com",
|
||||
Endpoints: SignedSessionEndpoints{Refresh: "/session/refresh"},
|
||||
})
|
||||
record, err := runtime.loadSignedSession(config)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
record.SessionID = "session-1"
|
||||
record.SessionSecret = "old-secret"
|
||||
record.ExpiresAt = time.Now().Add(time.Minute).UTC().Format(time.RFC3339)
|
||||
if err := runtime.saveSignedSession(config, record); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
coordinator, err := runtime.signedSessionCoordinator(config)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
const workers = 12
|
||||
results := make(chan *signedSessionRecord, workers)
|
||||
errors := make(chan error, workers)
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(workers)
|
||||
for range workers {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
refreshed, refreshErr := runtime.refreshSignedSessionCoalesced(config, coordinator)
|
||||
results <- refreshed
|
||||
errors <- refreshErr
|
||||
}()
|
||||
}
|
||||
select {
|
||||
case <-requestStarted:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("refresh request did not start")
|
||||
}
|
||||
|
||||
mutexAvailable := make(chan struct{})
|
||||
go func() {
|
||||
coordinator.mu.Lock()
|
||||
coordinator.mu.Unlock()
|
||||
close(mutexAvailable)
|
||||
}()
|
||||
select {
|
||||
case <-mutexAvailable:
|
||||
case <-time.After(250 * time.Millisecond):
|
||||
t.Fatal("coordinator mutex was held during refresh HTTP")
|
||||
}
|
||||
|
||||
close(releaseRequest)
|
||||
wg.Wait()
|
||||
close(results)
|
||||
close(errors)
|
||||
for refreshErr := range errors {
|
||||
if refreshErr != nil {
|
||||
t.Fatalf("coalesced refresh failed: %v", refreshErr)
|
||||
}
|
||||
}
|
||||
for refreshed := range results {
|
||||
if refreshed == nil || refreshed.SessionSecret != "rotated-secret" {
|
||||
t.Fatalf("unexpected refreshed record: %+v", refreshed)
|
||||
}
|
||||
}
|
||||
if got := atomic.LoadInt32(&calls); got != 1 {
|
||||
t.Fatalf("refresh requests = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRefreshSignedSessionSharesFailureWithWaiters(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
requestStarted := make(chan struct{})
|
||||
releaseRequest := make(chan struct{})
|
||||
transport := roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
if calls.Add(1) == 1 {
|
||||
close(requestStarted)
|
||||
}
|
||||
<-releaseRequest
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusServiceUnavailable,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{}`)),
|
||||
Request: req,
|
||||
}, nil
|
||||
})
|
||||
runtime := newSignedSessionTestRuntime(t, "refresh-failure-coalesced", transport)
|
||||
config := signedSessionConfigWithDefaults(&SignedSessionConfig{
|
||||
Namespace: "refresh-failure-coalesced",
|
||||
BaseURL: "https://auth.example.com",
|
||||
Endpoints: SignedSessionEndpoints{Refresh: "/session/refresh"},
|
||||
})
|
||||
record, err := runtime.loadSignedSession(config)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
record.SessionID = "session-1"
|
||||
record.SessionSecret = "old-secret"
|
||||
record.ExpiresAt = time.Now().Add(time.Minute).UTC().Format(time.RFC3339)
|
||||
if err := runtime.saveSignedSession(config, record); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
coordinator, err := runtime.signedSessionCoordinator(config)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
const workers = 12
|
||||
errors := make(chan error, workers)
|
||||
for range workers {
|
||||
go func() {
|
||||
_, refreshErr := runtime.refreshSignedSessionCoalesced(config, coordinator)
|
||||
errors <- refreshErr
|
||||
}()
|
||||
}
|
||||
select {
|
||||
case <-requestStarted:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("refresh request did not start")
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
close(releaseRequest)
|
||||
for range workers {
|
||||
if refreshErr := <-errors; refreshErr == nil || !strings.Contains(refreshErr.Error(), "HTTP 503") {
|
||||
t.Fatalf("unexpected coalesced refresh error: %v", refreshErr)
|
||||
}
|
||||
}
|
||||
if got := calls.Load(); got != 1 {
|
||||
t.Fatalf("parallel failed refresh calls = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignedSessionCompleteGrant(t *testing.T) {
|
||||
t.Run("uses the grant argument when provided", func(t *testing.T) {
|
||||
transport := roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
|
||||
Reference in New Issue
Block a user