fix: recover after game session loss

This commit is contained in:
2026-10-03 00:44:32 +08:00
parent e1c6ce082c
commit ab1aee9d89
6 changed files with 200 additions and 7 deletions
+8 -3
View File
@@ -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")
+12 -2
View File
@@ -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)
+14
View File
@@ -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
+33
View File
@@ -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
View File
@@ -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
View File
@@ -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"
}
}