fix: recover after game session loss
This commit is contained in:
@@ -28,6 +28,8 @@ const (
|
||||
maxCookieHeaderLen = 8 << 10
|
||||
)
|
||||
|
||||
var errSessionRequired = errors.New("session login required")
|
||||
|
||||
type LoginService interface {
|
||||
Login(request, sessionKey []byte) ([]byte, error)
|
||||
}
|
||||
@@ -372,13 +374,16 @@ func (s *Server) dispatch(path string, request []byte) (int, []byte, error) {
|
||||
func (s *Server) authorize(cookie string) (*gameSession, error) {
|
||||
token, err := parseSessionCookie(cookie)
|
||||
if err != nil {
|
||||
if errors.Is(err, errSessionRequired) {
|
||||
return nil, fmt.Errorf("%w: %v", transport.ErrGameSessionExpired, err)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
now := s.now()
|
||||
s.pruneSessions(now)
|
||||
game, ok := s.sessions[sessionTokenKey(token)]
|
||||
if !ok {
|
||||
return nil, errors.New("invalid game session cookie")
|
||||
return nil, fmt.Errorf("%w: cookie no longer names a live session", transport.ErrGameSessionExpired)
|
||||
}
|
||||
game.lastUsed = now
|
||||
return game, nil
|
||||
@@ -386,7 +391,7 @@ func (s *Server) authorize(cookie string) (*gameSession, error) {
|
||||
|
||||
func parseSessionCookie(cookie string) (string, error) {
|
||||
if cookie == "" {
|
||||
return "", errors.New("session login required")
|
||||
return "", errSessionRequired
|
||||
}
|
||||
if len(cookie) > maxCookieHeaderLen {
|
||||
return "", errors.New("game session cookie header is too large")
|
||||
@@ -405,7 +410,7 @@ func parseSessionCookie(cookie string) (string, error) {
|
||||
token = candidate
|
||||
}
|
||||
if !seen {
|
||||
return "", errors.New("session login required")
|
||||
return "", errSessionRequired
|
||||
}
|
||||
if len(token) != 50 || token[48:] != "|1" {
|
||||
return "", errors.New("invalid game session cookie")
|
||||
|
||||
@@ -297,6 +297,8 @@ func TestGameSessionExpiresAndClearsKey(t *testing.T) {
|
||||
}
|
||||
if _, err := server.DispatchRaw("/EmptyInfo", []byte(body), "s="+reply.Cookie); err == nil {
|
||||
t.Fatal("expired game session was accepted")
|
||||
} else if !errors.Is(err, transport.ErrGameSessionExpired) {
|
||||
t.Fatalf("expired game session error=%v, want ErrGameSessionExpired", err)
|
||||
}
|
||||
if len(server.sessions) != 0 {
|
||||
t.Fatalf("expired game session remains in map: %d", len(server.sessions))
|
||||
@@ -395,8 +397,16 @@ func TestNativeLoginAndBatch(t *testing.T) {
|
||||
|
||||
func TestSessionRejectsMissingCookieAndUnknownPath(t *testing.T) {
|
||||
server, _ := NewServer(fakeLogin{}, fakeDomain{})
|
||||
if _, err := server.DispatchRaw("/EmptyInfo", nil, ""); err == nil {
|
||||
t.Fatal("authenticated endpoint accepted missing cookie")
|
||||
if _, err := server.DispatchRaw("/EmptyInfo", nil, ""); !errors.Is(err, transport.ErrGameSessionExpired) {
|
||||
t.Fatalf("missing cookie error=%v, want ErrGameSessionExpired", err)
|
||||
}
|
||||
unknown := strings.Repeat("a", 48) + "|1"
|
||||
if _, err := server.DispatchRaw("/BatchRequest", nil, "s="+unknown); !errors.Is(err, transport.ErrGameSessionExpired) {
|
||||
t.Fatalf("unknown cookie error=%v, want ErrGameSessionExpired", err)
|
||||
}
|
||||
if _, err := server.DispatchRaw("/EmptyInfo", nil, "s=malformed"); err == nil ||
|
||||
errors.Is(err, transport.ErrGameSessionExpired) {
|
||||
t.Fatalf("malformed cookie error=%v, want ordinary rejection", err)
|
||||
}
|
||||
reply := login(t, server)
|
||||
request := wire.AppendVarint(nil, 1, 99)
|
||||
|
||||
@@ -59,6 +59,13 @@ type Bootstrap struct {
|
||||
|
||||
var ErrNotImplemented = errors.New("packet not implemented")
|
||||
|
||||
// ErrGameSessionExpired tells the HTTP adapter that a syntactically valid
|
||||
// game-session cookie no longer names a live session. A dedicated status and
|
||||
// header let the client distinguish this condition from a transient network
|
||||
// outage for both ordinary and batch requests; the server cannot encode a
|
||||
// response with the per-session key after that key has been lost.
|
||||
var ErrGameSessionExpired = errors.New("game session expired")
|
||||
|
||||
func (b Bootstrap) Dispatch(path string, request []byte) (Reply, error) {
|
||||
switch path {
|
||||
case "/MaintenanceInfo":
|
||||
@@ -207,6 +214,13 @@ func (h HTTP) game(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
reply, err := h.Raw.DispatchRaw(path, body, r.Header.Get("Cookie"))
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrGameSessionExpired) {
|
||||
h.logger().Info("game session expired", "path", path, "duration_ms", elapsedMilliseconds(started))
|
||||
w.Header().Set("X-BD2-Session-Expired", "1")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
http.Error(w, "game session expired", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
h.logger().Warn("session packet rejected", "path", path, "duration_ms", elapsedMilliseconds(started), "error", err)
|
||||
http.Error(w, "session packet rejected", http.StatusBadRequest)
|
||||
return
|
||||
|
||||
@@ -3,6 +3,7 @@ package transport
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
@@ -68,6 +69,38 @@ func (cookieRawDispatcher) DispatchRaw(string, []byte, string) (RawReply, error)
|
||||
return RawReply{Body: []byte(`{}`), Cookie: "0123456789abcdef0123456789abcdef0123456789abcdef|1"}, nil
|
||||
}
|
||||
|
||||
type failedRawDispatcher struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (d failedRawDispatcher) DispatchRaw(string, []byte, string) (RawReply, error) {
|
||||
return RawReply{}, d.err
|
||||
}
|
||||
|
||||
func TestExpiredGameSessionUsesDedicatedHTTPMarker(t *testing.T) {
|
||||
h := HTTP{Raw: failedRawDispatcher{err: ErrGameSessionExpired}}.Handler()
|
||||
request := httptest.NewRequest(http.MethodPut, "/game/BatchRequest", strings.NewReader("encrypted"))
|
||||
request.Header.Set("Cookie", "s=0123456789abcdef0123456789abcdef0123456789abcdef|1")
|
||||
response := httptest.NewRecorder()
|
||||
h.ServeHTTP(response, request)
|
||||
|
||||
if response.Code != http.StatusUnauthorized || response.Header().Get("X-BD2-Session-Expired") != "1" ||
|
||||
response.Header().Get("Cache-Control") != "no-store" {
|
||||
t.Fatalf("status=%d marker=%q cache=%q body=%q", response.Code,
|
||||
response.Header().Get("X-BD2-Session-Expired"), response.Header().Get("Cache-Control"), response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDomainFailureDoesNotUseExpiredSessionMarker(t *testing.T) {
|
||||
h := HTTP{Raw: failedRawDispatcher{err: errors.New("mail seed is invalid")}}.Handler()
|
||||
response := httptest.NewRecorder()
|
||||
h.ServeHTTP(response, httptest.NewRequest(http.MethodPut, "/game/MailInfo", strings.NewReader("encrypted")))
|
||||
if response.Code != http.StatusBadRequest || response.Header().Get("X-BD2-Session-Expired") != "" {
|
||||
t.Fatalf("status=%d marker=%q body=%q", response.Code,
|
||||
response.Header().Get("X-BD2-Session-Expired"), response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthGameSessionCookieIsHostOnlySecureAndGameScoped(t *testing.T) {
|
||||
h := HTTP{
|
||||
Raw: cookieRawDispatcher{},
|
||||
|
||||
+132
-1
@@ -3,6 +3,7 @@ using System.Collections;
|
||||
using System.IO;
|
||||
using System.Reflection;
|
||||
using System.Text;
|
||||
using System.Threading;
|
||||
using BepInEx;
|
||||
using BepInEx.Logging;
|
||||
using HarmonyLib;
|
||||
@@ -39,6 +40,7 @@ public sealed class Plugin : BaseUnityPlugin
|
||||
private static MethodInfo OpenPCLoginPopup;
|
||||
private static MemoryAccessTokenStore AccessTokens;
|
||||
private static IRefreshCredentialStore RefreshCredentials;
|
||||
private static int SessionRecoveryInProgress;
|
||||
|
||||
private void Awake()
|
||||
{
|
||||
@@ -71,8 +73,24 @@ public sealed class Plugin : BaseUnityPlugin
|
||||
OpenPCLoginPopup = FindOpenPCLoginPopup();
|
||||
MethodInfo accessTokenGetter = FindAccessTokenGetter();
|
||||
MethodInfo clearPCLocalData = FindClearPCLocalData();
|
||||
MethodInfo sendWebRequest = typeof(UnityWebRequest).GetMethod(
|
||||
nameof(UnityWebRequest.SendWebRequest),
|
||||
BindingFlags.Instance | BindingFlags.Public,
|
||||
null,
|
||||
Type.EmptyTypes,
|
||||
null);
|
||||
Type networkManager = FindType("BDNetwork.NetworkManager");
|
||||
MethodInfo clientNetworkError = networkManager?.GetMethod(
|
||||
"ClientNetworkError",
|
||||
BindingFlags.Instance | BindingFlags.Public);
|
||||
MethodInfo exponentialBackOff = networkManager?.GetMethod(
|
||||
"ὧὥὡὠὮὦὯὥὭὣὩ",
|
||||
BindingFlags.Instance | BindingFlags.NonPublic) ?? networkManager?.GetMethod(
|
||||
"ExponetialBackOff",
|
||||
BindingFlags.Instance | BindingFlags.NonPublic);
|
||||
if (SendMaintenance == null || SetIntroState == null || OpenPCLoginPopup == null ||
|
||||
accessTokenGetter == null || clearPCLocalData == null)
|
||||
accessTokenGetter == null || clearPCLocalData == null || sendWebRequest == null ||
|
||||
clientNetworkError == null || exponentialBackOff == null)
|
||||
{
|
||||
throw new MissingMethodException("IntroUI authentication transition methods were not found (client version mismatch)");
|
||||
}
|
||||
@@ -88,6 +106,15 @@ public sealed class Plugin : BaseUnityPlugin
|
||||
harmony.Patch(
|
||||
clearPCLocalData,
|
||||
postfix: new HarmonyMethod(typeof(Plugin), nameof(ClearPCLocalDataPostfix)));
|
||||
harmony.Patch(
|
||||
sendWebRequest,
|
||||
postfix: new HarmonyMethod(typeof(Plugin), nameof(SendWebRequestPostfix)));
|
||||
harmony.Patch(
|
||||
clientNetworkError,
|
||||
prefix: new HarmonyMethod(typeof(Plugin), nameof(SuppressNetworkErrorDuringRecovery)));
|
||||
harmony.Patch(
|
||||
exponentialBackOff,
|
||||
prefix: new HarmonyMethod(typeof(Plugin), nameof(SuppressNetworkErrorDuringRecovery)));
|
||||
Logger.LogInfo("Server-authoritative Discord and Google login UI patch installed");
|
||||
}
|
||||
catch (Exception ex)
|
||||
@@ -101,6 +128,103 @@ public sealed class Plugin : BaseUnityPlugin
|
||||
ConfigureLoginPanel(__instance);
|
||||
}
|
||||
|
||||
private static void SendWebRequestPostfix(
|
||||
UnityWebRequest __instance,
|
||||
UnityWebRequestAsyncOperation __result)
|
||||
{
|
||||
if (__instance == null || __result == null)
|
||||
{
|
||||
return;
|
||||
}
|
||||
__result.completed += delegate
|
||||
{
|
||||
InspectCompletedGameRequest(__instance);
|
||||
};
|
||||
}
|
||||
|
||||
private static void InspectCompletedGameRequest(UnityWebRequest request)
|
||||
{
|
||||
try
|
||||
{
|
||||
if (!IsCurrentGameRequest(request, out Uri requestUri))
|
||||
{
|
||||
return;
|
||||
}
|
||||
if (requestUri.AbsolutePath.Equals("/game/LoginUser", StringComparison.Ordinal) &&
|
||||
request.responseCode >= 200 && request.responseCode < 300)
|
||||
{
|
||||
Interlocked.Exchange(ref SessionRecoveryInProgress, 0);
|
||||
return;
|
||||
}
|
||||
if (request.responseCode != 401 ||
|
||||
!string.Equals(request.GetResponseHeader("X-BD2-Session-Expired"), "1", StringComparison.Ordinal))
|
||||
{
|
||||
return;
|
||||
}
|
||||
RecoverExpiredGameSession();
|
||||
}
|
||||
catch (Exception ex)
|
||||
{
|
||||
Log?.LogError("Could not inspect the completed game request: " + ex);
|
||||
}
|
||||
}
|
||||
|
||||
private static bool IsCurrentGameRequest(UnityWebRequest request, out Uri requestUri)
|
||||
{
|
||||
requestUri = null;
|
||||
if (ServerRoot == null || request == null ||
|
||||
!Uri.TryCreate(request.url, UriKind.Absolute, out Uri parsed) ||
|
||||
!SameOrigin(ServerRoot, parsed) ||
|
||||
!parsed.AbsolutePath.StartsWith("/game/", StringComparison.Ordinal))
|
||||
{
|
||||
return false;
|
||||
}
|
||||
requestUri = parsed;
|
||||
return true;
|
||||
}
|
||||
|
||||
private static void RecoverExpiredGameSession()
|
||||
{
|
||||
if (Interlocked.CompareExchange(ref SessionRecoveryInProgress, 1, 0) != 0)
|
||||
{
|
||||
return;
|
||||
}
|
||||
try
|
||||
{
|
||||
object network = FindSGSingleton("Net");
|
||||
object app = FindSGSingleton("App");
|
||||
MethodInfo refresh = network?.GetType().GetMethod(
|
||||
"Refresh",
|
||||
BindingFlags.Instance | BindingFlags.Public,
|
||||
null,
|
||||
Type.EmptyTypes,
|
||||
null);
|
||||
MethodInfo restart = app?.GetType().GetMethod(
|
||||
"AppReStart",
|
||||
BindingFlags.Instance | BindingFlags.Public,
|
||||
null,
|
||||
Type.EmptyTypes,
|
||||
null);
|
||||
if (refresh == null || restart == null)
|
||||
{
|
||||
throw new MissingMethodException("client game-session recovery methods were not found");
|
||||
}
|
||||
Log?.LogWarning("Game session expired; returning to login and creating a new session");
|
||||
refresh.Invoke(network, null);
|
||||
restart.Invoke(app, null);
|
||||
}
|
||||
catch
|
||||
{
|
||||
Interlocked.Exchange(ref SessionRecoveryInProgress, 0);
|
||||
throw;
|
||||
}
|
||||
}
|
||||
|
||||
private static bool SuppressNetworkErrorDuringRecovery()
|
||||
{
|
||||
return Volatile.Read(ref SessionRecoveryInProgress) == 0;
|
||||
}
|
||||
|
||||
private static bool AccessTokenPrefix(ref string __result)
|
||||
{
|
||||
if (Authentication != null && Authentication.mode == "oauth")
|
||||
@@ -973,6 +1097,13 @@ public sealed class Plugin : BaseUnityPlugin
|
||||
return null;
|
||||
}
|
||||
|
||||
private static object FindSGSingleton(string propertyName)
|
||||
{
|
||||
Type sg = FindType("SG");
|
||||
PropertyInfo property = sg?.GetProperty(propertyName, BindingFlags.Static | BindingFlags.Public);
|
||||
return property?.GetValue(null);
|
||||
}
|
||||
|
||||
private static bool ProviderEnabled(string provider)
|
||||
{
|
||||
if (Authentication?.providers == null)
|
||||
|
||||
+1
-1
@@ -6,6 +6,6 @@
|
||||
"plugins": {
|
||||
"local_identity": "0.6.0",
|
||||
"capture_environment": "0.2.0",
|
||||
"login_ui": "0.1.1"
|
||||
"login_ui": "0.1.2"
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user