Skip to content
Open
Show file tree
Hide file tree
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
78 changes: 63 additions & 15 deletions xal/nsal/transport.go
Original file line number Diff line number Diff line change
Expand Up @@ -55,30 +55,78 @@ func (t *Transport) RoundTrip(req *http.Request) (*http.Response, error) {
return t.baseTransport().RoundTrip(req)
}

token, policy, err := t.TokenAndSignature(ctx, req.URL)
if err != nil {
return nil, fmt.Errorf("request XSTS token and signature: %w", err)
var data []byte
if req.Body != nil {
var err error
data, err = io.ReadAll(req.Body)
if err != nil {
return nil, fmt.Errorf("read request body: %w", err)
}
}

req2 := req.Clone(ctx)
token.SetAuthHeader(req2)
return t.roundTripAuthenticated(req, exclusion, data)
}

// roundTripAuthenticated signs and sends req, retrying once when Xbox reports
// that the XSTS token expired before its advertised lifetime.
func (t *Transport) roundTripAuthenticated(req *http.Request, exclusion headerExclusionSet, data []byte) (*http.Response, error) {
ctx := req.Context()
for attempt := 0; ; attempt++ {
token, policy, err := t.TokenAndSignature(ctx, req.URL)
if err != nil {
return nil, fmt.Errorf("request XSTS token and signature: %w", err)
}

if req2.Header.Get("Signature") == "" && !exclusion.signature() {
var data []byte
req2 := req.Clone(ctx)
if req.Body != nil {
signingBuffer := &bytes.Buffer{}
if _, err := signingBuffer.ReadFrom(req.Body); err != nil {
signingBuffer.Reset()
return nil, fmt.Errorf("clone request body: %w", err)
req2.Body = io.NopCloser(bytes.NewReader(data))
req2.GetBody = func() (io.ReadCloser, error) {
return io.NopCloser(bytes.NewReader(data)), nil
}
data, req2.Body = signingBuffer.Bytes(), io.NopCloser(signingBuffer)
}
if err := policy.Sign(req2, data, t.Resolver.src.ProofKey(), timestamp.Now()); err != nil {
return nil, fmt.Errorf("sign request: %w", err)
token.SetAuthHeader(req2)

if req2.Header.Get("Signature") == "" && !exclusion.signature() {
if err := policy.Sign(req2, data, t.Resolver.src.ProofKey(), timestamp.Now()); err != nil {
return nil, fmt.Errorf("sign request: %w", err)
}
}

resp, err := t.baseTransport().RoundTrip(req2)
if err != nil {
return nil, err
}
invalidator, ok := t.Resolver.src.(xstsTokenInvalidator)
if attempt > 0 || !ok || !tokenExpired(resp) {
return resp, nil
}
if resp.Body != nil {
_ = resp.Body.Close()
}
invalidator.InvalidateXSTSToken(token)
}
}

// tokenExpired reports whether Xbox explicitly rejected an expired token.
func tokenExpired(resp *http.Response) bool {
if resp == nil || resp.StatusCode != http.StatusUnauthorized {
return false
}
for _, value := range resp.Header.Values("WWW-Authenticate") {
for part := range strings.SplitSeq(value, ",") {
part = strings.TrimSpace(strings.ToLower(part))
part = strings.TrimSpace(strings.TrimPrefix(part, "token "))
if part == "error='token_expired'" {
return true
}
}
}
return false
}

return t.baseTransport().RoundTrip(req2)
// xstsTokenInvalidator lets token sources discard a token rejected upstream.
type xstsTokenInvalidator interface {
InvalidateXSTSToken(*xsts.Token)
}

// TokenAndSignature resolves an XSTS token and signature policy for the given URL.
Expand Down
105 changes: 105 additions & 0 deletions xal/nsal/transport_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,98 @@ func TestTransportRoundTripSignsRequest(t *testing.T) {
}
}

func TestTransportRoundTripWithNilResolverReturnsError(t *testing.T) {
transport := &Transport{}
req, err := http.NewRequest(http.MethodGet, "https://multiplayer.minecraft.net/authentication", nil)
if err != nil {
t.Fatalf("new request: %v", err)
}

_, err = transport.RoundTrip(req)
if err == nil || !strings.Contains(err.Error(), "xal/nsal: nil Resolver") {
t.Fatalf("RoundTrip error = %v, want nil Resolver error", err)
}
}

func TestTransportRoundTripRefreshesExpiredXSTSToken(t *testing.T) {
key := mustGenerateKey(t)
stale := authorizationToken("stale")
fresh := authorizationToken("fresh")
src := &refreshingTransportTokenSource{
transportTokenSource: transportTokenSource{token: stale, proofKey: key},
fresh: fresh,
}
firstResponseBody := &trackingBody{ReadCloser: http.NoBody}
var requests int
transport := &Transport{
Base: roundTripFunc(func(req *http.Request) (*http.Response, error) {
requests++
body, err := io.ReadAll(req.Body)
if err != nil || string(body) != "payload" {
t.Fatalf("request %d body = %q, err = %v", requests, body, err)
}
if requests == 1 {
return &http.Response{
StatusCode: http.StatusUnauthorized,
Header: http.Header{"Www-Authenticate": {"Token error='token_expired'"}},
Body: firstResponseBody,
}, nil
}
if requests != 2 {
t.Fatalf("unexpected request %d", requests)
}
if req.GetBody == nil {
t.Fatal("retried request is not replayable")
}
if got := req.Header.Get("Authorization"); got != "XBL3.0 x=uhs;fresh" {
t.Fatalf("Authorization = %q, want fresh token", got)
}
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
}),
Resolver: testResolver(src),
}

req, err := http.NewRequest(http.MethodPut, "https://multiplayer.minecraft.net/authentication", strings.NewReader("payload"))
if err != nil {
t.Fatalf("new request: %v", err)
}
req.GetBody = nil
resp, err := transport.RoundTrip(req)
if err != nil {
t.Fatalf("RoundTrip: %v", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusOK)
}
if src.invalidated != stale {
t.Fatal("invalidated token was not the rejected token")
}
if !firstResponseBody.closed {
t.Fatal("first response body was not closed before retry")
}
}

func TestTokenExpired(t *testing.T) {
for name, tc := range map[string]struct {
headers []string
want bool
}{
"first parameter": {[]string{"Token error='token_expired'"}, true},
"no comma space": {[]string{"Token realm='xboxlive.com',error='token_expired'"}, true},
"later header": {[]string{"Token error='token_required'", "Token error='token_expired'"}, true},
"other error": {[]string{"Token error='token_required'"}, false},
"provider error": {[]string{"Token provider_error='token_expired'"}, false},
} {
t.Run(name, func(t *testing.T) {
resp := &http.Response{StatusCode: http.StatusUnauthorized, Header: http.Header{"Www-Authenticate": tc.headers}}
if got := tokenExpired(resp); got != tc.want {
t.Fatalf("tokenExpired = %t, want %t", got, tc.want)
}
})
}
}

func TestTransportRoundTripUsesExistingAuthorization(t *testing.T) {
src := &transportTokenSource{token: authorizationToken("unexpected")}
transport := &Transport{
Expand Down Expand Up @@ -193,6 +285,19 @@ type transportTokenSource struct {
err error
}

type refreshingTransportTokenSource struct {
transportTokenSource
fresh *xsts.Token
invalidated *xsts.Token
}

func (src *refreshingTransportTokenSource) InvalidateXSTSToken(token *xsts.Token) {
src.invalidated = token
if src.token == token {
src.token = src.fresh
}
}

func (src *transportTokenSource) XSTSToken(_ context.Context, relyingParty string) (*xsts.Token, error) {
src.called = true
src.calls++
Expand Down
76 changes: 61 additions & 15 deletions xal/sisu/session.go
Original file line number Diff line number Diff line change
Expand Up @@ -147,6 +147,9 @@ type Session struct {
xsts map[string]*xsts.Token
// xstsMu guards xsts tokens from concurrent read-write access.
xstsMu sync.Mutex
// xstsGeneration changes whenever a relying party rejects a token. Token
// acquisitions use it to avoid restoring a token fetched before invalidation.
xstsGeneration uint64

// resp is the last known response for SISU authorization request.
// It contains title, user, and an XSTS token that relies on the
Expand Down Expand Up @@ -242,30 +245,73 @@ func (s *Session) Snapshot() *Snapshot {
//
// XSTS tokens are cached per relying party and reused until expiration.
func (s *Session) XSTSToken(ctx context.Context, relyingParty string) (*xsts.Token, error) {
s.xstsMu.Lock()
token, ok := s.xsts[relyingParty]
if ok && token.Valid() {
// Re-use the cached XSTS token as possible.
return s.xstsToken(ctx, relyingParty, s.requestXSTS)
}

// xstsToken avoids caching an acquisition that overlaps an invalidation.
func (s *Session) xstsToken(ctx context.Context, relyingParty string, request func(context.Context, string) (*xsts.Token, error)) (*xsts.Token, error) {
for {
s.xstsMu.Lock()
token, ok := s.xsts[relyingParty]
if ok && token.Valid() {
// Re-use the cached XSTS token as possible.
s.xstsMu.Unlock()
return token, nil
}
generation := s.xstsGeneration
s.xstsMu.Unlock()

token, err := request(ctx, relyingParty)
if err != nil {
return nil, err
}
if !token.Valid() {
return nil, errors.New("xal/sisu: invalid XSTS token data")
}

s.xstsMu.Lock()
if cached, ok := s.xsts[relyingParty]; ok && cached.Valid() {
s.xstsMu.Unlock()
return cached, nil
}
if s.xstsGeneration != generation {
s.xstsMu.Unlock()
continue
}
s.xsts[relyingParty] = token
s.xstsMu.Unlock()
return token, nil
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.
s.xstsMu.Unlock()
}

token, err := s.requestXSTS(ctx, relyingParty)
if err != nil {
return nil, err
}
if !token.Valid() {
return nil, errors.New("xal/sisu: invalid XSTS token data")
// InvalidateXSTSToken removes a rejected XSTS token from the session caches if
// it has not already been replaced. Xbox services may invalidate a token before
// its signed NotAfter time.
func (s *Session) InvalidateXSTSToken(rejected *xsts.Token) {
if rejected == nil {
return
}

// Keep both caches behind one invalidation boundary. Otherwise a default-RP
// acquisition can observe the new generation while still reusing s.resp.
s.xstsMu.Lock()
s.respMu.Lock()
defer s.respMu.Unlock()
defer s.xstsMu.Unlock()
if cached, ok := s.xsts[relyingParty]; ok && cached.Valid() {
return cached, nil
s.xstsGeneration++

for relyingParty, token := range s.xsts {
if sameXSTSToken(token, rejected) {
delete(s.xsts, relyingParty)
}
}
if s.resp != nil && sameXSTSToken(s.resp.AuthorizationToken, rejected) {
s.resp = nil
}
s.xsts[relyingParty] = token
return token, nil
}

func sameXSTSToken(a, b *xsts.Token) bool {
return a != nil && b != nil && a.Token == b.Token
}

// requestXSTS obtains a new XSTS token for the relying party.
Expand Down
72 changes: 72 additions & 0 deletions xal/sisu/session_unit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,78 @@ func TestXSTSTokenDoesNotHoldCacheLockDuringTokenRequest(t *testing.T) {
}
}

func TestSessionInvalidatesOnlyRejectedXSTSToken(t *testing.T) {
rejected := &xsts.Token{Token: "rejected"}
replacement := &xsts.Token{Token: "replacement"}
session := (Config{}).New(staticMSATokenSource{}, nil)
session.xsts[defaultRelyingParty] = rejected
session.resp = &authorizationResponse{AuthorizationToken: rejected}

session.InvalidateXSTSToken(rejected)
if _, ok := session.xsts[defaultRelyingParty]; ok {
t.Fatal("rejected token remained cached")
}
if session.resp != nil {
t.Fatal("response containing rejected token remained cached")
}

session.xsts[defaultRelyingParty] = replacement
session.resp = &authorizationResponse{AuthorizationToken: replacement}
session.InvalidateXSTSToken(rejected)
if session.xsts[defaultRelyingParty] != replacement || session.resp == nil || session.resp.AuthorizationToken != replacement {
t.Fatal("replacement token was removed")
}
}

func TestSessionDoesNotCacheTokenFetchedDuringInvalidation(t *testing.T) {
rejected := validXSTSToken("rejected")
replacement := validXSTSToken("replacement")
session := (Config{}).New(staticMSATokenSource{}, nil)
fetchStarted, allowFetch := make(chan struct{}), make(chan struct{})
var calls int
fetch := func(context.Context, string) (*xsts.Token, error) {
calls++
if calls == 1 {
close(fetchStarted)
<-allowFetch
return rejected, nil
}
return replacement, nil
}

type result struct {
token *xsts.Token
err error
}
done := make(chan result, 1)
go func() {
token, err := session.xstsToken(context.Background(), defaultRelyingParty, fetch)
done <- result{token: token, err: err}
}()

<-fetchStarted
session.InvalidateXSTSToken(rejected)
close(allowFetch)
got := <-done
if got.err != nil {
t.Fatalf("XSTSToken: %v", got.err)
}
if got.token != replacement {
t.Fatal("XSTSToken returned the token fetched before invalidation")
}
if cached := session.xsts[defaultRelyingParty]; cached != replacement {
t.Fatal("token fetched before invalidation was restored to the cache")
}
}

func validXSTSToken(value string) *xsts.Token {
return &xsts.Token{
Token: value,
NotAfter: time.Now().Add(time.Hour),
DisplayClaims: xsts.DisplayClaims{UserInfo: []xsts.UserInfo{{}}},
}
}

type staticMSATokenSource struct{}

func (staticMSATokenSource) Token() (*oauth2.Token, error) {
Expand Down