feat(all): split client tooling and add OAuth server login

This commit is contained in:
2026-09-30 17:32:14 +08:00
parent 7cb5e36c2a
commit 50e2572385
230 changed files with 9875 additions and 853 deletions
+128
View File
@@ -0,0 +1,128 @@
//go:build !release
package main
import (
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
clientconfig "bd2server/internal/client/config"
clientlayout "bd2server/internal/client/layout"
)
func developmentRunOptions(args []string) ([]string, clientRunOptions, error) {
if len(args) == 0 || args[0] != "--dev" {
return args, clientRunOptions{}, nil
}
if len(args) < 2 || args[1] != "run" {
return nil, clientRunOptions{}, errors.New("usage: bd2client --dev run [client options]")
}
root, err := findClientDevelopmentRoot()
if err != nil {
return nil, clientRunOptions{}, err
}
clientArgs := append([]string(nil), args[2:]...)
gameDir, err := clientDevelopmentGameDirectory(clientArgs)
if err != nil {
return nil, clientRunOptions{}, err
}
if gameDir == "" {
preferences, preferenceErr := clientconfig.LoadPreferences()
if preferenceErr == nil {
gameDir = preferences.GameDirectory
}
}
if gameDir != "" {
if err := buildDevelopmentPlugins(root, gameDir); err != nil {
return nil, clientRunOptions{}, err
}
}
return clientArgs, clientRunOptions{
versionConfigPath: filepath.Join(root, "versions.json"),
logExecutablePath: filepath.Join(root, "data", "bd2client-dev"),
localIdentityPlugin: filepath.Join(root, "plugins", "LocalIdentity", "bin", "Release", "netstandard2.1", "BD2LocalIdentity.dll"),
loginUIPlugin: filepath.Join(root, "plugins", "LoginUI", "bin", "Release", "netstandard2.1", "BD2LoginUI.dll"),
}, nil
}
func findClientDevelopmentRoot() (string, error) {
working, err := os.Getwd()
if err != nil {
return "", fmt.Errorf("resolve development working directory: %w", err)
}
for directory := filepath.Clean(working); ; directory = filepath.Dir(directory) {
if clientDevelopmentFile(filepath.Join(directory, "versions.json")) &&
clientDevelopmentFile(filepath.Join(directory, "go", "go.mod")) &&
clientDevelopmentFile(filepath.Join(directory, "plugins", "LocalIdentity", "LocalIdentity.csproj")) &&
clientDevelopmentFile(filepath.Join(directory, "plugins", "LoginUI", "LoginUI.csproj")) {
return directory, nil
}
parent := filepath.Dir(directory)
if parent == directory {
break
}
}
return "", errors.New("development repository root not found; run from the bd2 repository")
}
func clientDevelopmentFile(path string) bool {
info, err := os.Stat(path)
return err == nil && info.Mode().IsRegular()
}
func clientDevelopmentGameDirectory(args []string) (string, error) {
for index := 0; index < len(args); index++ {
arg := args[index]
if arg == "--game-dir" {
if index+1 >= len(args) || strings.TrimSpace(args[index+1]) == "" {
return "", errors.New("--game-dir requires a directory")
}
return filepath.Clean(args[index+1]), nil
}
if strings.HasPrefix(arg, "--game-dir=") {
value := strings.TrimSpace(strings.TrimPrefix(arg, "--game-dir="))
if value == "" {
return "", errors.New("--game-dir requires a directory")
}
return filepath.Clean(value), nil
}
}
return "", nil
}
func buildDevelopmentPlugins(root, gameDir string) error {
installation, err := clientlayout.Resolve(gameDir)
if err != nil {
return nil
}
if !clientDevelopmentFile(filepath.Join(installation.BepInEx, "core", "BepInEx.dll")) {
return nil
}
projects := []string{
filepath.Join(root, "plugins", "LocalIdentity", "LocalIdentity.csproj"),
filepath.Join(root, "plugins", "LoginUI", "LoginUI.csproj"),
}
for _, project := range projects {
command := exec.Command(
"dotnet", "build", project, "-c", "Release",
"-p:GameDir="+filepath.Clean(gameDir),
"-p:BD2ManagedDir="+filepath.Join(installation.Data, "Managed"),
"-p:BD2BepInExDir="+installation.BepInEx,
"--nologo",
)
command.Dir = root
output, err := command.CombinedOutput()
if err != nil {
message := strings.TrimSpace(string(output))
if message == "" {
message = err.Error()
}
return fmt.Errorf("build development plugin %s: %s", filepath.Base(project), message)
}
}
return nil
}
+7
View File
@@ -0,0 +1,7 @@
//go:build release
package main
func developmentRunOptions(args []string) ([]string, clientRunOptions, error) {
return args, clientRunOptions{}, nil
}
@@ -0,0 +1,15 @@
//go:build release
package main
import "testing"
func TestReleaseBuildRejectsDevelopmentFlag(t *testing.T) {
args, options, err := developmentRunOptions([]string{"--dev", "run"})
if err != nil {
t.Fatal(err)
}
if err := runClient(args, options); err == nil {
t.Fatal("release build unexpectedly accepted --dev run")
}
}
+66
View File
@@ -0,0 +1,66 @@
//go:build !release
package main
import (
"path/filepath"
"testing"
)
func TestClientDevelopmentGameDirectory(t *testing.T) {
for _, test := range []struct {
name string
args []string
want string
}{
{name: "separate", args: []string{"--game-dir", filepath.Join("some", "game")}, want: filepath.Join("some", "game")},
{name: "equals", args: []string{"--game-dir=" + filepath.Join("other", "game")}, want: filepath.Join("other", "game")},
{name: "absent", args: []string{"--no-browser"}},
} {
t.Run(test.name, func(t *testing.T) {
got, err := clientDevelopmentGameDirectory(test.args)
if err != nil {
t.Fatal(err)
}
if got != test.want {
t.Fatalf("game directory = %q, want %q", got, test.want)
}
})
}
}
func TestClientDevelopmentGameDirectoryRequiresValue(t *testing.T) {
for _, args := range [][]string{{"--game-dir"}, {"--game-dir="}} {
if _, err := clientDevelopmentGameDirectory(args); err == nil {
t.Fatalf("args %v unexpectedly succeeded", args)
}
}
}
func TestDevelopmentRunOptionsUsesRepositoryFiles(t *testing.T) {
t.Setenv("APPDATA", t.TempDir())
t.Setenv("XDG_CONFIG_HOME", t.TempDir())
args, options, err := developmentRunOptions([]string{"--dev", "run", "--no-browser"})
if err != nil {
t.Fatal(err)
}
if len(args) != 1 || args[0] != "--no-browser" {
t.Fatalf("client args = %v", args)
}
for name, path := range map[string]string{
"versions": options.versionConfigPath,
"log executable": options.logExecutablePath,
"local identity": options.localIdentityPlugin,
"login UI": options.loginUIPlugin,
} {
if !filepath.IsAbs(path) {
t.Errorf("%s path is not absolute: %q", name, path)
}
}
}
func TestDevelopmentRunOptionsRequiresRun(t *testing.T) {
if _, _, err := developmentRunOptions([]string{"--dev"}); err == nil {
t.Fatal("development command without run unexpectedly succeeded")
}
}
+83
View File
@@ -0,0 +1,83 @@
package main
import (
"flag"
"fmt"
"os"
clientapp "bd2server/internal/client/app"
clientconfig "bd2server/internal/client/config"
)
type clientRunOptions struct {
versionConfigPath string
logExecutablePath string
localIdentityPlugin string
loginUIPlugin string
}
func main() {
args, options, err := developmentRunOptions(os.Args[1:])
if err == nil {
err = runClient(args, options)
}
if err != nil {
clientapp.ShowFatalError(err)
fmt.Fprintln(os.Stderr, "bd2client:", err)
os.Exit(1)
}
}
func runClient(args []string, options clientRunOptions) error {
fs := flag.NewFlagSet("bd2client", flag.ContinueOnError)
listen := fs.String("listen", "127.0.0.1:0", "loopback address for the local setup interface")
noBrowser := fs.Bool("no-browser", false, "print the interface URL without opening a browser")
gameDir := fs.String("game-dir", "", "initial Brown Dust II installation directory")
if err := fs.Parse(args); err != nil {
return err
}
if fs.NArg() != 0 {
return fmt.Errorf("unexpected argument %q", fs.Arg(0))
}
logger, logCloser, logPath, err := clientapp.OpenPersistentLogger(options.logExecutablePath)
if err != nil {
return err
}
defer logCloser.Close()
logger.Info("bd2client starting", "log_path", logPath)
var versions clientconfig.ReleaseVersions
if options.versionConfigPath == "" {
versions, err = clientconfig.ReleaseVersionsBesideExecutable()
} else {
versions, err = clientconfig.LoadReleaseVersions(options.versionConfigPath)
}
if err != nil {
logger.Error("client release version manifest is invalid", "error", err)
return err
}
if *gameDir == "" {
preferences, preferenceErr := clientconfig.LoadPreferences()
if preferenceErr != nil {
logger.Warn("could not load client preferences", "error", preferenceErr)
} else {
*gameDir = preferences.GameDirectory
}
}
if err := clientapp.Run(clientapp.Options{
Listen: *listen,
NoBrowser: *noBrowser,
InitialGameDir: *gameDir,
Logger: logger,
LogPath: logPath,
Versions: versions,
LocalIdentityPlugin: options.localIdentityPlugin,
LoginUIPlugin: options.loginUIPlugin,
}); err != nil {
logger.Error("bd2client stopped with an error", "error", err)
return err
}
logger.Info("bd2client exited")
return nil
}
+72
View File
@@ -0,0 +1,72 @@
//go:build !release
package main
import (
"errors"
"fmt"
"os"
"path/filepath"
"strings"
)
func runDevelopmentCommand(args []string) (bool, error) {
if len(args) == 0 || args[0] != "--dev" {
return false, nil
}
if len(args) < 2 || args[1] != "run" {
return true, errors.New("usage: bd2server --dev run [serve options]")
}
root, err := findDevelopmentRoot()
if err != nil {
return true, err
}
serveArgs := append([]string(nil), args[2:]...)
serveArgs = appendDefaultFlag(serveArgs, "--version-config", filepath.Join(root, "versions.json"))
serveArgs = appendDefaultFlag(serveArgs, "--authentication-config", filepath.Join(root, "authentication.json"))
serveArgs = appendDefaultFlag(serveArgs, "--resource-config", filepath.Join(root, "resources.json"))
serveArgs = appendDefaultFlag(serveArgs, "--data-dir", filepath.Join(root, "data"))
serveArgs = appendDefaultFlag(serveArgs, "--state", filepath.Join(root, "data", "state", "state.db"))
return true, serve(serveArgs)
}
func findDevelopmentRoot() (string, error) {
working, err := os.Getwd()
if err != nil {
return "", fmt.Errorf("resolve development working directory: %w", err)
}
for directory := filepath.Clean(working); ; directory = filepath.Dir(directory) {
if regularDevelopmentFile(filepath.Join(directory, "versions.json")) &&
regularDevelopmentFile(filepath.Join(directory, "go", "go.mod")) &&
regularDevelopmentFile(filepath.Join(directory, "authentication.json")) &&
regularDevelopmentFile(filepath.Join(directory, "resources.json")) {
return directory, nil
}
parent := filepath.Dir(directory)
if parent == directory {
break
}
}
return "", errors.New("development repository root not found; run from the bd2 repository")
}
func regularDevelopmentFile(path string) bool {
info, err := os.Stat(path)
return err == nil && info.Mode().IsRegular()
}
func appendDefaultFlag(args []string, name, value string) []string {
for index, arg := range args {
if arg == name || strings.HasPrefix(arg, name+"=") {
return args
}
if index > 0 && args[index-1] == name {
return args
}
}
return append(args, name, value)
}
func developmentUsage() string {
return "\n\nDevelopment build only:\n\tbd2server --dev run [serve options]"
}
+7
View File
@@ -0,0 +1,7 @@
//go:build release
package main
func runDevelopmentCommand([]string) (bool, error) { return false, nil }
func developmentUsage() string { return "" }
@@ -0,0 +1,15 @@
//go:build release
package main
import "testing"
func TestReleaseBuildDoesNotHandleDevelopmentCommand(t *testing.T) {
handled, err := runDevelopmentCommand([]string{"--dev", "run"})
if err != nil {
t.Fatal(err)
}
if handled {
t.Fatal("release build unexpectedly handled --dev run")
}
}
+18
View File
@@ -0,0 +1,18 @@
//go:build !release
package main
import "testing"
func TestAppendDefaultFlagPreservesExplicitOverride(t *testing.T) {
for _, args := range [][]string{{"--data-dir", "custom"}, {"--data-dir=custom"}} {
got := appendDefaultFlag(append([]string(nil), args...), "--data-dir", "default")
if len(got) != len(args) {
t.Fatalf("args=%v got=%v", args, got)
}
}
got := appendDefaultFlag(nil, "--data-dir", "default")
if len(got) != 2 || got[0] != "--data-dir" || got[1] != "default" {
t.Fatalf("default args=%v", got)
}
}
+161 -123
View File
@@ -9,32 +9,42 @@ import (
"net/http"
"os"
"path/filepath"
"strings"
"time"
"bd2server/internal/account"
"bd2server/internal/accountstate"
"bd2server/internal/battle"
"bd2server/internal/bootstrap"
"bd2server/internal/clientplugin"
"bd2server/internal/deck"
"bd2server/internal/feature"
"bd2server/internal/gacha"
"bd2server/internal/gamedata"
"bd2server/internal/introdb"
"bd2server/internal/mail"
"bd2server/internal/missions"
"bd2server/internal/pictorial"
"bd2server/internal/player"
"bd2server/internal/progress"
"bd2server/internal/readonly"
"bd2server/internal/schedule"
"bd2server/internal/session"
"bd2server/internal/transport"
"bd2server/internal/versionconfig"
"bd2server/internal/world"
"bd2server/internal/server/account"
"bd2server/internal/server/accountstate"
"bd2server/internal/server/auth"
"bd2server/internal/server/authconfig"
"bd2server/internal/server/battle"
"bd2server/internal/server/bootstrap"
"bd2server/internal/server/deck"
"bd2server/internal/server/feature"
"bd2server/internal/server/gacha"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/mail"
"bd2server/internal/server/missions"
"bd2server/internal/server/pictorial"
"bd2server/internal/server/player"
"bd2server/internal/server/progress"
"bd2server/internal/server/readonly"
"bd2server/internal/server/resourcefetch"
"bd2server/internal/server/resourcepolicy"
"bd2server/internal/server/schedule"
"bd2server/internal/server/session"
"bd2server/internal/server/transport"
"bd2server/internal/server/versionconfig"
"bd2server/internal/server/world"
)
func main() {
if handled, err := runDevelopmentCommand(os.Args[1:]); handled {
if err != nil {
slog.Error("development command failed", "error", err)
os.Exit(1)
}
return
}
if len(os.Args) < 2 {
usage()
os.Exit(2)
@@ -43,10 +53,10 @@ func main() {
switch os.Args[1] {
case "serve":
err = serve(os.Args[2:])
case "patch-client":
err = patchClient(os.Args[2:])
case "state":
err = stateCommand(os.Args[2:])
case "resources":
err = resourcesCommand(os.Args[2:])
case "help", "-h", "--help":
usage()
return
@@ -62,21 +72,20 @@ func main() {
func serve(args []string) (serveErr error) {
fs := flag.NewFlagSet("serve", flag.ContinueOnError)
versionConfigPath := fs.String("version-config", "", "repository versions.json override")
authConfigPath := fs.String("authentication-config", "", "authentication.json override for development")
resourceConfigPath := fs.String("resource-config", "", "resources.json override for development")
listen := fs.String("listen", "127.0.0.1:8080", "local listen address")
cdn := fs.String("cdn", "", "ServerData root (required)")
gameData := fs.String("game-data", "", "versioned GameData root (required)")
dataDir := fs.String("data-dir", "", "server data directory (defaults beside the executable)")
gameDataVersion := fs.String("game-data-version", "", "validated GameData version (defaults to versions.json)")
gameDataOrigin := fs.String("game-data-origin", "https://dl.bd2.pmang.cloud/GameData", "official repair source used only when local validation fails")
gameDataOrigin := fs.String("game-data-origin", resourcepolicy.OfficialGameDataURL, "official GameData repair source override for development")
accountSeed := fs.String("account-seed", "", "versioned local account seed")
playerSeed := fs.String("player-seed", "", "versioned starter inventory and characters")
readonlySeed := fs.String("readonly-seed", "", "versioned server schedules and optional feature defaults")
mailSeed := fs.String("mail-seed", "", "versioned starter mailbox")
stateFile := fs.String("state", `..\data\state\state.db`, "local account SQLite database")
stateFile := fs.String("state", "", "account SQLite database override")
deckSeed := fs.String("deck-seed", "", "versioned starter deck")
worldSeed := fs.String("world-seed", "", "versioned starter world")
gachaScheduleSeed := fs.String("gacha-schedule-seed", "", "versioned dynamic gacha schedule")
gameDir := fs.String("game-dir", "", "Brown Dust II client directory (required)")
identityPlugin := fs.String("identity-plugin", "", "optional BD2LocalIdentity.dll override for development")
devToolsConfig := fs.String("dev-tools-config", "", "optional local development-tool settings JSON")
if err := fs.Parse(args); err != nil {
return err
@@ -102,6 +111,49 @@ func serve(args []string) (serveErr error) {
}
}
versionconfig.Use(versions)
if *authConfigPath == "" {
*authConfigPath, err = authconfig.BesideExecutable()
if err != nil {
return err
}
}
authentication, err := authconfig.Load(*authConfigPath)
if err != nil {
return err
}
authRuntime, err := authentication.ResolveEnvironment()
if err != nil {
return err
}
defer clear(authRuntime.MasterKey)
if *resourceConfigPath == "" {
*resourceConfigPath, err = resourcepolicy.BesideExecutable()
if err != nil {
return err
}
}
resources, err := resourcepolicy.Load(*resourceConfigPath)
if err != nil {
return err
}
if *dataDir == "" {
executable, executableErr := os.Executable()
if executableErr != nil {
return fmt.Errorf("resolve server data directory: %w", executableErr)
}
*dataDir = filepath.Join(filepath.Dir(executable), "data")
}
*dataDir, err = filepath.Abs(filepath.Clean(*dataDir))
if err != nil {
return fmt.Errorf("resolve server data directory: %w", err)
}
gameData := filepath.Join(*dataDir, "resources", "GameData")
if *stateFile == "" {
*stateFile = filepath.Join(*dataDir, "state", "state.db")
}
if err := os.MkdirAll(filepath.Dir(filepath.Clean(*stateFile)), 0o755); err != nil {
return fmt.Errorf("create server state directory: %w", err)
}
seedRoot := versions.Resolve(versions.SeedDirectory)
for target, name := range map[*string]string{
accountSeed: "login_user.json", playerSeed: "starter_player.json", readonlySeed: "readonly.json",
@@ -111,38 +163,24 @@ func serve(args []string) (serveErr error) {
*target = filepath.Join(seedRoot, name)
}
}
if *cdn == "" || *gameData == "" || *gameDir == "" {
return errors.New("serve requires --game-dir, --cdn, and --game-data")
clientOrigin := "http://" + *listen
if authentication.Mode == "oauth" {
clientOrigin = strings.TrimSuffix(authentication.PublicURL, "/")
}
packagedPlugin, err := clientplugin.ResolvePackaged(*identityPlugin)
if err != nil {
return err
}
pluginResult, err := clientplugin.Install(*gameDir, packagedPlugin)
if err != nil {
return err
}
if pluginResult.Changed {
slog.Info("installed local identity plugin", "path", pluginResult.Destination)
} else {
slog.Info("local identity plugin is current", "path", pluginResult.Destination)
}
base := "http://" + *listen + "/game/"
base := clientOrigin + "/game/"
publicResources := resources.Public(versions.BundleVersion, *gameDataVersion)
cfg := bootstrap.Config{
BaseURL: base,
CDNURL: "http://" + *listen + "/assets/ServerData",
CDNURL: publicResources.ServerDataURL,
Version: versions.ClientVersion,
BundleVer: versions.BundleVersion,
GameDataURL: "http://" + *listen + "/assets/GameData",
GameDataURL: publicResources.GameDataURL,
GameDataVer: *gameDataVersion,
}
if err := cfg.Validate(); err != nil {
return err
}
if info, err := os.Stat(*cdn); err != nil || !info.IsDir() {
return fmt.Errorf("CDN directory is unavailable: %q", *cdn)
}
verifiedGameData, downloaded, err := gamedata.Ensure(context.Background(), nil, filepath.Clean(*gameData), *gameDataVersion, *gameDataOrigin)
verifiedGameData, downloaded, err := gamedata.Ensure(context.Background(), nil, gameData, *gameDataVersion, *gameDataOrigin)
if err != nil {
return fmt.Errorf("refuse to advertise unavailable or unverified GameData: %w", err)
}
@@ -157,8 +195,8 @@ func serve(args []string) (serveErr error) {
if err != nil {
return fmt.Errorf("load starter player: %w", err)
}
if login.Version != versions.ProtocolVersion || starter.Version != versions.ProtocolVersion {
return fmt.Errorf("protocol version %s requires matching account and player seeds (got %s and %s)", versions.ProtocolVersion, login.Version, starter.Version)
if login.Version != versions.ClientVersion || starter.Version != versions.ClientVersion {
return fmt.Errorf("client version %s requires matching account and player seeds (got %s and %s)", versions.ClientVersion, login.Version, starter.Version)
}
gachaSchedule, err := gacha.LoadScheduleSeed(filepath.Clean(*gachaScheduleSeed), versions.ClientVersion)
if err != nil {
@@ -171,7 +209,7 @@ func serve(args []string) (serveErr error) {
for _, window := range gachaSchedule.StepUps {
stepUpGroupIDs = append(stepUpGroupIDs, window.GroupID)
}
regularGacha, equipmentGacha, err := gamedata.LoadActiveGachaForSchedules(filepath.Clean(*gameData), *gameDataVersion, scheduleGroupIDs, stepUpGroupIDs)
regularGacha, equipmentGacha, err := gamedata.LoadActiveGachaForSchedules(gameData, *gameDataVersion, scheduleGroupIDs, stepUpGroupIDs)
if err != nil {
return fmt.Errorf("load active gacha GameData: %w", err)
}
@@ -180,6 +218,18 @@ func serve(args []string) (serveErr error) {
return fmt.Errorf("open account state database: %w", err)
}
defer stateRepository.Close()
var authService *auth.Service
if authentication.Mode == "oauth" {
authStore, err := auth.Open(filepath.Join(filepath.Dir(filepath.Clean(*stateFile)), "auth.db"), authRuntime.MasterKey)
if err != nil {
return fmt.Errorf("open authentication database: %w", err)
}
defer authStore.Close()
authService, err = auth.New(authRuntime, authStore)
if err != nil {
return err
}
}
accountDomains := []string{"characters", "collection", "deck", "equipment", "items", "mail", "missions", "progress", "wallet"}
if !stateRepository.IsNew() {
if err := stateRepository.RequireDomains(accountDomains...); err != nil {
@@ -226,7 +276,7 @@ func serve(args []string) (serveErr error) {
if err != nil {
return fmt.Errorf("load owned inventory: %w", err)
}
randomBoxes, err := gamedata.LoadRandomBoxDesign(filepath.Clean(*gameData), *gameDataVersion)
randomBoxes, err := gamedata.LoadRandomBoxDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load deterministic random-box GameData: %w", err)
}
@@ -259,7 +309,7 @@ func serve(args []string) (serveErr error) {
if err := login.AttachCurrencies(wallet); err != nil {
return fmt.Errorf("attach wallet to login: %w", err)
}
slotDesign, err := gamedata.LoadInventorySlotDesign(filepath.Clean(*gameData), *gameDataVersion)
slotDesign, err := gamedata.LoadInventorySlotDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load inventory slot GameData: %w", err)
}
@@ -286,7 +336,7 @@ func serve(args []string) (serveErr error) {
if err := mailService.AttachSeedPath(filepath.Clean(*mailSeed)); err != nil {
return fmt.Errorf("watch mail seed: %w", err)
}
missionDesign, err := gamedata.LoadMissionDesign(filepath.Clean(*gameData), *gameDataVersion)
missionDesign, err := gamedata.LoadMissionDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load mission GameData: %w", err)
}
@@ -310,39 +360,39 @@ func serve(args []string) (serveErr error) {
if err != nil {
return fmt.Errorf("load owned equipment: %w", err)
}
equipmentSlots, err := gamedata.LoadEquipmentSlots(filepath.Clean(*gameData), *gameDataVersion)
equipmentSlots, err := gamedata.LoadEquipmentSlots(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load equipment slot GameData: %w", err)
}
if err := ownedEquipment.AttachSlots(equipmentSlots); err != nil {
return fmt.Errorf("attach equipment slot GameData: %w", err)
}
equipmentUpgrade, err := gamedata.LoadEquipmentUpgradeDesign(filepath.Clean(*gameData), *gameDataVersion)
equipmentUpgrade, err := gamedata.LoadEquipmentUpgradeDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load equipment upgrade GameData: %w", err)
}
if err := ownedEquipment.AttachUpgrade(equipmentUpgrade, wallet, ownedItems); err != nil {
return fmt.Errorf("attach equipment upgrade GameData: %w", err)
}
equipmentCraft, err := gamedata.LoadEquipmentCraftDesign(filepath.Clean(*gameData), *gameDataVersion)
equipmentCraft, err := gamedata.LoadEquipmentCraftDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load equipment crafting GameData: %w", err)
}
talentGrowth, err := gamedata.LoadTalentGrowthDesign(filepath.Clean(*gameData), *gameDataVersion)
talentGrowth, err := gamedata.LoadTalentGrowthDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load talent growth GameData: %w", err)
}
if err := ownedEquipment.AttachCraft(equipmentCraft); err != nil {
return fmt.Errorf("attach equipment crafting GameData: %w", err)
}
equipmentSmelting, err := gamedata.LoadEquipmentSmeltingDesign(filepath.Clean(*gameData), *gameDataVersion)
equipmentSmelting, err := gamedata.LoadEquipmentSmeltingDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load equipment smelting GameData: %w", err)
}
if err := ownedEquipment.AttachSmelting(equipmentSmelting, wallet, ownedItems); err != nil {
return fmt.Errorf("attach equipment smelting GameData: %w", err)
}
equipmentOptionReroll, err := gamedata.LoadEquipmentOptionRerollDesign(filepath.Clean(*gameData), *gameDataVersion)
equipmentOptionReroll, err := gamedata.LoadEquipmentOptionRerollDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load equipment option reroll GameData: %w", err)
}
@@ -353,7 +403,7 @@ func serve(args []string) (serveErr error) {
if err != nil {
return fmt.Errorf("load owned collection: %w", err)
}
infiniteGacha, err := gamedata.LoadInfiniteGacha(filepath.Clean(*gameData), *gameDataVersion)
infiniteGacha, err := gamedata.LoadInfiniteGacha(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load infinite gacha GameData: %w", err)
}
@@ -382,7 +432,7 @@ func serve(args []string) (serveErr error) {
gachaService.AttachPreviewMission(func() error {
return missionService.CompleteMission(gamedata.MissionKey{GroupType: 0, GroupID: 1, ID: 111})
})
worldService, err := world.Load(filepath.Clean(*worldSeed), filepath.Clean(*gameData), *gameDataVersion,
worldService, err := world.Load(filepath.Clean(*worldSeed), gameData, *gameDataVersion,
stateRepository, progressState, starter, ownedEquipment, ownedItems, wallet)
if err != nil {
return fmt.Errorf("load world state: %w", err)
@@ -407,12 +457,12 @@ func serve(args []string) (serveErr error) {
if err := deckStateStore.AttachPresetRuntime(wallet, worldService.CharacterService(), ownedEquipment, collection); err != nil {
return fmt.Errorf("attach ordinary preset runtime: %w", err)
}
pictorialDesign, err := gamedata.LoadPictorialDesign(filepath.Clean(*gameData), *gameDataVersion)
pictorialDesign, err := gamedata.LoadPictorialDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load pictorial GameData: %w", err)
}
pictorialService := &pictorial.Service{Design: pictorialDesign, Owned: worldService}
charAwakeDesign, err := gamedata.LoadCharAwakeDesign(filepath.Clean(*gameData), *gameDataVersion)
charAwakeDesign, err := gamedata.LoadCharAwakeDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load character awakening GameData: %w", err)
}
@@ -430,7 +480,7 @@ func serve(args []string) (serveErr error) {
if err := worldService.CharacterService().AttachTalentGrowth(talentGrowth); err != nil {
return fmt.Errorf("attach character talent growth: %w", err)
}
costumePotentialDesign, err := gamedata.LoadCostumePotentialDesign(filepath.Clean(*gameData), *gameDataVersion)
costumePotentialDesign, err := gamedata.LoadCostumePotentialDesign(gameData, *gameDataVersion)
if err != nil {
return fmt.Errorf("load costume potential GameData: %w", err)
}
@@ -438,7 +488,7 @@ func serve(args []string) (serveErr error) {
if err != nil {
return err
}
battleService := battle.NewService(filepath.Clean(*gameData), *gameDataVersion, ownedItems, worldService.CurrentPackID)
battleService := battle.NewService(gameData, *gameDataVersion, ownedItems, worldService.CurrentPackID)
battleService.AttachTutorialWin(func() error {
return missionService.CompleteMission(gamedata.MissionKey{GroupType: 0, GroupID: 1, ID: 113})
})
@@ -469,6 +519,11 @@ func serve(args []string) (serveErr error) {
if err != nil {
return err
}
if authService != nil {
if err := game.AttachLoginAuthenticator(authService); err != nil {
return err
}
}
if stateRepository.IsNew() {
if err := ensureAccountStateInitialized(
progressState, deckStateStore, ownedItems, ownedEquipment,
@@ -495,7 +550,14 @@ func serve(args []string) (serveErr error) {
return err
}
dispatcher := transport.Bootstrap{Config: cfg}
handler := transport.HTTP{Dispatcher: dispatcher, Raw: game, CDNDir: filepath.Clean(*cdn), GameDataDir: filepath.Clean(*gameData)}.Handler()
var authHandler http.Handler
if authService != nil {
authHandler = authService.Handler()
}
handler := transport.HTTP{
Dispatcher: dispatcher, Raw: game, Authentication: authentication,
AuthenticationHandler: authHandler, ResourcePolicy: publicResources,
}.Handler()
server := &http.Server{
Addr: *listen,
Handler: handler,
@@ -504,7 +566,7 @@ func serve(args []string) (serveErr error) {
WriteTimeout: 20 * time.Second,
IdleTimeout: 60 * time.Second,
}
slog.Info("BD2 local server listening", "address", *listen, "client", cfg.Version, "bundle", cfg.BundleVer, "cdn", *cdn, "gameData", verifiedGameData.ArchivePath, "gameDataEntries", verifiedGameData.EntryCount, "accountSeed", *accountSeed)
slog.Info("BD2 server listening", "address", *listen, "client", cfg.Version, "bundle", cfg.BundleVer, "resourceMode", publicResources.Mode, "gameData", verifiedGameData.ArchivePath, "gameDataEntries", verifiedGameData.EntryCount, "accountSeed", *accountSeed)
return server.ListenAndServe()
}
@@ -524,73 +586,49 @@ func ensureAccountStateInitialized(stores ...accountStateInitializer) error {
return nil
}
func patchClient(args []string) error {
fs := flag.NewFlagSet("patch-client", flag.ContinueOnError)
gameDir := fs.String("game-dir", "", "BrownDust II game directory (required)")
serverURL := fs.String("url", "http://127.0.0.1:8080/game/", "exactly 27-byte replacement LIVE_URL")
verify := fs.Bool("verify", false, "inspect the embedded LIVE_URL without writing")
if err := fs.Parse(args); err != nil {
func resourcesCommand(args []string) error {
if len(args) == 0 || args[0] != "fetch" {
return errors.New("resources requires the fetch subcommand")
}
fs := flag.NewFlagSet("resources fetch", flag.ContinueOnError)
versionConfigPath := fs.String("version-config", "", "repository versions.json override")
output := fs.String("output", "", "resource mirror output directory (required)")
platform := fs.String("platform", "StandaloneWindows64", "official ServerData platform")
if err := fs.Parse(args[1:]); err != nil {
return err
}
if *gameDir == "" {
return errors.New("patch-client requires --game-dir")
if *output == "" {
return errors.New("resources fetch requires --output")
}
if *verify {
result, err := introdb.VerifyClient(*gameDir)
if err != nil {
return err
}
fmt.Printf("resources.assets: %s\nTextAsset pathID: %d\nLIVE_URL: %s\n", result.AssetsPath, result.ObjectPath, result.URL)
return nil
var versions versionconfig.Config
var err error
if *versionConfigPath == "" {
versions, err = versionconfig.Find()
} else {
versions, err = versionconfig.Load(*versionConfigPath)
}
result, err := introdb.PatchClient(*gameDir, *serverURL)
if err != nil {
return err
}
disabled, err := disableLegacyPlugin(*gameDir)
manifest, err := resourcefetch.Fetch(context.Background(), resourcefetch.Options{
OutputRoot: *output, Platform: *platform, BundleVersion: versions.BundleVersion,
GameDataVersion: versions.GameDataVersion,
Progress: func(message string) { slog.Info(message) },
})
if err != nil {
return err
}
verified, err := introdb.VerifyClient(*gameDir)
if err != nil {
return fmt.Errorf("post-patch verification: %w", err)
}
fmt.Printf("patched: %s\nbackup: %s\nLIVE_URL: %s\n", result.AssetsPath, result.BackupPath, verified.URL)
if disabled != "" {
fmt.Printf("legacy plugin disabled: %s\n", disabled)
}
slog.Info("official resource mirror complete", "output", *output, "bundles", manifest.ServerData.Bundles, "bytes", manifest.ServerData.Bytes)
return nil
}
func disableLegacyPlugin(gameDir string) (string, error) {
source := filepath.Join(gameDir, "BepInEx", "plugins", "PluginLocalRes.dll")
destination := filepath.Join(gameDir, "BepInEx", "disabled", "PluginLocalRes.dll")
if _, err := os.Stat(source); errors.Is(err, os.ErrNotExist) {
return "", nil
} else if err != nil {
return "", fmt.Errorf("inspect legacy local-resource plugin: %w", err)
}
if _, err := os.Stat(destination); err == nil {
return "", fmt.Errorf("legacy plugin exists at both active and disabled paths")
} else if !errors.Is(err, os.ErrNotExist) {
return "", err
}
if err := os.MkdirAll(filepath.Dir(destination), 0o755); err != nil {
return "", err
}
if err := os.Rename(source, destination); err != nil {
return "", fmt.Errorf("disable legacy local-resource plugin: %w", err)
}
return destination, nil
}
func usage() {
fmt.Fprintln(os.Stderr, `bd2server - BrownDust II local development server
fmt.Fprintln(os.Stderr, `bd2server - BrownDust II server
Usage:
bd2server serve --game-dir DIR --cdn DIR --game-data DIR [--version-config FILE] [options]
bd2server patch-client --game-dir DIR [options]
bd2server state check [options]
bd2server serve [--data-dir DIR] [--version-config FILE] [options]
bd2server resources fetch --output DIR [--version-config FILE]
bd2server state check [options]
The server binds to loopback by default and is intended for local research.`)
The server binds to loopback by default.`+developmentUsage())
}
+1 -1
View File
@@ -6,7 +6,7 @@ import (
"fmt"
"path/filepath"
"bd2server/internal/accountstate"
"bd2server/internal/server/accountstate"
)
func stateCommand(args []string) error {
+1 -1
View File
@@ -4,7 +4,7 @@ import (
"fmt"
"strings"
"bd2server/internal/accountstate"
"bd2server/internal/server/accountstate"
)
func stateProblemsError(prefix string, problems []accountstate.Problem) error {
+519
View File
@@ -0,0 +1,519 @@
// Package app serves bd2client's embedded, loopback-only setup interface.
package app
import (
"context"
"crypto/rand"
"crypto/subtle"
"embed"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"html/template"
"io"
"log/slog"
"net"
"net/http"
"os"
"path/filepath"
"runtime"
"sync"
"time"
clientconfig "bd2server/internal/client/config"
clientlayout "bd2server/internal/client/layout"
clientsetup "bd2server/internal/client/setup"
)
//go:embed web/index.html
var webFS embed.FS
var errGameAlreadyRunning = errors.New("Brown Dust II is already running")
type Options struct {
Listen string
NoBrowser bool
InitialGameDir string
Logger *slog.Logger
LogPath string
Versions clientconfig.ReleaseVersions
LocalIdentityPlugin string
LoginUIPlugin string
}
type request struct {
GameDirectory string `json:"game_directory"`
ServerOrigin string `json:"server_origin"`
CDNMode clientconfig.CDNMode `json:"cdn_mode"`
LocalResourceDirectory string `json:"local_resource_directory"`
UILanguage string `json:"ui_language"`
}
type response struct {
OK bool `json:"ok"`
Message string `json:"message,omitempty"`
Data any `json:"data,omitempty"`
}
type handler struct {
token string
origin string
initialGameDir string
browse func(string) (string, error)
browseResources func(string) (string, error)
shutdown func()
quitOnce sync.Once
logger *slog.Logger
logPath string
versions clientconfig.ReleaseVersions
initialSettings clientconfig.Settings
autoOpen bool
localIdentityPlugin string
loginUIPlugin string
}
func Run(options Options) error {
logger := options.Logger
if logger == nil {
logger = slog.Default()
}
listen := options.Listen
if listen == "" {
listen = "127.0.0.1:0"
}
listener, err := net.Listen("tcp", listen)
if err != nil {
logger.Error("could not start local interface", "error", err)
return fmt.Errorf("start bd2client interface: %w", err)
}
address := listener.Addr().(*net.TCPAddr)
if !address.IP.IsLoopback() {
_ = listener.Close()
logger.Error("refused non-loopback interface", "address", listener.Addr().String())
return errors.New("bd2client interface must listen on a loopback address")
}
token, err := newToken()
if err != nil {
_ = listener.Close()
logger.Error("could not create local interface session", "error", err)
return err
}
origin := "http://" + listener.Addr().String()
initialSettings := clientconfig.Settings{ServerOrigin: "http://127.0.0.1:8080", CDNMode: clientconfig.CDNOfficial}
autoOpen := false
if options.InitialGameDir != "" {
if _, inspectErr := clientsetup.Inspect(options.InitialGameDir, options.Versions); inspectErr != nil {
logger.Warn("saved game directory is no longer valid", "error", inspectErr)
} else if loaded, loadErr := clientconfig.Load(options.InitialGameDir); loadErr != nil {
logger.Warn("saved client connection settings are unavailable", "error", loadErr)
} else {
initialSettings = loaded
autoOpen = true
}
}
server := &http.Server{
ReadHeaderTimeout: 5 * time.Second,
ReadTimeout: 15 * time.Second,
WriteTimeout: 30 * time.Second,
IdleTimeout: 60 * time.Second,
}
done := make(chan struct{})
h := &handler{
token: token,
origin: origin,
initialGameDir: options.InitialGameDir,
browse: browseForGameDirectory,
browseResources: browseForResourceDirectory,
logger: logger,
logPath: options.LogPath,
versions: options.Versions,
initialSettings: initialSettings,
autoOpen: autoOpen,
localIdentityPlugin: options.LocalIdentityPlugin,
loginUIPlugin: options.LoginUIPlugin,
shutdown: func() {
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
_ = server.Shutdown(ctx)
}()
},
}
server.Handler = h.routes()
logger.Info("local interface listening", "address", listener.Addr().String())
go func() {
err := server.Serve(listener)
if err != nil && !errors.Is(err, http.ErrServerClosed) {
logger.Error("local interface stopped unexpectedly", "error", err)
}
close(done)
}()
pageURL := origin + "/?session=" + token
fmt.Fprintf(os.Stdout, "BD2 Client Studio: %s\n", pageURL)
if !options.NoBrowser {
if err := openBrowser(pageURL); err != nil {
logger.Warn("could not open client window automatically", "error", err)
} else {
logger.Info("client window opened")
}
} else {
logger.Info("automatic client window disabled")
}
<-done
logger.Info("local interface stopped")
return nil
}
func (h *handler) routes() http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("GET /", h.index)
mux.HandleFunc("POST /api/browse", h.observe("browse game directory", h.authorize(h.browseDirectory)))
mux.HandleFunc("POST /api/browse-resources", h.observe("browse resource directory", h.authorize(h.browseResourceDirectory)))
mux.HandleFunc("POST /api/inspect", h.observe("inspect game directory", h.authorize(h.inspect)))
mux.HandleFunc("POST /api/resources", h.observe("check resource policy", h.authorize(h.resources)))
mux.HandleFunc("POST /api/save", h.observe("save settings", h.authorize(h.save)))
mux.HandleFunc("POST /api/patch", h.observe("patch client", h.authorize(h.patch)))
mux.HandleFunc("POST /api/install", h.observe("install plugins", h.authorize(h.install)))
mux.HandleFunc("POST /api/launch", h.observe("launch game", h.authorize(h.launch)))
mux.HandleFunc("POST /api/quit", h.observe("quit", h.authorize(h.quit)))
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "no-referrer")
w.Header().Set("X-Frame-Options", "DENY")
mux.ServeHTTP(w, r)
})
}
func (h *handler) index(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/" || r.URL.Query().Get("session") != h.token {
http.NotFound(w, r)
return
}
data, err := webFS.ReadFile("web/index.html")
if err != nil {
http.Error(w, "embedded interface unavailable", http.StatusInternalServerError)
return
}
tmpl, err := template.New("index").Parse(string(data))
if err != nil {
http.Error(w, "embedded interface invalid", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "text/html; charset=utf-8")
w.Header().Set("Content-Security-Policy", "default-src 'none'; style-src 'unsafe-inline'; script-src 'unsafe-inline'; img-src data:; connect-src 'self'; font-src 'self'")
_ = tmpl.Execute(w, map[string]string{
"Token": h.token, "GameDirectory": h.initialGameDir, "LogPath": h.logPath, "Platform": runtime.GOOS,
"ServerOrigin": h.initialSettings.ServerOrigin, "CDNMode": string(h.initialSettings.CDNMode),
"LocalResourceDirectory": h.initialSettings.LocalResourceDirectory,
"AutoOpen": fmt.Sprintf("%t", h.autoOpen),
})
}
func (h *handler) authorize(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
provided := r.Header.Get("X-BD2-Session")
if len(provided) != len(h.token) || subtle.ConstantTimeCompare([]byte(provided), []byte(h.token)) != 1 {
h.log().Warn("API request rejected", "operation", r.URL.Path, "reason", "invalid session")
http.Error(w, "forbidden", http.StatusForbidden)
return
}
if origin := r.Header.Get("Origin"); origin != "" && origin != h.origin {
h.log().Warn("API request rejected", "operation", r.URL.Path, "reason", "foreign origin")
http.Error(w, "forbidden origin", http.StatusForbidden)
return
}
next(w, r)
}
}
type statusWriter struct {
http.ResponseWriter
status int
}
func (w *statusWriter) WriteHeader(status int) {
w.status = status
w.ResponseWriter.WriteHeader(status)
}
func (h *handler) observe(operation string, next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
tracked := &statusWriter{ResponseWriter: w, status: http.StatusOK}
next(tracked, r)
if tracked.status >= http.StatusBadRequest {
h.log().Error("client API operation failed", "operation", operation, "status", tracked.status)
return
}
h.log().Info("client API operation completed", "operation", operation, "status", tracked.status)
}
}
func (h *handler) log() *slog.Logger {
if h.logger != nil {
return h.logger
}
return slog.Default()
}
func (h *handler) browseDirectory(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
h.log().Info("directory selection opened", "kind", "game")
dir, err := h.browse(input.UILanguage)
if err != nil {
h.writeError(w, err)
return
}
if dir == "" {
h.log().Info("directory selection cancelled", "kind", "game")
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Selection cancelled"})
return
}
status, err := clientsetup.Inspect(dir, h.versions)
if err != nil {
h.writeError(w, err)
return
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Game client found", Data: status})
}
func (h *handler) browseResourceDirectory(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
h.log().Info("directory selection opened", "kind", "resources")
dir, err := h.browseResources(input.UILanguage)
if err != nil {
h.writeError(w, err)
return
}
if dir == "" {
h.log().Info("directory selection cancelled", "kind", "resources")
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Selection cancelled"})
return
}
policy, err := clientsetup.FetchResourcePolicy(context.Background(), nil, clientconfig.Settings{
ServerOrigin: "http://127.0.0.1",
CDNMode: clientconfig.CDNLocal,
LocalResourceDirectory: dir,
}, h.versions)
if err != nil {
h.writeError(w, err)
return
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Local resource directory found", Data: policy})
}
func (h *handler) inspect(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
status, err := clientsetup.Inspect(input.GameDirectory, h.versions)
if err != nil {
h.writeError(w, err)
return
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Game directory is valid", Data: status})
}
func (h *handler) resources(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
ctx, cancel := context.WithTimeout(r.Context(), 12*time.Second)
defer cancel()
policy, err := clientsetup.FetchResourcePolicy(ctx, nil, input.settings(), h.versions)
if err != nil {
h.writeError(w, err)
return
}
message := "The client will use the release-locked official CDN"
if policy.Mode == clientconfig.CDNLocal {
message = "Local resources verified"
} else if policy.Mode == clientconfig.CDNServer {
message = "Server resource policy verified"
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: message, Data: policy})
}
func (h *handler) save(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
settings, err := clientsetup.SaveSettings(input.GameDirectory, input.settings(), h.versions)
if err != nil {
h.writeError(w, err)
return
}
if err := clientconfig.SavePreferences(input.GameDirectory); err != nil {
h.writeError(w, err)
return
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Connection settings saved", Data: settings})
}
func (h *handler) patch(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
result, err := clientsetup.Patch(input.GameDirectory, input.settings(), h.versions)
if err != nil {
h.writeError(w, err)
return
}
if err := clientconfig.SavePreferences(input.GameDirectory); err != nil {
h.writeError(w, err)
return
}
message := "Client entry point patched; the original backup was retained"
if !result.Changed {
message = "Client patch is already complete; no asset file was rewritten"
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: message, Data: result})
}
func (h *handler) install(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
result, err := clientsetup.InstallPlugins(
input.GameDirectory,
input.settings(),
h.versions,
h.localIdentityPlugin,
h.loginUIPlugin,
)
if err != nil {
h.writeError(w, err)
return
}
if err := clientconfig.SavePreferences(input.GameDirectory); err != nil {
h.writeError(w, err)
return
}
message := "BD2 client plugins installed or updated"
if !result.LocalIdentity.Changed && !result.LoginUI.Changed {
message = "BD2 client plugins are already up to date; no DLL was rewritten"
}
h.writeJSON(w, http.StatusOK, response{OK: true, Message: message, Data: result})
}
func (h *handler) launch(w http.ResponseWriter, r *http.Request) {
input, ok := h.decode(w, r)
if !ok {
return
}
status, err := clientsetup.Inspect(input.GameDirectory, h.versions)
if err != nil {
h.writeError(w, err)
return
}
if status.PatchedURL != clientsetup.PatchPlaceholder {
h.writeError(w, errors.New("apply the client patch before launching the game"))
return
}
if !status.BepInEx {
h.writeError(w, errors.New("install BepInEx before launching the game"))
return
}
ctx, cancel := context.WithTimeout(r.Context(), 12*time.Second)
defer cancel()
if _, err := clientsetup.FetchResourcePolicy(ctx, nil, input.settings(), h.versions); err != nil {
h.writeError(w, err)
return
}
if _, err := clientsetup.SaveSettings(input.GameDirectory, input.settings(), h.versions); err != nil {
h.writeError(w, err)
return
}
if err := clientconfig.SavePreferences(input.GameDirectory); err != nil {
h.writeError(w, err)
return
}
installation, err := clientlayout.Resolve(status.GameDirectory)
if err != nil {
h.writeError(w, err)
return
}
if !installation.SupportedOnHost() {
h.writeError(w, fmt.Errorf("cannot launch a %s game client from this operating system", installation.Kind))
return
}
for _, pluginName := range []string{"BD2LocalIdentity.dll", "BD2LoginUI.dll"} {
info, statErr := os.Stat(filepath.Join(installation.Plugins, pluginName))
if statErr != nil || !info.Mode().IsRegular() {
h.writeError(w, fmt.Errorf("install or update the client plugins before launching; %s is missing", pluginName))
return
}
}
if err := launchGame(installation.LaunchTarget()); err != nil {
if errors.Is(err, errGameAlreadyRunning) {
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Brown Dust II is already running"})
return
}
h.writeError(w, fmt.Errorf("launch Brown Dust II: %w", err))
return
}
h.log().Info("game launch requested", "platform", installation.Kind, "client_version", status.ClientVersion)
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Brown Dust II started"})
}
func (h *handler) quit(w http.ResponseWriter, _ *http.Request) {
h.log().Info("client exit requested")
h.writeJSON(w, http.StatusOK, response{OK: true, Message: "Client tool exited"})
h.quitOnce.Do(h.shutdown)
}
func (h *handler) decode(w http.ResponseWriter, r *http.Request) (request, bool) {
r.Body = http.MaxBytesReader(w, r.Body, 64<<10)
decoder := json.NewDecoder(r.Body)
decoder.DisallowUnknownFields()
var input request
if err := decoder.Decode(&input); err != nil {
h.writeError(w, fmt.Errorf("invalid request: %w", err))
return request{}, false
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
h.writeError(w, errors.New("request must contain exactly one JSON object"))
return request{}, false
}
return input, true
}
func (r request) settings() clientconfig.Settings {
return clientconfig.Settings{
ServerOrigin: r.ServerOrigin,
CDNMode: r.CDNMode,
LocalResourceDirectory: r.LocalResourceDirectory,
}
}
func (h *handler) writeError(w http.ResponseWriter, err error) {
h.log().Error("client operation error", "error", err)
h.writeJSON(w, http.StatusBadRequest, response{OK: false, Message: err.Error()})
}
func (h *handler) writeJSON(w http.ResponseWriter, status int, value response) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(value)
}
func newToken() (string, error) {
buffer := make([]byte, 32)
if _, err := rand.Read(buffer); err != nil {
return "", fmt.Errorf("generate UI session: %w", err)
}
return base64.RawURLEncoding.EncodeToString(buffer), nil
}
+129
View File
@@ -0,0 +1,129 @@
package app
import (
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
clientconfig "bd2server/internal/client/config"
)
func TestIndexRequiresSessionAndServesEmbeddedStudio(t *testing.T) {
h := &handler{token: "test-session", origin: "http://127.0.0.1"}
server := httptest.NewServer(h.routes())
defer server.Close()
response, err := http.Get(server.URL + "/")
if err != nil {
t.Fatal(err)
}
_ = response.Body.Close()
if response.StatusCode != http.StatusNotFound {
t.Fatalf("without session status=%d", response.StatusCode)
}
response, err = http.Get(server.URL + "/?session=test-session")
if err != nil {
t.Fatal(err)
}
defer response.Body.Close()
buffer, err := io.ReadAll(response.Body)
if err != nil {
t.Fatal(err)
}
page := string(buffer)
for _, marker := range []string{
"BD2 Client Studio", "test-session", `name="bd2-platform"`,
`id="directoryScene"`, `id="serverScene"`, `id="deskScene"`,
`id="gameDir"`, `id="origin"`, `id="patch"`, `id="install"`, `id="launch"`,
`value="official"`, `value="local"`, `value="server"`,
`id="localResourceDir"`, `id="browseResources"`,
`Asia/Shanghai`, `Asia/Hong_Kong`, `Asia/Macau`, `Asia/Taipei`,
`const zhCN=CHINA_TIME_ZONES.has(detectedTimeZone)`,
"opening-curtain", "is-entering", "@keyframes reveal", "prefers-reduced-motion",
} {
if !strings.Contains(page, marker) {
t.Errorf("page lacks %q", marker)
}
}
if response.Header.Get("Content-Security-Policy") == "" || response.Header.Get("Cache-Control") != "no-store" {
t.Fatalf("security headers=%v", response.Header)
}
}
func TestAPIRejectsMalformedAndTrailingJSON(t *testing.T) {
h := &handler{token: "test-session", origin: "http://local.invalid", versions: clientconfig.ReleaseVersions{ClientVersion: "2.35.10"}}
for name, body := range map[string]string{
"malformed": `{`,
"trailing": `{}` + `{}`,
"unknown": `{"unexpected":true}`,
} {
t.Run(name, func(t *testing.T) {
request := httptest.NewRequest(http.MethodPost, "/api/inspect", strings.NewReader(body))
request.Header.Set("X-BD2-Session", "test-session")
request.Header.Set("Origin", "http://local.invalid")
response := httptest.NewRecorder()
h.routes().ServeHTTP(response, request)
if response.Code != http.StatusBadRequest {
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
}
})
}
}
func TestAPIRejectsMissingTokenAndForeignOrigin(t *testing.T) {
h := &handler{token: "test-session", origin: "http://local.invalid", browse: func(string) (string, error) { return "", nil }}
for name, values := range map[string][2]string{
"missing token": {"", "http://local.invalid"},
"foreign origin": {"test-session", "https://attacker.invalid"},
} {
t.Run(name, func(t *testing.T) {
request := httptest.NewRequest(http.MethodPost, "/api/browse", strings.NewReader("{}"))
request.Header.Set("X-BD2-Session", values[0])
request.Header.Set("Origin", values[1])
response := httptest.NewRecorder()
h.routes().ServeHTTP(response, request)
if response.Code != http.StatusForbidden {
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
}
})
}
}
func TestInspectAPI(t *testing.T) {
dir := t.TempDir()
for path, data := range map[string][]byte{
filepath.Join(dir, "BrownDust II.exe"): []byte("exe"),
filepath.Join(dir, "BrownDust II_Data", "resources.assets"): []byte("not a real Unity file"),
filepath.Join(dir, "BrownDust II_Data", "globalgamemanagers"): []byte("\x002.35.10\x00"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
h := &handler{token: "test-session", origin: "http://local.invalid", versions: clientconfig.ReleaseVersions{ClientVersion: "2.35.10"}}
body, _ := json.Marshal(request{GameDirectory: dir})
req := httptest.NewRequest(http.MethodPost, "/api/inspect", strings.NewReader(string(body)))
req.Header.Set("X-BD2-Session", "test-session")
req.Header.Set("Origin", "http://local.invalid")
recorder := httptest.NewRecorder()
h.routes().ServeHTTP(recorder, req)
if recorder.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", recorder.Code, recorder.Body.String())
}
var result response
if err := json.Unmarshal(recorder.Body.Bytes(), &result); err != nil {
t.Fatal(err)
}
if !result.OK {
t.Fatalf("response=%+v", result)
}
}
+147
View File
@@ -0,0 +1,147 @@
//go:build windows
package app
import (
"fmt"
"runtime"
"syscall"
"unsafe"
"golang.org/x/sys/windows"
)
const (
coinitApartmentThreaded = 0x2
clsctxInprocServer = 0x1
fosNoChangeDir = 0x00000008
fosPickFolders = 0x00000020
fosForceFileSystem = 0x00000040
fosPathMustExist = 0x00000800
fosDontAddToRecent = 0x02000000
sigdnFileSystemPath = 0x80058000
errorCancelled = 0x800704c7
)
var (
ole32DLL = windows.NewLazySystemDLL("ole32.dll")
user32DLL = windows.NewLazySystemDLL("user32.dll")
coInitializeEx = ole32DLL.NewProc("CoInitializeEx")
coUninitialize = ole32DLL.NewProc("CoUninitialize")
coCreateInstance = ole32DLL.NewProc("CoCreateInstance")
coTaskMemFree = ole32DLL.NewProc("CoTaskMemFree")
getForegroundWindow = user32DLL.NewProc("GetForegroundWindow")
clsidFileOpenDialog = windows.GUID{Data1: 0xdc1c5a9c, Data2: 0xe88a, Data3: 0x4dde, Data4: [8]byte{0xa5, 0xa1, 0x60, 0xf8, 0x2a, 0x20, 0xae, 0xf7}}
iidIFileOpenDialog = windows.GUID{Data1: 0xd57c7288, Data2: 0xd4ad, Data3: 0x4768, Data4: [8]byte{0xbe, 0x02, 0x9d, 0x96, 0x95, 0x32, 0xd9, 0x60}}
)
// comObject is sufficient for IFileOpenDialog and IShellItem because COM
// interfaces begin with a pointer to a vtable. The methods used below are
// selected by their documented vtable positions.
type comObject struct {
vtable *[29]uintptr
}
func browseForDirectory(titleText string) (string, error) {
// COM apartment state belongs to an OS thread. Keep this handler on one
// thread from initialization until every interface has been released.
runtime.LockOSThread()
defer runtime.UnlockOSThread()
result, _, _ := coInitializeEx.Call(0, coinitApartmentThreaded)
if hresultFailed(result) {
return "", hresultError("initialize Windows directory picker", result)
}
defer coUninitialize.Call()
var dialog *comObject
result, _, _ = coCreateInstance.Call(
uintptr(unsafe.Pointer(&clsidFileOpenDialog)),
0,
clsctxInprocServer,
uintptr(unsafe.Pointer(&iidIFileOpenDialog)),
uintptr(unsafe.Pointer(&dialog)),
)
if hresultFailed(result) {
return "", hresultError("create Windows directory picker", result)
}
if dialog == nil {
return "", fmt.Errorf("create Windows directory picker: the system returned no dialog")
}
defer comRelease(dialog)
var options uint32
result = comCall(dialog, 10, uintptr(unsafe.Pointer(&options))) // IFileDialog::GetOptions
if hresultFailed(result) {
return "", hresultError("read Windows directory picker options", result)
}
options |= fosNoChangeDir | fosPickFolders | fosForceFileSystem | fosPathMustExist | fosDontAddToRecent
result = comCall(dialog, 9, uintptr(options)) // IFileDialog::SetOptions
if hresultFailed(result) {
return "", hresultError("set Windows directory picker options", result)
}
title, err := windows.UTF16PtrFromString(titleText)
if err != nil {
return "", fmt.Errorf("set Windows directory picker title: %w", err)
}
result = comCall(dialog, 17, uintptr(unsafe.Pointer(title))) // IFileDialog::SetTitle
runtime.KeepAlive(title)
if hresultFailed(result) {
return "", hresultError("set Windows directory picker title", result)
}
owner, _, _ := getForegroundWindow.Call()
result = comCall(dialog, 3, owner) // IModalWindow::Show
if uint32(result) == errorCancelled {
return "", nil
}
if hresultFailed(result) {
return "", hresultError("show Windows directory picker", result)
}
var item *comObject
result = comCall(dialog, 20, uintptr(unsafe.Pointer(&item))) // IFileDialog::GetResult
if hresultFailed(result) {
return "", hresultError("read selected directory", result)
}
if item == nil {
return "", fmt.Errorf("read selected directory: the system returned no directory")
}
defer comRelease(item)
var path *uint16
result = comCall(item, 5, sigdnFileSystemPath, uintptr(unsafe.Pointer(&path))) // IShellItem::GetDisplayName
if hresultFailed(result) {
return "", hresultError("read selected directory path", result)
}
if path == nil {
return "", fmt.Errorf("read selected directory path: the system returned an empty path")
}
defer coTaskMemFree.Call(uintptr(unsafe.Pointer(path)))
return windows.UTF16PtrToString(path), nil
}
func comCall(object *comObject, method int, args ...uintptr) uintptr {
callArgs := make([]uintptr, 1, len(args)+1)
callArgs[0] = uintptr(unsafe.Pointer(object))
callArgs = append(callArgs, args...)
result, _, _ := syscall.SyscallN(object.vtable[method], callArgs...)
return result
}
func comRelease(object *comObject) {
if object != nil {
comCall(object, 2) // IUnknown::Release
}
}
func hresultFailed(result uintptr) bool {
return int32(uint32(result)) < 0
}
func hresultError(action string, result uintptr) error {
return fmt.Errorf("%s: HRESULT 0x%08X", action, uint32(result))
}
+128
View File
@@ -0,0 +1,128 @@
package app
import (
"fmt"
"io"
"log/slog"
"os"
"path/filepath"
"runtime"
"sync"
)
const (
clientLogName = "bd2client.log"
clientLogBackupName = "bd2client.log.1"
clientLogMaxBytes = 2 << 20
)
// OpenPersistentLogger creates the GUI client's bounded, persistent log next
// to the executable. The active log is capped at 2 MiB and one previous log is
// retained, so a client left installed for a long time cannot grow without
// limit.
func OpenPersistentLogger(executablePath string) (*slog.Logger, io.Closer, string, error) {
if executablePath == "" {
var err error
executablePath, err = os.Executable()
if err != nil {
return nil, nil, "", fmt.Errorf("locate bd2client executable: %w", err)
}
}
absolute, err := filepath.Abs(executablePath)
if err != nil {
return nil, nil, "", fmt.Errorf("resolve bd2client executable path: %w", err)
}
logDirectory := filepath.Join(filepath.Dir(absolute), "logs")
if runtime.GOOS == "darwin" {
home, homeErr := os.UserHomeDir()
if homeErr != nil {
return nil, nil, "", fmt.Errorf("locate macOS user home for logs: %w", homeErr)
}
logDirectory = filepath.Join(home, "Library", "Logs", "BD2 Client Studio")
}
if err := os.MkdirAll(logDirectory, 0o700); err != nil {
return nil, nil, "", fmt.Errorf("create bd2client log directory: %w", err)
}
path := filepath.Join(logDirectory, clientLogName)
writer, err := openRollingLog(path, filepath.Join(logDirectory, clientLogBackupName), clientLogMaxBytes)
if err != nil {
return nil, nil, "", err
}
logger := slog.New(slog.NewTextHandler(writer, &slog.HandlerOptions{Level: slog.LevelInfo}))
return logger, writer, path, nil
}
type rollingLog struct {
mu sync.Mutex
path string
backupPath string
maxBytes int64
file *os.File
size int64
}
func openRollingLog(path, backupPath string, maxBytes int64) (*rollingLog, error) {
file, err := os.OpenFile(path, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600)
if err != nil {
return nil, fmt.Errorf("open bd2client log: %w", err)
}
info, err := file.Stat()
if err != nil {
_ = file.Close()
return nil, fmt.Errorf("inspect bd2client log: %w", err)
}
return &rollingLog{
path: path,
backupPath: backupPath,
maxBytes: maxBytes,
file: file,
size: info.Size(),
}, nil
}
func (w *rollingLog) Write(data []byte) (int, error) {
w.mu.Lock()
defer w.mu.Unlock()
if w.file == nil {
return 0, os.ErrClosed
}
if w.size > 0 && w.size+int64(len(data)) > w.maxBytes {
if err := w.rotate(); err != nil {
return 0, err
}
}
written, err := w.file.Write(data)
w.size += int64(written)
return written, err
}
func (w *rollingLog) rotate() error {
if err := w.file.Close(); err != nil {
return fmt.Errorf("close bd2client log for rotation: %w", err)
}
w.file = nil
if err := os.Remove(w.backupPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("replace bd2client log backup: %w", err)
}
if err := os.Rename(w.path, w.backupPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("rotate bd2client log: %w", err)
}
file, err := os.OpenFile(w.path, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600)
if err != nil {
return fmt.Errorf("create rotated bd2client log: %w", err)
}
w.file = file
w.size = 0
return nil
}
func (w *rollingLog) Close() error {
w.mu.Lock()
defer w.mu.Unlock()
if w.file == nil {
return nil
}
err := w.file.Close()
w.file = nil
return err
}
+107
View File
@@ -0,0 +1,107 @@
package app
import (
"bytes"
"log/slog"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"unicode"
"unicode/utf8"
)
func TestOpenPersistentLoggerUsesExecutableLogDirectory(t *testing.T) {
root := t.TempDir()
executable := filepath.Join(root, "bd2client.exe")
logger, closer, logPath, err := OpenPersistentLogger(executable)
if err != nil {
t.Fatal(err)
}
logger.Info("test entry")
if err := closer.Close(); err != nil {
t.Fatal(err)
}
wantPath := filepath.Join(root, "logs", clientLogName)
if logPath != wantPath {
t.Fatalf("log path=%q want=%q", logPath, wantPath)
}
data, err := os.ReadFile(logPath)
if err != nil {
t.Fatal(err)
}
if !bytes.Contains(data, []byte("test entry")) {
t.Fatalf("log does not contain test entry: %s", data)
}
}
func TestPersistentOperationalLogsUseEnglish(t *testing.T) {
var output bytes.Buffer
logger := slog.New(slog.NewTextHandler(&output, nil))
h := &handler{token: "test-session", origin: "http://local.invalid", logger: logger}
body := `{"game_directory":"Z:\\missing","server_origin":"http://127.0.0.1:8080","cdn_mode":"official","local_resource_directory":"","ui_language":"zh-CN"}`
request := httptest.NewRequest(http.MethodPost, "/api/inspect", strings.NewReader(body))
request.Header.Set("X-BD2-Session", "test-session")
request.Header.Set("Origin", "http://local.invalid")
response := httptest.NewRecorder()
h.routes().ServeHTTP(response, request)
for len(output.Bytes()) > 0 {
r, size := utf8.DecodeRune(output.Bytes())
if unicode.Is(unicode.Han, r) {
t.Fatalf("operational log contains Han character %q: %s", r, output.String())
}
output.Next(size)
}
}
func TestRollingLogRetainsOneBackup(t *testing.T) {
root := t.TempDir()
path := filepath.Join(root, clientLogName)
backup := filepath.Join(root, clientLogBackupName)
writer, err := openRollingLog(path, backup, 8)
if err != nil {
t.Fatal(err)
}
if _, err := writer.Write([]byte("first")); err != nil {
t.Fatal(err)
}
if _, err := writer.Write([]byte("second")); err != nil {
t.Fatal(err)
}
if err := writer.Close(); err != nil {
t.Fatal(err)
}
active, err := os.ReadFile(path)
if err != nil {
t.Fatal(err)
}
previous, err := os.ReadFile(backup)
if err != nil {
t.Fatal(err)
}
if string(active) != "second" || string(previous) != "first" {
t.Fatalf("active=%q backup=%q", active, previous)
}
}
func TestAuthorizationLogDoesNotIncludeSessionValues(t *testing.T) {
var output bytes.Buffer
logger := slog.New(slog.NewTextHandler(&output, nil))
h := &handler{token: "expected-session-secret", origin: "http://local.invalid", logger: logger}
request := httptest.NewRequest(http.MethodPost, "/api/inspect", strings.NewReader("{}"))
request.Header.Set("X-BD2-Session", "provided-session-secret")
request.Header.Set("Origin", "http://local.invalid")
response := httptest.NewRecorder()
h.routes().ServeHTTP(response, request)
if response.Code != http.StatusForbidden {
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
}
logged := output.String()
for _, secret := range []string{"expected-session-secret", "provided-session-secret"} {
if strings.Contains(logged, secret) {
t.Fatalf("log contains session value %q: %s", secret, logged)
}
}
}
+65
View File
@@ -0,0 +1,65 @@
//go:build darwin
package app
import (
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
)
func ShowFatalError(err error) {
if err == nil {
return
}
message := strings.ReplaceAll(err.Error(), `"`, `\"`)
_ = exec.Command("osascript", "-e", `display alert "BD2 Client Studio" message "`+message+`" as critical`).Run()
}
func openBrowser(url string) error {
return exec.Command("open", url).Start()
}
func browseForGameDirectory(language string) (string, error) {
prompt := "Select the Brown Dust II.app bundle or its parent folder"
if language == "zh-CN" {
prompt = "选择 Brown Dust II.app 或其所在文件夹"
}
return macDirectoryPicker(prompt)
}
func browseForResourceDirectory(language string) (string, error) {
prompt := "Select the CDN directory containing ServerData and GameData"
if language == "zh-CN" {
prompt = "选择包含 ServerData 和 GameData 的 CDN 目录"
}
return macDirectoryPicker(prompt)
}
func macDirectoryPicker(prompt string) (string, error) {
prompt = strings.ReplaceAll(prompt, `"`, `\"`)
command := exec.Command("osascript", "-e", `POSIX path of (choose folder with prompt "`+prompt+`")`)
output, err := command.Output()
if err != nil {
var exitErr *exec.ExitError
if errors.As(err, &exitErr) && exitErr.ExitCode() == 1 {
return "", nil
}
return "", fmt.Errorf("open macOS directory picker: %w", err)
}
return strings.TrimSpace(string(output)), nil
}
func launchGame(target string) error {
info, err := os.Stat(target)
if err != nil || !info.IsDir() || !strings.EqualFold(filepath.Ext(target), ".app") {
return fmt.Errorf("invalid macOS application bundle %q", target)
}
if err := exec.Command("pgrep", "-x", "BrownDust II").Run(); err == nil {
return errGameAlreadyRunning
}
return exec.Command("open", target).Start()
}
+36
View File
@@ -0,0 +1,36 @@
//go:build !windows && !darwin
package app
import (
"errors"
"fmt"
"os"
"os/exec"
"runtime"
)
func ShowFatalError(err error) {
if err != nil {
_, _ = fmt.Fprintln(os.Stderr, "BD2 Client Studio:", err)
}
}
func openBrowser(url string) error {
if runtime.GOOS == "darwin" {
return exec.Command("open", url).Start()
}
return exec.Command("xdg-open", url).Start()
}
func browseForGameDirectory(string) (string, error) {
return "", errors.New("the native directory picker is unavailable on this platform; enter the Windows client directory manually")
}
func browseForResourceDirectory(string) (string, error) {
return "", errors.New("the native directory picker is unavailable on this platform; enter the resource directory manually")
}
func launchGame(string) error {
return errors.New("the Brown Dust II client is not supported on Linux")
}
+182
View File
@@ -0,0 +1,182 @@
//go:build windows
package app
import (
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
"syscall"
"time"
"unsafe"
"golang.org/x/sys/windows"
)
var (
platformUser32DLL = syscall.NewLazyDLL("user32.dll")
messageBoxW = platformUser32DLL.NewProc("MessageBoxW")
enumWindowsProc = platformUser32DLL.NewProc("EnumWindows")
getWindowThreadProcessIDProc = platformUser32DLL.NewProc("GetWindowThreadProcessId")
isWindowVisibleProc = platformUser32DLL.NewProc("IsWindowVisible")
showWindowAsyncProc = platformUser32DLL.NewProc("ShowWindowAsync")
setForegroundWindowProc = platformUser32DLL.NewProc("SetForegroundWindow")
)
// ShowFatalError keeps startup failures visible even though the release
// executable uses the Windows GUI subsystem and therefore has no console.
func ShowFatalError(err error) {
if err == nil {
return
}
message, conversionErr := syscall.UTF16PtrFromString(fmt.Sprintf(
"BD2 Client Studio could not start:\n\n%s\n\nSee the logs directory next to bd2client.exe for details.", err,
))
if conversionErr != nil {
return
}
title, conversionErr := syscall.UTF16PtrFromString("BD2 Client Studio")
if conversionErr != nil {
return
}
messageBoxW.Call(0, uintptr(unsafe.Pointer(message)), uintptr(unsafe.Pointer(title)), 0x10)
}
// CREATE_NO_WINDOW prevents console-subsystem helpers such as powershell.exe
// from allocating a visible console when bd2client is built as a Windows GUI
// executable. HideWindow also covers helpers that elect to create a window
// despite inheriting no console from the parent process.
const createNoWindow = 0x08000000
func hiddenCommand(name string, args ...string) *exec.Cmd {
command := exec.Command(name, args...)
command.SysProcAttr = &syscall.SysProcAttr{
HideWindow: true,
CreationFlags: createNoWindow,
}
return command
}
// visibleCommand suppresses a console allocation without hiding the GUI
// window created by the child process. It must be used for the game itself;
// hiddenCommand is reserved for background helper processes.
func visibleCommand(name string, args ...string) *exec.Cmd {
command := exec.Command(name, args...)
command.SysProcAttr = &syscall.SysProcAttr{CreationFlags: createNoWindow}
return command
}
func openBrowser(url string) error {
for _, edge := range edgeCandidates() {
if info, err := os.Stat(edge); err == nil && !info.IsDir() {
return hiddenCommand(edge, "--app="+url, "--window-size=1100,760", "--no-first-run").Start()
}
}
return hiddenCommand("rundll32.exe", "url.dll,FileProtocolHandler", url).Start()
}
func edgeCandidates() []string {
var candidates []string
if edge, err := exec.LookPath("msedge.exe"); err == nil {
candidates = append(candidates, edge)
}
for _, root := range []string{os.Getenv("ProgramFiles(x86)"), os.Getenv("ProgramFiles"), os.Getenv("LOCALAPPDATA")} {
if root != "" {
candidates = append(candidates, filepath.Join(root, "Microsoft", "Edge", "Application", "msedge.exe"))
}
}
return candidates
}
func browseForGameDirectory(language string) (string, error) {
title := "Select the Brown Dust II installation directory"
if language == "zh-CN" {
title = "选择 Brown Dust II 安装目录"
}
return browseForDirectory(title)
}
func browseForResourceDirectory(language string) (string, error) {
title := "Select the CDN directory containing ServerData and GameData"
if language == "zh-CN" {
title = "选择包含 ServerData 和 GameData 的 CDN 目录"
}
return browseForDirectory(title)
}
func launchGame(target string) error {
if processID, running, err := windowsExecutableProcessID(filepath.Base(target)); err != nil {
return err
} else if running {
if !activateProcessWindow(processID, 5*time.Second) {
return fmt.Errorf("Brown Dust II is running, but its window could not be restored")
}
return errGameAlreadyRunning
}
command := visibleCommand(target)
command.Dir = filepath.Dir(target)
if err := command.Start(); err != nil {
return err
}
// Unity creates the top-level window asynchronously. Best-effort foreground
// activation prevents the new window from opening behind Client Studio.
activateProcessWindow(uint32(command.Process.Pid), 15*time.Second)
return nil
}
func windowsExecutableProcessID(name string) (uint32, bool, error) {
snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0)
if err != nil {
return 0, false, err
}
defer windows.CloseHandle(snapshot)
entry := windows.ProcessEntry32{Size: uint32(unsafe.Sizeof(windows.ProcessEntry32{}))}
if err := windows.Process32First(snapshot, &entry); err != nil {
return 0, false, err
}
for {
if strings.EqualFold(windows.UTF16ToString(entry.ExeFile[:]), name) {
return entry.ProcessID, true, nil
}
if err := windows.Process32Next(snapshot, &entry); err != nil {
if err == windows.ERROR_NO_MORE_FILES {
return 0, false, nil
}
return 0, false, err
}
}
}
func activateProcessWindow(processID uint32, timeout time.Duration) bool {
deadline := time.Now().Add(timeout)
for {
if window := topLevelWindowForProcess(processID); window != 0 {
const swRestore = 9
showWindowAsyncProc.Call(window, swRestore)
setForegroundWindowProc.Call(window)
return true
}
if time.Now().After(deadline) {
return false
}
time.Sleep(100 * time.Millisecond)
}
}
func topLevelWindowForProcess(processID uint32) uintptr {
var found uintptr
callback := syscall.NewCallback(func(window uintptr, _ uintptr) uintptr {
var owner uint32
getWindowThreadProcessIDProc.Call(window, uintptr(unsafe.Pointer(&owner)))
visible, _, _ := isWindowVisibleProc.Call(window)
if owner == processID && visible != 0 {
found = window
return 0
}
return 1
})
enumWindowsProc.Call(callback, 0)
return found
}
@@ -0,0 +1,40 @@
//go:build windows
package app
import "testing"
func TestHiddenCommandNeverAllocatesVisibleConsole(t *testing.T) {
command := hiddenCommand("powershell.exe", "-NoProfile")
if command.SysProcAttr == nil {
t.Fatal("hidden command has no Windows process attributes")
}
if !command.SysProcAttr.HideWindow {
t.Fatal("hidden command does not request a hidden window")
}
if command.SysProcAttr.CreationFlags&createNoWindow == 0 {
t.Fatalf("hidden command creation flags %#x omit CREATE_NO_WINDOW", command.SysProcAttr.CreationFlags)
}
}
func TestVisibleCommandDoesNotHideGUIWindow(t *testing.T) {
command := visibleCommand("Brown Dust II.exe")
if command.SysProcAttr == nil {
t.Fatal("visible command has no Windows process attributes")
}
if command.SysProcAttr.HideWindow {
t.Fatal("visible game command requests a hidden window")
}
if command.SysProcAttr.CreationFlags&createNoWindow == 0 {
t.Fatalf("visible command creation flags %#x omit CREATE_NO_WINDOW", command.SysProcAttr.CreationFlags)
}
}
func TestHRESULTFailureClassification(t *testing.T) {
if hresultFailed(0) || hresultFailed(1) {
t.Fatal("successful HRESULT classified as failure")
}
if !hresultFailed(errorCancelled) || !hresultFailed(0x80004005) {
t.Fatal("failed HRESULT classified as success")
}
}
File diff suppressed because one or more lines are too long
+163
View File
@@ -0,0 +1,163 @@
// Package config owns the client-side connection settings consumed by the
// standalone setup tool and the Local Identity plugin.
package config
import (
"encoding/json"
"errors"
"fmt"
"net"
"net/url"
"os"
"path/filepath"
"strings"
clientlayout "bd2server/internal/client/layout"
)
const (
SchemaVersion = 2
FileName = "bd2.client.json"
CDNOfficial CDNMode = "official"
CDNLocal CDNMode = "local"
CDNServer CDNMode = "server"
)
type CDNMode string
type Settings struct {
SchemaVersion int `json:"schema_version"`
ServerOrigin string `json:"server_origin"`
CDNMode CDNMode `json:"cdn_mode"`
LocalResourceDirectory string `json:"local_resource_directory,omitempty"`
}
func Path(gameDir string) string {
if installation, err := clientlayout.Resolve(gameDir); err == nil {
return filepath.Join(installation.Config, FileName)
}
return filepath.Join(filepath.Clean(gameDir), "BepInEx", "config", FileName)
}
func Normalize(in Settings) (Settings, error) {
origin, err := NormalizeOrigin(in.ServerOrigin)
if err != nil {
return Settings{}, err
}
localDirectory := strings.TrimSpace(in.LocalResourceDirectory)
switch in.CDNMode {
case CDNOfficial, CDNServer:
if localDirectory != "" {
return Settings{}, errors.New("client config: local_resource_directory is only valid in local mode")
}
case CDNLocal:
if localDirectory == "" {
return Settings{}, errors.New("client config: local mode requires local_resource_directory")
}
localDirectory, err = filepath.Abs(filepath.Clean(localDirectory))
if err != nil {
return Settings{}, fmt.Errorf("client config: resolve local resource directory: %w", err)
}
default:
return Settings{}, fmt.Errorf("client config: unsupported CDN mode %q", in.CDNMode)
}
return Settings{
SchemaVersion: SchemaVersion,
ServerOrigin: origin,
CDNMode: in.CDNMode,
LocalResourceDirectory: localDirectory,
}, nil
}
func NormalizeOrigin(raw string) (string, error) {
raw = strings.TrimSpace(raw)
parsed, err := url.Parse(raw)
if err != nil {
return "", fmt.Errorf("client config: parse server origin: %w", err)
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return "", errors.New("client config: server address must use http or https")
}
if parsed.Host == "" {
return "", errors.New("client config: server address must include a host")
}
if parsed.User != nil {
return "", errors.New("client config: credentials are not allowed in the server address")
}
if parsed.RawQuery != "" || parsed.Fragment != "" {
return "", errors.New("client config: server address cannot contain a query or fragment")
}
if parsed.Path != "" && parsed.Path != "/" {
return "", errors.New("client config: enter only the server origin, without /game or another path")
}
if parsed.Scheme == "http" && !isLoopback(parsed.Hostname()) {
return "", errors.New("client config: non-loopback servers must use https")
}
parsed.Path = ""
parsed.RawPath = ""
return strings.TrimSuffix(parsed.String(), "/"), nil
}
func isLoopback(host string) bool {
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func Save(gameDir string, in Settings) (Settings, error) {
settings, err := Normalize(in)
if err != nil {
return Settings{}, err
}
path := Path(gameDir)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return Settings{}, fmt.Errorf("client config: create config directory: %w", err)
}
data, err := json.MarshalIndent(settings, "", " ")
if err != nil {
return Settings{}, err
}
data = append(data, '\n')
temporary, err := os.CreateTemp(filepath.Dir(path), ".bd2-client-*.tmp")
if err != nil {
return Settings{}, fmt.Errorf("client config: create temporary config: %w", err)
}
temporaryPath := temporary.Name()
defer os.Remove(temporaryPath)
if err = temporary.Chmod(0o600); err == nil {
_, err = temporary.Write(data)
}
if err == nil {
err = temporary.Sync()
}
if closeErr := temporary.Close(); err == nil {
err = closeErr
}
if err != nil {
return Settings{}, fmt.Errorf("client config: stage config: %w", err)
}
if err := replaceFile(temporaryPath, path); err != nil {
return Settings{}, fmt.Errorf("client config: install config: %w", err)
}
return settings, nil
}
func Load(gameDir string) (Settings, error) {
data, err := os.ReadFile(Path(gameDir))
if err != nil {
return Settings{}, err
}
var settings Settings
decoder := json.NewDecoder(strings.NewReader(string(data)))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&settings); err != nil {
return Settings{}, fmt.Errorf("client config: decode: %w", err)
}
if settings.SchemaVersion != SchemaVersion {
return Settings{}, fmt.Errorf("client config: unsupported schema_version %d", settings.SchemaVersion)
}
return Normalize(settings)
}
+98
View File
@@ -0,0 +1,98 @@
package config
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestNormalize(t *testing.T) {
got, err := Normalize(Settings{ServerOrigin: " https://example.com:8443/ ", CDNMode: CDNServer})
if err != nil {
t.Fatal(err)
}
if got.SchemaVersion != SchemaVersion || got.ServerOrigin != "https://example.com:8443" || got.CDNMode != CDNServer {
t.Fatalf("normalized=%+v", got)
}
for _, bad := range []string{"example.com", "ftp://example.com", "http://192.168.1.8:8080", "https://u:p@example.com", "https://example.com/game/", "https://example.com?q=1"} {
if _, err := Normalize(Settings{ServerOrigin: bad, CDNMode: CDNOfficial}); err == nil {
t.Errorf("accepted origin %q", bad)
}
}
localRoot := t.TempDir()
local, err := Normalize(Settings{ServerOrigin: "http://127.0.0.1:8080", CDNMode: CDNLocal, LocalResourceDirectory: localRoot})
if err != nil || local.LocalResourceDirectory != localRoot {
t.Fatalf("local=%+v err=%v", local, err)
}
if _, err := Normalize(Settings{ServerOrigin: "http://127.0.0.1:8080", CDNMode: CDNLocal}); err == nil {
t.Fatal("accepted local mode without a resource directory")
}
if _, err := Normalize(Settings{ServerOrigin: "http://127.0.0.1:8080", CDNMode: CDNOfficial, LocalResourceDirectory: localRoot}); err == nil {
t.Fatal("accepted a local resource directory in official mode")
}
}
func TestSaveLoad(t *testing.T) {
dir := t.TempDir()
want := Settings{ServerOrigin: "http://127.0.0.1:8080", CDNMode: CDNLocal, LocalResourceDirectory: t.TempDir()}
if _, err := Save(dir, want); err != nil {
t.Fatal(err)
}
got, err := Load(dir)
if err != nil {
t.Fatal(err)
}
if got.ServerOrigin != want.ServerOrigin || got.CDNMode != want.CDNMode || got.SchemaVersion != SchemaVersion || got.LocalResourceDirectory != want.LocalResourceDirectory {
t.Fatalf("loaded=%+v", got)
}
data, err := os.ReadFile(Path(dir))
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(data), "secret") || strings.Contains(string(data), "token") {
t.Fatalf("client config unexpectedly stores a credential: %s", data)
}
updated := Settings{ServerOrigin: "https://friends.example:8443", CDNMode: CDNServer}
if _, err := Save(dir, updated); err != nil {
t.Fatalf("replace config: %v", err)
}
got, err = Load(dir)
if err != nil || got.ServerOrigin != updated.ServerOrigin || got.CDNMode != updated.CDNMode {
t.Fatalf("replaced=%+v err=%v", got, err)
}
}
func TestSaveOmitsLocalDirectoryOutsideLocalMode(t *testing.T) {
dir := t.TempDir()
if _, err := Save(dir, Settings{ServerOrigin: "https://example.com", CDNMode: CDNOfficial}); err != nil {
t.Fatal(err)
}
data, err := os.ReadFile(Path(dir))
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(data), "local_resource_directory") {
t.Fatalf("official config contains local directory field: %s", data)
}
}
func TestPathUsesMacAppSiblingBepInEx(t *testing.T) {
parent := t.TempDir()
app := filepath.Join(parent, "BrownDust II.app")
for _, path := range []string{
filepath.Join(app, "Contents", "MacOS", "BrownDust II"),
filepath.Join(app, "Contents", "Resources", "Data", "resources.assets"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte("test"), 0o700); err != nil {
t.Fatal(err)
}
}
want := filepath.Join(parent, "BepInEx", "config", FileName)
if got := Path(app); got != want {
t.Fatalf("Path()=%q want=%q", got, want)
}
}
+96
View File
@@ -0,0 +1,96 @@
package config
import (
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
)
const preferencesSchemaVersion = 1
type Preferences struct {
SchemaVersion int `json:"schema_version"`
GameDirectory string `json:"game_directory"`
}
func PreferencesPath() (string, error) {
root, err := os.UserConfigDir()
if err != nil {
return "", fmt.Errorf("locate user configuration directory: %w", err)
}
return filepath.Join(root, "BD2 Client Studio", "preferences.json"), nil
}
func LoadPreferences() (Preferences, error) {
path, err := PreferencesPath()
if err != nil {
return Preferences{}, err
}
data, err := os.ReadFile(path)
if errors.Is(err, os.ErrNotExist) {
return Preferences{}, nil
}
if err != nil {
return Preferences{}, fmt.Errorf("read client preferences: %w", err)
}
decoder := json.NewDecoder(strings.NewReader(string(data)))
decoder.DisallowUnknownFields()
var preferences Preferences
if err := decoder.Decode(&preferences); err != nil {
return Preferences{}, fmt.Errorf("decode client preferences: %w", err)
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return Preferences{}, errors.New("client preferences must contain exactly one JSON object")
}
if preferences.SchemaVersion != preferencesSchemaVersion || strings.TrimSpace(preferences.GameDirectory) == "" {
return Preferences{}, errors.New("client preferences are invalid")
}
preferences.GameDirectory = filepath.Clean(preferences.GameDirectory)
return preferences, nil
}
func SavePreferences(gameDirectory string) error {
abs, err := filepath.Abs(filepath.Clean(strings.TrimSpace(gameDirectory)))
if err != nil || strings.TrimSpace(gameDirectory) == "" {
return errors.New("client preferences require a valid game directory")
}
path, err := PreferencesPath()
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
return fmt.Errorf("create client preferences directory: %w", err)
}
data, err := json.MarshalIndent(Preferences{SchemaVersion: preferencesSchemaVersion, GameDirectory: abs}, "", " ")
if err != nil {
return err
}
data = append(data, '\n')
temporary, err := os.CreateTemp(filepath.Dir(path), ".preferences-*.tmp")
if err != nil {
return fmt.Errorf("stage client preferences: %w", err)
}
temporaryPath := temporary.Name()
defer os.Remove(temporaryPath)
if err = temporary.Chmod(0o600); err == nil {
_, err = temporary.Write(data)
}
if err == nil {
err = temporary.Sync()
}
if closeErr := temporary.Close(); err == nil {
err = closeErr
}
if err != nil {
return fmt.Errorf("stage client preferences: %w", err)
}
if err := replaceFile(temporaryPath, path); err != nil {
return fmt.Errorf("install client preferences: %w", err)
}
return nil
}
@@ -0,0 +1,31 @@
package config
import (
"os"
"path/filepath"
"testing"
)
func TestPreferencesRoundTrip(t *testing.T) {
root := t.TempDir()
t.Setenv("APPDATA", root)
game := filepath.Join(root, "game")
if err := SavePreferences(game); err != nil {
t.Fatal(err)
}
got, err := LoadPreferences()
if err != nil {
t.Fatal(err)
}
want, _ := filepath.Abs(game)
if got.SchemaVersion != preferencesSchemaVersion || got.GameDirectory != want {
t.Fatalf("preferences=%+v", got)
}
path, err := PreferencesPath()
if err != nil {
t.Fatal(err)
}
if _, err := os.Stat(path); err != nil {
t.Fatal(err)
}
}
+70
View File
@@ -0,0 +1,70 @@
package config
import (
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"strings"
)
const ReleaseFileName = "versions.json"
var (
clientVersionPattern = regexp.MustCompile(`^[0-9]+\.[0-9]+\.[0-9]+$`)
resourceVersionPattern = regexp.MustCompile(`^[0-9]{14}$`)
)
// ReleaseVersions is the exact client/resource tuple supported by one
// bd2client distribution. The release package carries the authoritative
// versions.json next to bd2client.exe.
type ReleaseVersions struct {
ClientVersion string `json:"client_version"`
GameDataVersion string `json:"game_data_version"`
BundleVersion string `json:"bundle_version"`
SeedDirectory string `json:"seed_directory"`
Plugins struct {
LocalIdentity string `json:"local_identity"`
CaptureEnvironment string `json:"capture_environment"`
LoginUI string `json:"login_ui"`
} `json:"plugins"`
}
func LoadReleaseVersions(path string) (ReleaseVersions, error) {
data, err := os.ReadFile(filepath.Clean(path))
if err != nil {
return ReleaseVersions{}, fmt.Errorf("read client release versions: %w", err)
}
decoder := json.NewDecoder(strings.NewReader(string(data)))
decoder.DisallowUnknownFields()
var versions ReleaseVersions
if err := decoder.Decode(&versions); err != nil {
return ReleaseVersions{}, fmt.Errorf("decode client release versions: %w", err)
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return ReleaseVersions{}, errors.New("client release versions must contain exactly one JSON object")
}
if !clientVersionPattern.MatchString(versions.ClientVersion) {
return ReleaseVersions{}, fmt.Errorf("invalid client_version %q", versions.ClientVersion)
}
if !resourceVersionPattern.MatchString(versions.BundleVersion) || !resourceVersionPattern.MatchString(versions.GameDataVersion) {
return ReleaseVersions{}, errors.New("bundle_version and game_data_version must be 14-digit timestamps")
}
if versions.SeedDirectory == "" ||
versions.Plugins.LocalIdentity == "" || versions.Plugins.LoginUI == "" || versions.Plugins.CaptureEnvironment == "" {
return ReleaseVersions{}, errors.New("client release versions are incomplete")
}
return versions, nil
}
func ReleaseVersionsBesideExecutable() (ReleaseVersions, error) {
executable, err := os.Executable()
if err != nil {
return ReleaseVersions{}, fmt.Errorf("locate bd2client executable: %w", err)
}
return LoadReleaseVersions(filepath.Join(filepath.Dir(executable), ReleaseFileName))
}
+39
View File
@@ -0,0 +1,39 @@
package config
import (
"os"
"path/filepath"
"testing"
)
func TestLoadReleaseVersions(t *testing.T) {
path := filepath.Join(t.TempDir(), ReleaseFileName)
data := `{"client_version":"2.35.10","game_data_version":"20260923193640","bundle_version":"20260921135230","seed_directory":"go/seed/v2_35_10","plugins":{"local_identity":"0.6.0","capture_environment":"0.2.0","login_ui":"0.1.0"}}`
if err := os.WriteFile(path, []byte(data), 0o600); err != nil {
t.Fatal(err)
}
got, err := LoadReleaseVersions(path)
if err != nil {
t.Fatal(err)
}
if got.ClientVersion != "2.35.10" || got.BundleVersion != "20260921135230" || got.GameDataVersion != "20260923193640" {
t.Fatalf("versions=%+v", got)
}
}
func TestLoadReleaseVersionsRejectsUnknownAndTrailingData(t *testing.T) {
for name, data := range map[string]string{
"unknown": `{"client_version":"2.35.10","unknown":true}`,
"trailing": `{}` + `{}`,
} {
t.Run(name, func(t *testing.T) {
path := filepath.Join(t.TempDir(), ReleaseFileName)
if err := os.WriteFile(path, []byte(data), 0o600); err != nil {
t.Fatal(err)
}
if _, err := LoadReleaseVersions(path); err == nil {
t.Fatal("accepted invalid release versions")
}
})
}
}
@@ -1,6 +1,6 @@
//go:build !windows
package clientplugin
package config
import "os"
@@ -1,6 +1,6 @@
//go:build windows
package clientplugin
package config
import (
"os"
+63
View File
@@ -0,0 +1,63 @@
package introdb
import (
"crypto/aes"
"crypto/cipher"
"crypto/hmac"
"crypto/sha1"
"fmt"
)
const PageSize = 4096
var Header = []byte("SQLite format 3\x00")
func decryptPages(in []byte) ([]byte, error) { return cryptPages(in, false) }
func encryptPages(in []byte) ([]byte, error) { return cryptPages(in, true) }
func cryptPages(in []byte, encrypt bool) ([]byte, error) {
if len(in) == 0 || len(in)%PageSize != 0 {
return nil, fmt.Errorf("dbcrypt: database length %d is not a non-zero multiple of %d", len(in), PageSize)
}
block, err := aes.NewCipher(deriveKey())
if err != nil {
return nil, err
}
out := make([]byte, len(in))
for start := 0; start < len(in); start += PageSize {
var mode cipher.BlockMode = cipher.NewCBCEncrypter(block, Header)
if !encrypt {
mode = cipher.NewCBCDecrypter(block, Header)
}
mode.CryptBlocks(out[start:start+PageSize], in[start:start+PageSize])
}
return out, nil
}
func deriveKey() []byte {
password := []byte(fmt.Sprintf("%X", sha1.Sum([]byte("spdhdnlwmrpavmtm"))))
return pbkdf2SHA1(password, Header, 2010, 32)
}
func pbkdf2SHA1(password, salt []byte, iterations, length int) []byte {
var result []byte
for block := uint32(1); len(result) < length; block++ {
message := append(append([]byte{}, salt...), byte(block>>24), byte(block>>16), byte(block>>8), byte(block))
u := hmacSHA1(password, message)
t := append([]byte{}, u...)
for i := 1; i < iterations; i++ {
u = hmacSHA1(password, u)
for j := range t {
t[j] ^= u[j]
}
}
result = append(result, t...)
}
return result[:length]
}
func hmacSHA1(key, message []byte) []byte {
h := hmac.New(sha1.New, key)
_, _ = h.Write(message)
return h.Sum(nil)
}
@@ -9,16 +9,15 @@ import (
"fmt"
"io"
"os"
"path/filepath"
"bd2server/internal/dbcrypt"
clientlayout "bd2server/internal/client/layout"
)
const (
oldURL = "https://mt.bd2.pmang.cloud/"
)
var salt = dbcrypt.Header
var salt = Header
// Result describes a completed in-place client patch. BackupPath is the
// immutable pre-patch copy and is never overwritten by a later invocation.
@@ -29,6 +28,7 @@ type Result struct {
ObjectSize uint32
OldURL string
NewURL string
Changed bool
}
// VerifyResult is useful to patch-client's --verify mode and to diagnostics.
@@ -74,7 +74,7 @@ func PatchClient(gameDir, newURL string) (Result, error) {
return Result{}, findErr
}
if current == newURL {
return Result{assets, assets + ".bak", entry.pathID, entry.size, current, newURL}, nil
return Result{assets, assets + ".bak", entry.pathID, entry.size, current, newURL, false}, nil
}
return Result{}, fmt.Errorf("introdb: LIVE_URL is already %q, not the expected official URL", current)
}
@@ -112,7 +112,7 @@ func PatchClient(gameDir, newURL string) (Result, error) {
if err := atomicWrite(assets, b); err != nil {
return Result{}, err
}
return Result{assets, backup, entry.pathID, entry.size, oldURL, newURL}, nil
return Result{assets, backup, entry.pathID, entry.size, oldURL, newURL, true}, nil
}
// VerifyClient reads and decrypts the embedded Intro TextAsset. It verifies
@@ -152,7 +152,11 @@ func ResourcesPath(gameDir string) (string, error) {
if gameDir == "" {
return "", errors.New("introdb: empty game directory")
}
p := filepath.Join(gameDir, "BrownDust II_Data", "resources.assets")
installation, err := clientlayout.Resolve(gameDir)
if err != nil {
return "", fmt.Errorf("introdb: resolve game layout: %w", err)
}
p := installation.Resources
st, err := os.Stat(p)
if err != nil {
return "", fmt.Errorf("introdb: resources.assets not found at %q: %w", p, err)
@@ -164,10 +168,10 @@ func ResourcesPath(gameDir string) (string, error) {
}
// DecryptPages decrypts the game's independent 4096-byte AES-CBC pages.
func DecryptPages(in []byte) ([]byte, error) { return dbcrypt.DecryptPages(in) }
func DecryptPages(in []byte) ([]byte, error) { return decryptPages(in) }
// EncryptPages encrypts the game's independent 4096-byte AES-CBC pages.
func EncryptPages(in []byte) ([]byte, error) { return dbcrypt.EncryptPages(in) }
func EncryptPages(in []byte) ([]byte, error) { return encryptPages(in) }
func validateIntroDB(p []byte, expected string) error {
if !bytes.HasPrefix(p, salt) {
@@ -5,8 +5,6 @@ import (
"os"
"path/filepath"
"testing"
"bd2server/internal/dbcrypt"
)
func referenceClientDir(t *testing.T) string {
@@ -19,7 +17,7 @@ func referenceClientDir(t *testing.T) string {
}
func TestPagesRoundTrip(t *testing.T) {
p := make([]byte, dbcrypt.PageSize*2)
p := make([]byte, PageSize*2)
copy(p, salt)
for i := 16; i < len(p); i++ {
p[i] = byte(i * 31)
@@ -37,7 +35,7 @@ func TestPagesRoundTrip(t *testing.T) {
}
}
func TestPagesRejectPartialPage(t *testing.T) {
if _, err := DecryptPages(make([]byte, dbcrypt.PageSize-1)); err == nil {
if _, err := DecryptPages(make([]byte, PageSize-1)); err == nil {
t.Fatal("accepted partial page")
}
}
@@ -109,6 +107,9 @@ func TestPatchClientTransaction(t *testing.T) {
if r.OldURL != oldURL || r.NewURL != local || r.BackupPath != dst+".bak" {
t.Fatalf("unexpected patch result: %#v", r)
}
if !r.Changed {
t.Fatal("first patch was not reported as changed")
}
if _, err := os.Stat(r.BackupPath); err != nil {
t.Fatalf("backup missing: %v", err)
}
@@ -119,7 +120,11 @@ func TestPatchClientTransaction(t *testing.T) {
if v.URL != local {
t.Fatalf("LIVE_URL=%q, want %q", v.URL, local)
}
if _, err := PatchClient(tmp, local); err != nil {
repeated, err := PatchClient(tmp, local)
if err != nil {
t.Fatalf("idempotent patch: %v", err)
}
if repeated.Changed {
t.Fatal("idempotent patch was reported as changed")
}
}
+120
View File
@@ -0,0 +1,120 @@
// Package layout resolves the supported Windows and macOS Brown Dust II
// installation layouts without relying on the host running bd2client.
package layout
import (
"errors"
"fmt"
"os"
"path/filepath"
"runtime"
"strings"
)
type Kind string
const (
Windows Kind = "windows"
MacOS Kind = "macos"
)
type Installation struct {
Kind Kind
Selected string
Root string
Executable string
Data string
Resources string
Managers string
BepInEx string
Config string
Plugins string
Disabled string
}
func Resolve(selected string) (Installation, error) {
if strings.TrimSpace(selected) == "" {
return Installation{}, errors.New("select the Brown Dust II installation directory")
}
abs, err := filepath.Abs(strings.TrimSpace(selected))
if err != nil {
return Installation{}, fmt.Errorf("resolve game directory: %w", err)
}
abs = filepath.Clean(abs)
if installation, ok := windowsLayout(abs); ok {
return installation, nil
}
if installation, ok := macLayout(abs); ok {
return installation, nil
}
return Installation{}, errors.New("the selected directory is not a complete Brown Dust II Windows or macOS client")
}
func windowsLayout(root string) (Installation, bool) {
executable := filepath.Join(root, "BrownDust II.exe")
data := filepath.Join(root, "BrownDust II_Data")
if !regularFile(executable) || !regularFile(filepath.Join(data, "resources.assets")) {
return Installation{}, false
}
return newInstallation(Windows, root, root, executable, data, filepath.Join(root, "BepInEx")), true
}
func macLayout(selected string) (Installation, bool) {
candidates := []string{selected}
if !strings.EqualFold(filepath.Ext(selected), ".app") {
candidates = append(candidates, filepath.Join(selected, "BrownDust II.app"))
}
for _, app := range candidates {
contents := filepath.Join(app, "Contents")
executable := filepath.Join(contents, "MacOS", "BrownDust II")
data := filepath.Join(contents, "Resources", "Data")
if regularFile(executable) && regularFile(filepath.Join(data, "resources.assets")) {
// BepInEx Unix distributions are normally extracted beside the
// .app bundle. Also accept an installation placed inside Contents.
bepInEx := filepath.Join(filepath.Dir(app), "BepInEx")
insideBundle := filepath.Join(contents, "BepInEx")
if directoryExists(insideBundle) && !directoryExists(bepInEx) {
bepInEx = insideBundle
}
return newInstallation(MacOS, selected, app, executable, data, bepInEx), true
}
}
return Installation{}, false
}
func newInstallation(kind Kind, selected, root, executable, data, bepInEx string) Installation {
return Installation{
Kind: kind,
Selected: selected,
Root: root,
Executable: executable,
Data: data,
Resources: filepath.Join(data, "resources.assets"),
Managers: filepath.Join(data, "globalgamemanagers"),
BepInEx: bepInEx,
Config: filepath.Join(bepInEx, "config"),
Plugins: filepath.Join(bepInEx, "plugins"),
Disabled: filepath.Join(bepInEx, "disabled"),
}
}
func (i Installation) LaunchTarget() string {
if i.Kind == MacOS {
return i.Root
}
return i.Executable
}
func (i Installation) SupportedOnHost() bool {
return (i.Kind == Windows && runtime.GOOS == "windows") || (i.Kind == MacOS && runtime.GOOS == "darwin")
}
func regularFile(path string) bool {
info, err := os.Stat(path)
return err == nil && info.Mode().IsRegular()
}
func directoryExists(path string) bool {
info, err := os.Stat(path)
return err == nil && info.IsDir()
}
+40
View File
@@ -0,0 +1,40 @@
package layout
import (
"os"
"path/filepath"
"testing"
)
func writeFile(t *testing.T, path string) {
t.Helper()
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte("test"), 0o700); err != nil {
t.Fatal(err)
}
}
func TestResolveWindows(t *testing.T) {
root := t.TempDir()
writeFile(t, filepath.Join(root, "BrownDust II.exe"))
writeFile(t, filepath.Join(root, "BrownDust II_Data", "resources.assets"))
got, err := Resolve(root)
if err != nil || got.Kind != Windows || got.Resources != filepath.Join(root, "BrownDust II_Data", "resources.assets") || got.Plugins != filepath.Join(root, "BepInEx", "plugins") {
t.Fatalf("layout=%+v err=%v", got, err)
}
}
func TestResolveMacAppAndParent(t *testing.T) {
parent := t.TempDir()
app := filepath.Join(parent, "BrownDust II.app")
writeFile(t, filepath.Join(app, "Contents", "MacOS", "BrownDust II"))
writeFile(t, filepath.Join(app, "Contents", "Resources", "Data", "resources.assets"))
for _, selected := range []string{app, parent} {
got, err := Resolve(selected)
if err != nil || got.Kind != MacOS || got.Root != app || got.BepInEx != filepath.Join(parent, "BepInEx") || got.LaunchTarget() != app {
t.Fatalf("selected=%q layout=%+v err=%v", selected, got, err)
}
}
}
@@ -1,4 +1,4 @@
package clientplugin
package plugin
import (
"bytes"
@@ -8,55 +8,80 @@ import (
"io"
"os"
"path/filepath"
clientlayout "bd2server/internal/client/layout"
)
const (
FileName = "BD2LocalIdentity.dll"
BepInExReleasesURL = "https://github.com/BepInEx/BepInEx/releases"
)
type Spec struct {
fileName string
}
var (
LocalIdentity = Spec{fileName: "BD2LocalIdentity.dll"}
LoginUI = Spec{fileName: "BD2LoginUI.dll"}
)
func (s Spec) FileName() string { return s.fileName }
func (s Spec) validate() error {
if s.fileName == "" || filepath.Base(s.fileName) != s.fileName || filepath.Ext(s.fileName) != ".dll" {
return errors.New("clientplugin: invalid plugin specification")
}
return nil
}
type Result struct {
Destination string
Changed bool
}
func ResolvePackaged(explicit string) (string, error) {
func ResolvePackaged(spec Spec, explicit string) (string, error) {
if err := spec.validate(); err != nil {
return "", err
}
if explicit != "" {
return filepath.Clean(explicit), nil
}
executable, err := os.Executable()
if err != nil {
return "", fmt.Errorf("clientplugin: resolve server executable: %w", err)
return "", fmt.Errorf("clientplugin: resolve client tool executable: %w", err)
}
return filepath.Join(filepath.Dir(executable), "plugins", FileName), nil
return filepath.Join(filepath.Dir(executable), "plugins", spec.fileName), nil
}
// Install verifies that the user installed BepInEx, then atomically stages the
// packaged local-identity plugin into its plugins directory. It never installs
// or downloads BepInEx itself.
func Install(gameDir, source string) (Result, error) {
// packaged plugin into its plugins directory. It never installs or downloads
// BepInEx itself.
func Install(spec Spec, gameDir, source string) (Result, error) {
if err := spec.validate(); err != nil {
return Result{}, err
}
if gameDir == "" || source == "" {
return Result{}, errors.New("clientplugin: game directory and plugin source are required")
}
gameDir = filepath.Clean(gameDir)
source = filepath.Clean(source)
gameExecutable := filepath.Join(gameDir, "BrownDust II.exe")
if info, err := os.Stat(gameExecutable); err != nil || info.IsDir() {
return Result{}, fmt.Errorf("clientplugin: game executable is unavailable at %q", gameExecutable)
installation, err := clientlayout.Resolve(gameDir)
if err != nil {
return Result{}, fmt.Errorf("clientplugin: resolve game layout: %w", err)
}
bepInEx := filepath.Join(gameDir, "BepInEx", "core", "BepInEx.dll")
bepInEx := filepath.Join(installation.BepInEx, "core", "BepInEx.dll")
if info, err := os.Stat(bepInEx); err != nil || info.IsDir() {
return Result{}, fmt.Errorf("clientplugin: BepInEx is not installed; install it manually from %s, then restart the server; %s was not copied", BepInExReleasesURL, FileName)
return Result{}, fmt.Errorf("clientplugin: BepInEx is not installed; install it manually from %s, then run the client tool again; %s was not copied", BepInExReleasesURL, spec.fileName)
}
sourceData, err := os.ReadFile(source)
if err != nil {
return Result{}, fmt.Errorf("clientplugin: read packaged %s: %w", FileName, err)
return Result{}, fmt.Errorf("clientplugin: read packaged %s: %w", spec.fileName, err)
}
if len(sourceData) == 0 {
return Result{}, fmt.Errorf("clientplugin: packaged %s is empty", FileName)
return Result{}, fmt.Errorf("clientplugin: packaged %s is empty", spec.fileName)
}
pluginDir := filepath.Join(gameDir, "BepInEx", "plugins")
destination := filepath.Join(pluginDir, FileName)
pluginDir := installation.Plugins
destination := filepath.Join(pluginDir, spec.fileName)
if installed, err := os.ReadFile(destination); err == nil {
if bytes.Equal(hash(installed), hash(sourceData)) {
return Result{Destination: destination}, nil
@@ -67,7 +92,7 @@ func Install(gameDir, source string) (Result, error) {
if err := os.MkdirAll(pluginDir, 0o755); err != nil {
return Result{}, fmt.Errorf("clientplugin: create plugin directory: %w", err)
}
temporary, err := os.CreateTemp(pluginDir, ".BD2LocalIdentity-*.tmp")
temporary, err := os.CreateTemp(pluginDir, "."+spec.fileName+"-*.tmp")
if err != nil {
return Result{}, fmt.Errorf("clientplugin: create temporary plugin: %w", err)
}
+113
View File
@@ -0,0 +1,113 @@
package plugin
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestInstallRequiresBepInExWithoutCopyingPlugin(t *testing.T) {
for _, spec := range []Spec{LocalIdentity, LoginUI} {
t.Run(spec.FileName(), func(t *testing.T) {
gameDir := t.TempDir()
for path, data := range map[string][]byte{
filepath.Join(gameDir, "BrownDust II.exe"): []byte("game"),
filepath.Join(gameDir, "BrownDust II_Data", "resources.assets"): []byte("assets"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
source := filepath.Join(t.TempDir(), spec.FileName())
if err := os.WriteFile(source, []byte("plugin"), 0o600); err != nil {
t.Fatal(err)
}
_, err := Install(spec, gameDir, source)
if err == nil || !strings.Contains(err.Error(), BepInExReleasesURL) {
t.Fatalf("missing BepInEx error=%v", err)
}
if _, statErr := os.Stat(filepath.Join(gameDir, "BepInEx", "plugins", spec.FileName())); !os.IsNotExist(statErr) {
t.Fatalf("plugin was copied without BepInEx: %v", statErr)
}
})
}
}
func TestInstallCopiesUpdatesAndSkipsIdenticalPlugin(t *testing.T) {
for _, spec := range []Spec{LocalIdentity, LoginUI} {
t.Run(spec.FileName(), func(t *testing.T) {
gameDir := t.TempDir()
for path, data := range map[string][]byte{
filepath.Join(gameDir, "BrownDust II.exe"): []byte("game"),
filepath.Join(gameDir, "BrownDust II_Data", "resources.assets"): []byte("assets"),
filepath.Join(gameDir, "BepInEx", "core", "BepInEx.dll"): []byte("bepinex"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
source := filepath.Join(t.TempDir(), spec.FileName())
if err := os.WriteFile(source, []byte("v1"), 0o600); err != nil {
t.Fatal(err)
}
first, err := Install(spec, gameDir, source)
if err != nil || !first.Changed {
t.Fatalf("first install=%+v err=%v", first, err)
}
second, err := Install(spec, gameDir, source)
if err != nil || second.Changed {
t.Fatalf("idempotent install=%+v err=%v", second, err)
}
if err := os.WriteFile(source, []byte("v2"), 0o600); err != nil {
t.Fatal(err)
}
third, err := Install(spec, gameDir, source)
if err != nil || !third.Changed {
t.Fatalf("update=%+v err=%v", third, err)
}
got, err := os.ReadFile(third.Destination)
if err != nil || string(got) != "v2" {
t.Fatalf("installed=%q err=%v", got, err)
}
})
}
}
func TestInstallKeepsPluginsSeparate(t *testing.T) {
gameDir := t.TempDir()
for path, data := range map[string][]byte{
filepath.Join(gameDir, "BrownDust II.exe"): []byte("game"),
filepath.Join(gameDir, "BrownDust II_Data", "resources.assets"): []byte("assets"),
filepath.Join(gameDir, "BepInEx", "core", "BepInEx.dll"): []byte("bepinex"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
for _, spec := range []Spec{LocalIdentity, LoginUI} {
source := filepath.Join(t.TempDir(), spec.FileName())
if err := os.WriteFile(source, []byte(spec.FileName()), 0o600); err != nil {
t.Fatal(err)
}
if _, err := Install(spec, gameDir, source); err != nil {
t.Fatal(err)
}
}
for _, spec := range []Spec{LocalIdentity, LoginUI} {
path := filepath.Join(gameDir, "BepInEx", "plugins", spec.FileName())
data, err := os.ReadFile(path)
if err != nil || string(data) != spec.FileName() {
t.Fatalf("%s=%q err=%v", spec.FileName(), data, err)
}
}
}
@@ -0,0 +1,9 @@
//go:build !windows
package plugin
import "os"
func replaceFile(source, destination string) error {
return os.Rename(source, destination)
}
@@ -0,0 +1,34 @@
//go:build windows
package plugin
import (
"os"
"syscall"
"unsafe"
)
var moveFileEx = syscall.NewLazyDLL("kernel32.dll").NewProc("MoveFileExW")
func replaceFile(source, destination string) error {
sourcePtr, err := syscall.UTF16PtrFromString(source)
if err != nil {
return err
}
destinationPtr, err := syscall.UTF16PtrFromString(destination)
if err != nil {
return err
}
result, _, callErr := moveFileEx.Call(
uintptr(unsafe.Pointer(sourcePtr)),
uintptr(unsafe.Pointer(destinationPtr)),
0x1|0x8, // MOVEFILE_REPLACE_EXISTING | MOVEFILE_WRITE_THROUGH
)
if result == 0 {
if callErr != syscall.Errno(0) {
return callErr
}
return os.ErrInvalid
}
return nil
}
+323
View File
@@ -0,0 +1,323 @@
// Package setup implements the filesystem and network operations exposed by
// bd2client. It contains no server runtime dependencies.
package setup
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"os"
"path/filepath"
"regexp"
"strings"
"time"
clientconfig "bd2server/internal/client/config"
"bd2server/internal/client/introdb"
clientlayout "bd2server/internal/client/layout"
"bd2server/internal/client/plugin"
)
const PatchPlaceholder = "http://127.0.0.1:8080/game/"
type GameStatus struct {
GameDirectory string `json:"game_directory"`
ClientVersion string `json:"client_version"`
Executable bool `json:"executable"`
Resources bool `json:"resources"`
BepInEx bool `json:"bepinex"`
Config bool `json:"config"`
PatchedURL string `json:"patched_url,omitempty"`
}
var clientVersionPattern = regexp.MustCompile(`(?:^|\x00)([0-9]{1,2}\.[0-9]{1,2}\.[0-9]{1,3})(?:\x00)`)
type ResourcePolicy struct {
Mode clientconfig.CDNMode `json:"mode"`
ServerDataURL string `json:"server_data_url"`
GameDataURL string `json:"game_data_url"`
BundleVersion string `json:"bundle_version"`
GameDataVersion string `json:"game_data_version"`
LocalDirectory string `json:"local_directory,omitempty"`
}
type InstallResult struct {
LocalIdentity plugin.Result `json:"local_identity"`
LoginUI plugin.Result `json:"login_ui"`
}
func Inspect(gameDir string, versions clientconfig.ReleaseVersions) (GameStatus, error) {
if strings.TrimSpace(gameDir) == "" {
return GameStatus{}, errors.New("select the Brown Dust II installation directory")
}
installation, err := clientlayout.Resolve(gameDir)
if err != nil {
return GameStatus{}, err
}
status := GameStatus{GameDirectory: installation.Root}
status.Executable = regularFile(installation.Executable)
status.Resources = regularFile(installation.Resources)
status.BepInEx = regularFile(filepath.Join(installation.BepInEx, "core", "BepInEx.dll"))
status.Config = regularFile(filepath.Join(installation.Config, clientconfig.FileName))
status.ClientVersion, err = detectClientVersion(installation.Managers)
if err != nil {
return status, err
}
if status.ClientVersion != versions.ClientVersion {
return status, fmt.Errorf("unsupported Brown Dust II client version %s; this bd2client release requires %s", status.ClientVersion, versions.ClientVersion)
}
if verified, verifyErr := introdb.VerifyClient(installation.Root); verifyErr == nil {
status.PatchedURL = verified.URL
}
return status, nil
}
func SaveSettings(gameDir string, settings clientconfig.Settings, versions clientconfig.ReleaseVersions) (clientconfig.Settings, error) {
if _, err := Inspect(gameDir, versions); err != nil {
return clientconfig.Settings{}, err
}
return clientconfig.Save(gameDir, settings)
}
func Patch(gameDir string, settings clientconfig.Settings, versions clientconfig.ReleaseVersions) (introdb.Result, error) {
if _, err := SaveSettings(gameDir, settings, versions); err != nil {
return introdb.Result{}, err
}
result, err := introdb.PatchClient(gameDir, PatchPlaceholder)
if err != nil {
return introdb.Result{}, err
}
if _, err := introdb.VerifyClient(gameDir); err != nil {
return introdb.Result{}, fmt.Errorf("verify patched client resources: %w", err)
}
if _, err := disableLegacyPlugin(gameDir); err != nil {
return introdb.Result{}, err
}
return result, nil
}
func InstallPlugins(
gameDir string,
settings clientconfig.Settings,
versions clientconfig.ReleaseVersions,
localIdentitySource string,
loginUISource string,
) (InstallResult, error) {
status, err := Inspect(gameDir, versions)
if err != nil {
return InstallResult{}, err
}
if !status.BepInEx {
return InstallResult{}, fmt.Errorf("BepInEx is not installed; install it from %s before installing the plugins", plugin.BepInExReleasesURL)
}
if _, err := clientconfig.Save(gameDir, settings); err != nil {
return InstallResult{}, err
}
localSource, err := plugin.ResolvePackaged(plugin.LocalIdentity, localIdentitySource)
if err != nil {
return InstallResult{}, err
}
loginSource, err := plugin.ResolvePackaged(plugin.LoginUI, loginUISource)
if err != nil {
return InstallResult{}, err
}
local, err := plugin.Install(plugin.LocalIdentity, gameDir, localSource)
if err != nil {
return InstallResult{}, err
}
login, err := plugin.Install(plugin.LoginUI, gameDir, loginSource)
if err != nil {
return InstallResult{}, err
}
return InstallResult{LocalIdentity: local, LoginUI: login}, nil
}
func FetchResourcePolicy(ctx context.Context, client *http.Client, settings clientconfig.Settings, versions clientconfig.ReleaseVersions) (ResourcePolicy, error) {
normalized, err := clientconfig.Normalize(settings)
if err != nil {
return ResourcePolicy{}, err
}
if normalized.CDNMode == clientconfig.CDNOfficial {
return ResourcePolicy{Mode: clientconfig.CDNOfficial}, nil
}
if normalized.CDNMode == clientconfig.CDNLocal {
root, err := inspectLocalResourceDirectory(normalized.LocalResourceDirectory, versions)
if err != nil {
return ResourcePolicy{}, err
}
return ResourcePolicy{
Mode: clientconfig.CDNLocal,
ServerDataURL: localResourceURL(filepath.Join(root, "ServerData")),
GameDataURL: localResourceURL(filepath.Join(root, "GameData")),
BundleVersion: versions.BundleVersion,
GameDataVersion: versions.GameDataVersion,
LocalDirectory: root,
}, nil
}
if client == nil {
client = &http.Client{Timeout: 10 * time.Second}
}
body, err := json.Marshal(map[string]clientconfig.CDNMode{"cdn_mode": normalized.CDNMode})
if err != nil {
return ResourcePolicy{}, err
}
endpoint := normalized.ServerOrigin + "/client/resources"
request, err := http.NewRequestWithContext(ctx, http.MethodPut, endpoint, bytes.NewReader(body))
if err != nil {
return ResourcePolicy{}, err
}
request.Header.Set("Content-Type", "application/json")
request.Header.Set("Accept", "application/json")
response, err := client.Do(request)
if err != nil {
return ResourcePolicy{}, fmt.Errorf("request server resource policy: %w", err)
}
defer response.Body.Close()
limited := io.LimitReader(response.Body, 64<<10)
responseBody, err := io.ReadAll(limited)
if err != nil {
return ResourcePolicy{}, fmt.Errorf("read server resource policy response: %w", err)
}
if response.StatusCode != http.StatusOK {
message := strings.TrimSpace(string(responseBody))
if len(message) > 300 {
message = message[:300]
}
return ResourcePolicy{}, fmt.Errorf("server rejected CDN mode %s (HTTP %d): %s", normalized.CDNMode, response.StatusCode, message)
}
var policy ResourcePolicy
if err := json.Unmarshal(responseBody, &policy); err != nil {
return ResourcePolicy{}, fmt.Errorf("decode server resource policy: %w", err)
}
if policy.Mode != normalized.CDNMode {
return ResourcePolicy{}, fmt.Errorf("server returned CDN mode %q, expected %q", policy.Mode, normalized.CDNMode)
}
if err := validatePublicURL("ServerData", policy.ServerDataURL); err != nil {
return ResourcePolicy{}, err
}
if err := validatePublicURL("GameData", policy.GameDataURL); err != nil {
return ResourcePolicy{}, err
}
if policy.BundleVersion == "" || policy.GameDataVersion == "" {
return ResourcePolicy{}, errors.New("server resource policy is missing bundle_version or game_data_version")
}
if policy.BundleVersion != versions.BundleVersion || policy.GameDataVersion != versions.GameDataVersion {
return ResourcePolicy{}, fmt.Errorf(
"server resource versions do not match this client release: bundle=%s (want %s), GameData=%s (want %s)",
policy.BundleVersion, versions.BundleVersion, policy.GameDataVersion, versions.GameDataVersion,
)
}
return policy, nil
}
func inspectLocalResourceDirectory(raw string, versions clientconfig.ReleaseVersions) (string, error) {
root, err := filepath.Abs(filepath.Clean(strings.TrimSpace(raw)))
if err != nil {
return "", fmt.Errorf("resolve local resource directory: %w", err)
}
for _, relative := range []string{
filepath.Join("ServerData", "StandaloneWindows64", "HD", versions.BundleVersion, "catalog_alpha.json"),
filepath.Join("ServerData", "StandaloneWindows64", "HD", versions.BundleVersion, "catalog_alpha.hash"),
filepath.Join("GameData", versions.GameDataVersion, "release", "common-dbdata.info"),
filepath.Join("GameData", versions.GameDataVersion, "release", "common-dbdata.bin"),
} {
info, statErr := os.Stat(filepath.Join(root, relative))
if statErr != nil || !info.Mode().IsRegular() {
return "", fmt.Errorf("local resource directory is missing %s", relative)
}
}
return root, nil
}
func detectClientVersion(path string) (string, error) {
info, err := os.Stat(path)
if err != nil {
return "", fmt.Errorf("read Brown Dust II client version metadata: %w", err)
}
if !info.Mode().IsRegular() {
return "", errors.New("Brown Dust II client version metadata is not a regular file")
}
if info.Size() <= 0 || info.Size() > 64<<20 {
return "", fmt.Errorf("Brown Dust II client version metadata has an invalid size: %d", info.Size())
}
data, err := os.ReadFile(path)
if err != nil {
return "", fmt.Errorf("read Brown Dust II client version metadata: %w", err)
}
matches := clientVersionPattern.FindAllSubmatch(data, -1)
versions := make(map[string]struct{})
for _, match := range matches {
versions[string(match[1])] = struct{}{}
}
if len(versions) != 1 {
return "", fmt.Errorf("could not identify one unambiguous Brown Dust II client version in %s", path)
}
for version := range versions {
return version, nil
}
panic("unreachable")
}
func localResourceURL(path string) string {
slashed := filepath.ToSlash(filepath.Clean(path))
if filepath.VolumeName(path) != "" && !strings.HasPrefix(slashed, "/") {
slashed = "/" + slashed
}
return (&url.URL{Scheme: "file", Path: slashed}).String()
}
func validatePublicURL(name, raw string) error {
parsed, err := url.Parse(raw)
if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
return fmt.Errorf("server returned an invalid %s URL", name)
}
if parsed.Scheme == "http" && !resourceLoopback(parsed.Hostname()) {
return fmt.Errorf("server returned an invalid %s URL", name)
}
return nil
}
func resourceLoopback(host string) bool {
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func regularFile(path string) bool {
info, err := os.Stat(path)
return err == nil && !info.IsDir()
}
func disableLegacyPlugin(gameDir string) (string, error) {
installation, err := clientlayout.Resolve(gameDir)
if err != nil {
return "", err
}
source := filepath.Join(installation.Plugins, "PluginLocalRes.dll")
destination := filepath.Join(installation.Disabled, "PluginLocalRes.dll")
if _, err := os.Stat(source); errors.Is(err, os.ErrNotExist) {
return "", nil
} else if err != nil {
return "", fmt.Errorf("inspect legacy local resource plugin: %w", err)
}
if _, err := os.Stat(destination); err == nil {
return "", errors.New("the legacy local resource plugin exists in both active and disabled directories; remove one copy manually")
} else if !errors.Is(err, os.ErrNotExist) {
return "", err
}
if err := os.MkdirAll(filepath.Dir(destination), 0o755); err != nil {
return "", err
}
if err := os.Rename(source, destination); err != nil {
return "", fmt.Errorf("disable legacy local resource plugin: %w", err)
}
return destination, nil
}
+173
View File
@@ -0,0 +1,173 @@
package setup
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
clientconfig "bd2server/internal/client/config"
)
func testVersions() clientconfig.ReleaseVersions {
return clientconfig.ReleaseVersions{
ClientVersion: "2.35.10", BundleVersion: "20260921135230", GameDataVersion: "20260923193640",
}
}
func TestInspectRequiresGameFiles(t *testing.T) {
dir := t.TempDir()
if _, err := Inspect(dir, testVersions()); err == nil {
t.Fatal("accepted empty directory")
}
for path, data := range map[string][]byte{
filepath.Join(dir, "BrownDust II.exe"): []byte("exe"),
filepath.Join(dir, "BrownDust II_Data", "resources.assets"): []byte("assets"),
filepath.Join(dir, "BrownDust II_Data", "globalgamemanagers"): []byte("\x002.35.10\x00"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
status, err := Inspect(dir, testVersions())
if err != nil {
t.Fatal(err)
}
if !status.Executable || !status.Resources || status.BepInEx {
t.Fatalf("status=%+v", status)
}
}
func TestInspectRejectsUnsupportedClientVersion(t *testing.T) {
dir := t.TempDir()
for path, data := range map[string][]byte{
filepath.Join(dir, "BrownDust II.exe"): []byte("exe"),
filepath.Join(dir, "BrownDust II_Data", "resources.assets"): []byte("assets"),
filepath.Join(dir, "BrownDust II_Data", "globalgamemanagers"): []byte("\x002.36.0\x00"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
status, err := Inspect(dir, testVersions())
if err == nil || status.ClientVersion != "2.36.0" {
t.Fatalf("status=%+v err=%v", status, err)
}
}
func TestInspectMacApp(t *testing.T) {
parent := t.TempDir()
app := filepath.Join(parent, "BrownDust II.app")
for path, data := range map[string][]byte{
filepath.Join(app, "Contents", "MacOS", "BrownDust II"): []byte("binary"),
filepath.Join(app, "Contents", "Resources", "Data", "resources.assets"): []byte("assets"),
filepath.Join(app, "Contents", "Resources", "Data", "globalgamemanagers"): []byte("\x002.35.10\x00"),
filepath.Join(parent, "BepInEx", "core", "BepInEx.dll"): []byte("bepinex"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o700); err != nil {
t.Fatal(err)
}
}
status, err := Inspect(app, testVersions())
if err != nil || status.GameDirectory != app || status.ClientVersion != "2.35.10" || !status.BepInEx {
t.Fatalf("status=%+v err=%v", status, err)
}
}
func TestFetchResourcePolicyRejectsServerVersionMismatch(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_ = json.NewEncoder(w).Encode(ResourcePolicy{
Mode: clientconfig.CDNServer, ServerDataURL: "https://cdn.example/ServerData",
GameDataURL: "https://cdn.example/GameData", BundleVersion: "wrong", GameDataVersion: "wrong",
})
}))
defer server.Close()
_, err := FetchResourcePolicy(context.Background(), server.Client(), clientconfig.Settings{ServerOrigin: server.URL, CDNMode: clientconfig.CDNServer}, testVersions())
if err == nil {
t.Fatal("accepted mismatched server resource versions")
}
}
func TestFetchResourcePolicy(t *testing.T) {
var method string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
method = r.Method
var request map[string]string
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
t.Error(err)
}
if request["cdn_mode"] != "server" {
t.Errorf("request=%v", request)
}
w.Header().Set("Cache-Control", "no-store")
_ = json.NewEncoder(w).Encode(ResourcePolicy{
Mode: clientconfig.CDNServer,
ServerDataURL: "https://cdn.example/ServerData",
GameDataURL: "https://cdn.example/GameData",
BundleVersion: "20260921135230",
GameDataVersion: "20260923193640",
})
}))
defer server.Close()
policy, err := FetchResourcePolicy(context.Background(), server.Client(), clientconfig.Settings{
ServerOrigin: server.URL,
CDNMode: clientconfig.CDNServer,
}, testVersions())
if err != nil {
t.Fatal(err)
}
if method != http.MethodPut || policy.Mode != clientconfig.CDNServer || policy.BundleVersion != "20260921135230" {
t.Fatalf("method=%s policy=%+v", method, policy)
}
}
func TestLocalResourcesDoNotContactServer(t *testing.T) {
root := t.TempDir()
for _, relative := range []string{
filepath.Join("ServerData", "StandaloneWindows64", "HD", "20260921135230", "catalog_alpha.json"),
filepath.Join("ServerData", "StandaloneWindows64", "HD", "20260921135230", "catalog_alpha.hash"),
filepath.Join("GameData", "20260923193640", "release", "common-dbdata.info"),
filepath.Join("GameData", "20260923193640", "release", "common-dbdata.bin"),
} {
path := filepath.Join(root, relative)
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, []byte("resource"), 0o600); err != nil {
t.Fatal(err)
}
}
policy, err := FetchResourcePolicy(context.Background(), nil, clientconfig.Settings{
ServerOrigin: "http://127.0.0.1:8080",
CDNMode: clientconfig.CDNLocal,
LocalResourceDirectory: root,
}, testVersions())
if err != nil {
t.Fatal(err)
}
if policy.Mode != clientconfig.CDNLocal || policy.LocalDirectory != root || policy.ServerDataURL == "" || policy.GameDataURL == "" {
t.Fatalf("policy=%+v", policy)
}
}
func TestOfficialDoesNotContactServer(t *testing.T) {
policy, err := FetchResourcePolicy(context.Background(), nil, clientconfig.Settings{
ServerOrigin: "https://example.com",
CDNMode: clientconfig.CDNOfficial,
}, testVersions())
if err != nil || policy.Mode != clientconfig.CDNOfficial {
t.Fatalf("policy=%+v err=%v", policy, err)
}
}
-64
View File
@@ -1,64 +0,0 @@
package clientplugin
import (
"os"
"path/filepath"
"strings"
"testing"
)
func TestInstallRequiresBepInExWithoutCopyingPlugin(t *testing.T) {
gameDir := t.TempDir()
if err := os.WriteFile(filepath.Join(gameDir, "BrownDust II.exe"), []byte("game"), 0o600); err != nil {
t.Fatal(err)
}
source := filepath.Join(t.TempDir(), FileName)
if err := os.WriteFile(source, []byte("plugin"), 0o600); err != nil {
t.Fatal(err)
}
_, err := Install(gameDir, source)
if err == nil || !strings.Contains(err.Error(), BepInExReleasesURL) {
t.Fatalf("missing BepInEx error=%v", err)
}
if _, statErr := os.Stat(filepath.Join(gameDir, "BepInEx", "plugins", FileName)); !os.IsNotExist(statErr) {
t.Fatalf("plugin was copied without BepInEx: %v", statErr)
}
}
func TestInstallCopiesUpdatesAndSkipsIdenticalPlugin(t *testing.T) {
gameDir := t.TempDir()
for path, data := range map[string][]byte{
filepath.Join(gameDir, "BrownDust II.exe"): []byte("game"),
filepath.Join(gameDir, "BepInEx", "core", "BepInEx.dll"): []byte("bepinex"),
} {
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(path, data, 0o600); err != nil {
t.Fatal(err)
}
}
source := filepath.Join(t.TempDir(), FileName)
if err := os.WriteFile(source, []byte("v1"), 0o600); err != nil {
t.Fatal(err)
}
first, err := Install(gameDir, source)
if err != nil || !first.Changed {
t.Fatalf("first install=%+v err=%v", first, err)
}
second, err := Install(gameDir, source)
if err != nil || second.Changed {
t.Fatalf("idempotent install=%+v err=%v", second, err)
}
if err := os.WriteFile(source, []byte("v2"), 0o600); err != nil {
t.Fatal(err)
}
third, err := Install(gameDir, source)
if err != nil || !third.Changed {
t.Fatalf("update=%+v err=%v", third, err)
}
got, err := os.ReadFile(third.Destination)
if err != nil || string(got) != "v2" {
t.Fatalf("installed=%q err=%v", got, err)
}
}
@@ -12,12 +12,12 @@ import (
"path/filepath"
"time"
"bd2server/internal/cryptox"
"bd2server/internal/versionconfig"
"bd2server/internal/wire"
"bd2server/internal/server/cryptox"
"bd2server/internal/server/versionconfig"
"bd2server/internal/server/wire"
)
func ProtocolVersion() string { return versionconfig.Protocol() }
func StateVersion() string { return versionconfig.State() }
var (
ErrInvalidSeed = errors.New("account: invalid LoginUser seed")
@@ -7,15 +7,15 @@ import (
"testing"
"time"
"bd2server/internal/cryptox"
"bd2server/internal/wire"
"bd2server/internal/server/cryptox"
"bd2server/internal/server/wire"
)
func TestEncodeUsesFreshLocalKey(t *testing.T) {
user := wire.AppendVarint(nil, 1, 42)
user = wire.AppendString(user, 2, "Guest_42")
user = wire.AppendVarint(user, 5, 100)
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 11, UserInfo: user}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 11, UserInfo: user}
const local = "0123456789abcdef0123456789abcdef"
body, err := seed.Encode(local, time.UnixMilli(1234))
if err != nil {
@@ -48,7 +48,7 @@ func TestEncodeUsesFreshLocalKey(t *testing.T) {
}
func TestLoginValidatesEncryptedRequest(t *testing.T) {
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 3, UserInfo: wire.AppendVarint(nil, 1, 1)}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 3, UserInfo: wire.AppendVarint(nil, 1, 1)}
if _, err := seed.Login([]byte("not protobuf"), []byte("0123456789abcdef0123456789abcdef")); err == nil {
t.Fatal("Login accepted invalid protobuf request")
}
@@ -57,11 +57,11 @@ func TestLoginValidatesEncryptedRequest(t *testing.T) {
func TestLoadRejectsSeedWithUserKey(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "bad.json")
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 11, UserInfo: wire.AppendString(nil, 3, "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 11, UserInfo: wire.AppendString(nil, 3, "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")}
if err := seed.Write(path); err == nil {
t.Fatal("Write accepted user_key")
}
if err := os.WriteFile(path, []byte(`{"version":"2.34.13","packet_code":11,"user_info_base64":"GgF4"}`), 0o644); err != nil {
if err := os.WriteFile(path, []byte(`{"version":"2.35.10","packet_code":11,"user_info_base64":"GgF4"}`), 0o644); err != nil {
t.Fatal(err)
}
if _, err := Load(path); err == nil {
@@ -70,7 +70,7 @@ func TestLoadRejectsSeedWithUserKey(t *testing.T) {
}
func TestCheckedInSeedBuildsLoginWithoutCapture(t *testing.T) {
seed, err := Load(filepath.Join("..", "..", "seed", "v2_34_13", "login_user.json"))
seed, err := Load(filepath.Join("..", "..", "..", "seed", "v2_35_10", "login_user.json"))
if err != nil {
t.Fatal(err)
}
@@ -96,7 +96,7 @@ func (loginCurrencyFixture) Currencies() (uint64, uint64, uint64, uint64) {
func (loginCurrencyFixture) EquipmentMileageBalances() (uint64, uint64) { return 17, 845 }
func TestLoginRestoresEquipmentMileageFromCurrencyProvider(t *testing.T) {
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 3, UserInfo: wire.AppendVarint(nil, 1, 1)}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 3, UserInfo: wire.AppendVarint(nil, 1, 1)}
if err := seed.AttachCurrencies(loginCurrencyFixture{}); err != nil {
t.Fatal(err)
}
@@ -126,7 +126,7 @@ func TestLoginReplacesSeedPurchaseCountsFromProvider(t *testing.T) {
stale := wire.AppendVarint(nil, 1, 999)
userTemplate := wire.AppendVarint(nil, 1, 1)
userTemplate = wire.AppendBytes(userTemplate, 26, stale)
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 3, UserInfo: userTemplate}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 3, UserInfo: userTemplate}
current := wire.AppendVarint(nil, 1, 1100001)
current = wire.AppendVarint(current, 2, 9100033)
@@ -175,7 +175,7 @@ func (f *loginPresetSlotFixture) PresetSlotCount() uint64 { return f.count }
func TestLoginReplacesSeedPresetSlotFromProvider(t *testing.T) {
userTemplate := wire.AppendVarint(nil, 1, 1)
userTemplate = wire.AppendVarint(userTemplate, 28, 6)
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 3, UserInfo: userTemplate}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 3, UserInfo: userTemplate}
provider := &loginPresetSlotFixture{count: 9}
if err := seed.AttachPresetSlots(provider); err != nil {
t.Fatal(err)
@@ -223,7 +223,7 @@ func TestLoginReplacesAllInventorySlotFieldsFromProvider(t *testing.T) {
for field, value := range map[int]uint64{5: 100, 6: 100, 10: 500, 15: 100} {
user = wire.AppendVarint(user, field, value)
}
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 3, UserInfo: user}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 3, UserInfo: user}
provider := &loginInventorySlotFixture{items: 500, storage: 100, equipment: 2000, equipmentStorage: 100}
if err := seed.AttachInventorySlots(provider); err != nil {
t.Fatal(err)
@@ -248,7 +248,7 @@ func TestSeedInventorySlotsReadsUserInfoFields(t *testing.T) {
for field, value := range map[int]uint64{5: 100, 6: 101, 10: 500, 15: 102} {
user = wire.AppendVarint(user, field, value)
}
seed := &LoginSeed{Version: ProtocolVersion(), PacketCode: 3, UserInfo: user}
seed := &LoginSeed{Version: StateVersion(), PacketCode: 3, UserInfo: user}
items, storage, equipment, equipmentStorage, err := seed.SeedInventorySlots()
if err != nil || items != 100 || storage != 101 || equipment != 500 || equipmentStorage != 102 {
t.Fatalf("slots=%d/%d/%d/%d err=%v", items, storage, equipment, equipmentStorage, err)
@@ -4,7 +4,7 @@ import (
"context"
"fmt"
"bd2server/internal/stateio"
"bd2server/internal/server/stateio"
)
var _ stateio.AtomicEntryStore = (*Repository)(nil)
@@ -5,7 +5,7 @@ import (
"context"
"testing"
"bd2server/internal/stateio"
"bd2server/internal/server/stateio"
)
func TestSaveWithEntriesAtomicAndEntryOnly(t *testing.T) {
@@ -6,7 +6,7 @@ import (
"errors"
"fmt"
"bd2server/internal/stateio"
"bd2server/internal/server/stateio"
)
var _ stateio.EntryStore = (*Repository)(nil)
@@ -13,7 +13,7 @@ import (
"strconv"
"sync"
"bd2server/internal/stateio"
"bd2server/internal/server/stateio"
_ "modernc.org/sqlite"
)
+852
View File
@@ -0,0 +1,852 @@
package auth
import (
"bytes"
"context"
"crypto/sha256"
"crypto/subtle"
"database/sql"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"sync"
"time"
"bd2server/internal/server/authconfig"
"bd2server/internal/server/wire"
)
type Service struct {
config authconfig.Runtime
store *Store
client *http.Client
limits requestLimiter
}
type limitWindow struct {
started time.Time
count int
}
type requestLimiter struct {
mu sync.Mutex
windows map[string]limitWindow
lastSweep time.Time
}
type deviceResult struct {
Provider string `json:"provider"`
AccessToken string `json:"access_token"`
AccessExpiresIn int64 `json:"access_expires_in"`
RefreshToken string `json:"refresh_token"`
RefreshExpiresIn int64 `json:"refresh_expires_in"`
}
func New(config authconfig.Runtime, store *Store) (*Service, error) {
if config.Mode != "oauth" || store == nil {
return nil, errors.New("auth: OAuth service requires oauth configuration and store")
}
// Store.Open has already derived its purpose-specific keys. Do not retain
// the environment master key in the long-lived HTTP service configuration.
clear(config.MasterKey)
config.MasterKey = nil
return &Service{config: config, store: store, client: &http.Client{Timeout: 15 * time.Second}, limits: requestLimiter{windows: make(map[string]limitWindow)}}, nil
}
func (s *Service) Handler() http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("POST /auth/device", s.createDevice)
mux.HandleFunc("GET /auth/{provider}/start", s.start)
mux.HandleFunc("GET /auth/{provider}/callback", s.callback)
mux.HandleFunc("POST /auth/device/{id}/poll", s.poll)
mux.HandleFunc("POST /auth/session/refresh", s.refresh)
mux.HandleFunc("POST /auth/session/revoke", s.revoke)
return securityHeaders(mux)
}
func securityHeaders(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "no-referrer")
w.Header().Set("Content-Security-Policy", "default-src 'none'; frame-ancestors 'none'")
w.Header().Set("X-Frame-Options", "DENY")
next.ServeHTTP(w, r)
})
}
func decodeJSON(w http.ResponseWriter, r *http.Request, target any) bool {
defer r.Body.Close()
data, err := io.ReadAll(io.LimitReader(r.Body, 16<<10+1))
if err != nil || len(data) > 16<<10 {
http.Error(w, "request too large", http.StatusRequestEntityTooLarge)
return false
}
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
http.Error(w, "invalid JSON", http.StatusBadRequest)
return false
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
http.Error(w, "invalid JSON", http.StatusBadRequest)
return false
}
return true
}
func writeJSON(w http.ResponseWriter, status int, value any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(value)
}
func (s *Service) createDevice(w http.ResponseWriter, r *http.Request) {
clientIP := remoteIP(r.RemoteAddr)
if !s.limits.allow("create:"+clientIP, s.store.now(), time.Minute, 10) {
w.Header().Set("Retry-After", "60")
http.Error(w, "too many login attempts", http.StatusTooManyRequests)
return
}
var request struct {
Provider string `json:"provider"`
}
if !decodeJSON(w, r, &request) {
return
}
if _, ok := s.config.Providers[request.Provider]; !ok {
http.Error(w, "provider is not enabled", http.StatusBadRequest)
return
}
id, err := randomToken(18)
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
secret, err := randomToken(32)
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
startTicket, err := randomToken(32)
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
now := s.store.now()
clientHash := s.store.digest("client-ip", clientIP)
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
defer tx.Rollback()
if err := cleanupExpired(tx, now.Unix()); err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
var pending int
if err := tx.QueryRow(`SELECT COUNT(*) FROM devices WHERE client_hash=? AND status IN ('created','authorizing') AND expires_at>?`, clientHash, now.Unix()).Scan(&pending); err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
if pending >= 5 {
w.Header().Set("Retry-After", strconv.FormatInt(int64(s.config.DeviceTTL.Seconds()), 10))
http.Error(w, "too many pending login transactions", http.StatusTooManyRequests)
return
}
_, err = tx.Exec(`INSERT INTO devices(id,client_hash,secret_hash,start_hash,provider,status,created_at,expires_at) VALUES(?,?,?,?,?,'created',?,?)`, id, clientHash, s.store.digest("device-secret", secret), s.store.digest("start-ticket", startTicket), request.Provider, now.Unix(), now.Add(s.config.DeviceTTL).Unix())
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
if err := tx.Commit(); err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
start := *s.config.PublicURLParsed
start.Path = "/auth/" + request.Provider + "/start"
query := start.Query()
query.Set("transaction_id", id)
query.Set("ticket", startTicket)
start.RawQuery = query.Encode()
writeJSON(w, http.StatusCreated, map[string]any{"transaction_id": id, "device_secret": secret, "start_url": start.String(), "expires_in": int64(s.config.DeviceTTL.Seconds()), "poll_interval": 2})
}
func (s *Service) start(w http.ResponseWriter, r *http.Request) {
provider := r.PathValue("provider")
if _, ok := s.config.Providers[provider]; !ok {
http.Error(w, "provider is not enabled", http.StatusNotFound)
return
}
id, ticket := r.URL.Query().Get("transaction_id"), r.URL.Query().Get("ticket")
var storedHash []byte
var storedProvider, status string
var expires int64
err := s.store.db.QueryRow(`SELECT start_hash,provider,status,expires_at FROM devices WHERE id=?`, id).Scan(&storedHash, &storedProvider, &status, &expires)
if err != nil || subtle.ConstantTimeCompare(storedHash, s.store.digest("start-ticket", ticket)) != 1 || storedProvider != provider || status != "created" {
http.Error(w, "invalid login transaction", http.StatusForbidden)
return
}
if s.store.now().Unix() >= expires {
http.Error(w, "login transaction expired", http.StatusGone)
return
}
state, err := randomToken(32)
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
verifier, err := randomToken(32)
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
nonce, err := randomToken(24)
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
verifierCipher, err := s.store.seal(id, "pkce", []byte(verifier))
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
nonceCipher, err := s.store.seal(id, "nonce", []byte(nonce))
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
result, err := s.store.db.Exec(`UPDATE devices SET state_hash=?,verifier_cipher=?,nonce_cipher=?,start_hash=X'',status='authorizing' WHERE id=? AND status='created'`, s.store.digest("oauth-state", state), verifierCipher, nonceCipher, id)
count, affectedErr := rowsAffected(result)
if err != nil || affectedErr != nil || count != 1 {
http.Error(w, "could not start authorization", http.StatusConflict)
return
}
redirect := s.redirectURL(provider)
challenge := sha256.Sum256([]byte(verifier))
values := url.Values{"client_id": {s.config.Providers[provider].ClientID}, "redirect_uri": {redirect}, "response_type": {"code"}, "scope": {providerScope(provider)}, "state": {state}, "code_challenge": {base64.RawURLEncoding.EncodeToString(challenge[:])}, "code_challenge_method": {"S256"}}
if provider == "google" {
values.Set("nonce", nonce)
}
http.Redirect(w, r, providerAuthorizeURL(provider)+"?"+values.Encode(), http.StatusFound)
}
func (s *Service) callback(w http.ResponseWriter, r *http.Request) {
provider, state, code := r.PathValue("provider"), r.URL.Query().Get("state"), r.URL.Query().Get("code")
if _, ok := s.config.Providers[provider]; !ok {
http.Error(w, "provider is not enabled", http.StatusNotFound)
return
}
if state == "" {
http.Error(w, "authorization was not completed", http.StatusBadRequest)
return
}
if r.URL.Query().Get("error") != "" {
result, err := s.store.db.Exec(`UPDATE devices SET status='failed',error_code='provider_cancelled',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL
WHERE state_hash=? AND provider=? AND status='authorizing' AND expires_at>?`, s.store.digest("oauth-state", state), provider, s.store.now().Unix())
if err != nil {
http.Error(w, "authorization state unavailable", http.StatusInternalServerError)
return
}
if count, err := rowsAffected(result); err != nil || count != 1 {
http.Error(w, "invalid or expired authorization state", http.StatusForbidden)
return
}
http.Error(w, "authorization was cancelled", http.StatusBadRequest)
return
}
if code == "" {
http.Error(w, "authorization was not completed", http.StatusBadRequest)
return
}
var id, storedProvider, status string
var verifierCipher, nonceCipher []byte
var expires int64
err := s.store.db.QueryRow(`SELECT id,provider,status,verifier_cipher,nonce_cipher,expires_at FROM devices WHERE state_hash=?`, s.store.digest("oauth-state", state)).Scan(&id, &storedProvider, &status, &verifierCipher, &nonceCipher, &expires)
if err != nil || provider != storedProvider || status != "authorizing" || s.store.now().Unix() >= expires {
http.Error(w, "invalid or expired authorization state", http.StatusForbidden)
return
}
verifier, err := s.store.open(id, "pkce", verifierCipher)
if err != nil {
http.Error(w, "authorization state unavailable", http.StatusInternalServerError)
return
}
nonce, err := s.store.open(id, "nonce", nonceCipher)
if err != nil {
http.Error(w, "authorization state unavailable", http.StatusInternalServerError)
return
}
identity, err := s.exchangeIdentity(r.Context(), provider, code, string(verifier), string(nonce))
clear(verifier)
clear(nonce)
if err != nil {
_, _ = s.store.db.Exec(`UPDATE devices SET status='failed',error_code='provider_rejected' WHERE id=? AND status='authorizing'`, id)
http.Error(w, "provider authorization failed", http.StatusBadGateway)
return
}
if err := s.completeDevice(id, provider, identity); err != nil {
code := http.StatusInternalServerError
if errors.Is(err, ErrNotAllowed) {
code = http.StatusForbidden
} else if errors.Is(err, ErrConsumed) {
code = http.StatusConflict
}
http.Error(w, "authorization could not be completed", code)
return
}
w.Header().Set("Content-Security-Policy", "default-src 'none'; style-src 'unsafe-inline'")
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = io.WriteString(w, `<!doctype html><meta charset="utf-8"><title>BD2 login</title><p>Login complete. You can return to the game.</p>`)
}
type providerIdentity struct{ issuer, subject string }
func (s *Service) exchangeIdentity(ctx context.Context, provider, code, verifier, nonce string) (providerIdentity, error) {
values := url.Values{"client_id": {s.config.Providers[provider].ClientID}, "client_secret": {s.config.ProviderSecrets[provider]}, "grant_type": {"authorization_code"}, "code": {code}, "redirect_uri": {s.redirectURL(provider)}, "code_verifier": {verifier}}
request, _ := http.NewRequestWithContext(ctx, http.MethodPost, providerTokenURL(provider), strings.NewReader(values.Encode()))
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
response, err := s.client.Do(request)
if err != nil {
return providerIdentity{}, err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return providerIdentity{}, errors.New("token exchange rejected")
}
var token struct {
AccessToken string `json:"access_token"`
IDToken string `json:"id_token"`
}
if err := decodeProviderJSON(response.Body, &token); err != nil || token.AccessToken == "" {
return providerIdentity{}, errors.New("invalid token response")
}
if provider == "google" {
if token.IDToken == "" {
return providerIdentity{}, errors.New("Google ID token missing")
}
identity, err := s.verifyGoogleIDToken(ctx, token.IDToken, nonce)
token.AccessToken, token.IDToken = "", ""
return identity, err
}
userinfo, _ := http.NewRequestWithContext(ctx, http.MethodGet, providerUserURL(provider), nil)
userinfo.Header.Set("Authorization", "Bearer "+token.AccessToken)
response, err = s.client.Do(userinfo)
token.AccessToken = ""
if err != nil {
return providerIdentity{}, err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return providerIdentity{}, errors.New("userinfo rejected")
}
var user struct {
ID string `json:"id"`
Sub string `json:"sub"`
}
if err := decodeProviderJSON(response.Body, &user); err != nil {
return providerIdentity{}, err
}
if provider == "discord" && user.ID != "" {
return providerIdentity{issuer: "https://discord.com", subject: user.ID}, nil
}
return providerIdentity{}, errors.New("provider subject missing")
}
// verifyGoogleIDToken delegates signature and standard-claim verification to
// Google's HTTPS tokeninfo endpoint, then independently verifies this server's
// audience, nonce and expiry. The raw ID token is never persisted or logged.
func (s *Service) verifyGoogleIDToken(ctx context.Context, idToken, nonce string) (providerIdentity, error) {
endpoint := "https://oauth2.googleapis.com/tokeninfo?id_token=" + url.QueryEscape(idToken)
request, _ := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
response, err := s.client.Do(request)
if err != nil {
return providerIdentity{}, err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return providerIdentity{}, errors.New("Google ID token rejected")
}
var claims struct {
Issuer string `json:"iss"`
Audience string `json:"aud"`
Subject string `json:"sub"`
Nonce string `json:"nonce"`
Expires string `json:"exp"`
}
if err := decodeProviderJSON(response.Body, &claims); err != nil {
return providerIdentity{}, err
}
expires, err := strconv.ParseInt(claims.Expires, 10, 64)
validIssuer := claims.Issuer == "https://accounts.google.com" || claims.Issuer == "accounts.google.com"
if err != nil || !validIssuer || claims.Audience != s.config.Providers["google"].ClientID || claims.Subject == "" || claims.Nonce != nonce || s.store.now().Unix() >= expires {
return providerIdentity{}, errors.New("Google ID token claims rejected")
}
return providerIdentity{issuer: "https://accounts.google.com", subject: claims.Subject}, nil
}
func decodeProviderJSON(reader io.Reader, target any) error {
data, err := io.ReadAll(io.LimitReader(reader, 1<<20+1))
if err != nil {
return err
}
if len(data) > 1<<20 {
return errors.New("auth: provider response is too large")
}
decoder := json.NewDecoder(bytes.NewReader(data))
if err := decoder.Decode(target); err != nil {
return err
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return errors.New("auth: provider response has trailing JSON")
}
return nil
}
func (s *Service) completeDevice(deviceID, provider string, identity providerIdentity) error {
if err := validateProviderIdentity(provider, identity); err != nil {
return err
}
now := s.store.now()
tx, err := s.store.db.Begin()
if err != nil {
return err
}
defer tx.Rollback()
var accountID, status string
subjectHash := s.store.identityDigest(identity.issuer, identity.subject)
err = tx.QueryRow(`SELECT i.account_id,a.status FROM identities i JOIN accounts a ON a.id=i.account_id WHERE i.issuer=? AND i.subject_hash=?`, identity.issuer, subjectHash).Scan(&accountID, &status)
if errors.Is(err, sql.ErrNoRows) {
var count int
if err := tx.QueryRow(`SELECT COUNT(*) FROM accounts`).Scan(&count); err != nil {
return err
}
if count != 0 {
result, err := tx.Exec(`UPDATE devices SET status='failed',error_code='not_allowed',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND provider=? AND status='authorizing'`, deviceID, provider)
if err != nil {
return err
}
if count, err := rowsAffected(result); err != nil {
return err
} else if count != 1 {
return ErrConsumed
}
if err := tx.Commit(); err != nil {
return err
}
return ErrNotAllowed
}
accountID, err = randomToken(18)
if err != nil {
return err
}
if _, err = tx.Exec(`INSERT INTO accounts(id,status,created_at,last_login_at) VALUES(?,'active',?,?)`, accountID, now.Unix(), now.Unix()); err != nil {
return err
}
if _, err = tx.Exec(`INSERT INTO identities(provider,issuer,subject_hash,account_id,created_at,last_login_at) VALUES(?,?,?,?,?,?)`, provider, identity.issuer, subjectHash, accountID, now.Unix(), now.Unix()); err != nil {
return err
}
status = "active"
} else if err != nil {
return err
} else {
if _, err = tx.Exec(`UPDATE identities SET last_login_at=? WHERE issuer=? AND subject_hash=?`, now.Unix(), identity.issuer, subjectHash); err != nil {
return err
}
}
if status != "active" {
result, err := tx.Exec(`UPDATE devices SET status='failed',error_code='not_allowed',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND provider=? AND status='authorizing'`, deviceID, provider)
if err != nil {
return err
}
if count, err := rowsAffected(result); err != nil {
return err
} else if count != 1 {
return ErrConsumed
}
if err := tx.Commit(); err != nil {
return err
}
return ErrNotAllowed
}
result, familyID, err := s.issueTokens(tx, accountID, provider, now)
if err != nil {
return err
}
_ = familyID
payload, err := json.Marshal(result)
if err != nil {
return err
}
sealed, err := s.store.seal(deviceID, "result", payload)
clear(payload)
if err != nil {
return err
}
update, err := tx.Exec(`UPDATE devices SET result_cipher=?,status='complete',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND provider=? AND status='authorizing'`, sealed, deviceID, provider)
if err != nil {
return err
}
if count, err := rowsAffected(update); err != nil {
return err
} else if count != 1 {
return ErrConsumed
}
return tx.Commit()
}
func validateProviderIdentity(provider string, identity providerIdentity) error {
switch provider {
case "discord":
if identity.issuer != "https://discord.com" || len(identity.subject) == 0 || len(identity.subject) > 32 {
return errors.New("auth: invalid Discord identity")
}
for _, digit := range identity.subject {
if digit < '0' || digit > '9' {
return errors.New("auth: invalid Discord identity")
}
}
case "google":
if identity.issuer != "https://accounts.google.com" || len(identity.subject) == 0 || len(identity.subject) > 255 {
return errors.New("auth: invalid Google identity")
}
default:
return errors.New("auth: unsupported identity provider")
}
return nil
}
func (s *Service) issueTokens(tx *sql.Tx, accountID, provider string, now time.Time) (deviceResult, string, error) {
familyID, err := randomToken(18)
if err != nil {
return deviceResult{}, "", err
}
access, err := randomToken(32)
if err != nil {
return deviceResult{}, "", err
}
refresh, err := randomToken(32)
if err != nil {
return deviceResult{}, "", err
}
if _, err := tx.Exec(`INSERT INTO families(id,account_id,provider,created_at,expires_at) VALUES(?,?,?,?,?)`, familyID, accountID, provider, now.Unix(), now.Add(s.config.RefreshTTL).Unix()); err != nil {
return deviceResult{}, "", err
}
if _, err := tx.Exec(`INSERT INTO access_tokens(token_hash,family_id,account_id,created_at,expires_at) VALUES(?,?,?,?,?)`, s.store.digest("access-token", access), familyID, accountID, now.Unix(), now.Add(s.config.AccessTTL).Unix()); err != nil {
return deviceResult{}, "", err
}
if _, err := tx.Exec(`INSERT INTO refresh_tokens(token_hash,family_id,created_at,expires_at) VALUES(?,?,?,?)`, s.store.digest("refresh-token", refresh), familyID, now.Unix(), now.Add(s.config.RefreshTTL).Unix()); err != nil {
return deviceResult{}, "", err
}
return deviceResult{Provider: provider, AccessToken: access, AccessExpiresIn: int64(s.config.AccessTTL.Seconds()), RefreshToken: refresh, RefreshExpiresIn: int64(s.config.RefreshTTL.Seconds())}, familyID, nil
}
func (s *Service) poll(w http.ResponseWriter, r *http.Request) {
id := r.PathValue("id")
if !s.limits.allow("poll:"+remoteIP(r.RemoteAddr)+":"+id, s.store.now(), time.Minute, 60) {
w.Header().Set("Retry-After", "2")
http.Error(w, "poll rate exceeded", http.StatusTooManyRequests)
return
}
authorization := r.Header.Get("Authorization")
if !strings.HasPrefix(authorization, "Device ") {
http.Error(w, "invalid device transaction", http.StatusForbidden)
return
}
secret := strings.TrimPrefix(authorization, "Device ")
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "login result unavailable", http.StatusInternalServerError)
return
}
defer tx.Rollback()
var storedHash, sealed []byte
var status, errorCode string
var expires int64
err = tx.QueryRow(`SELECT secret_hash,status,COALESCE(result_cipher,X''),COALESCE(error_code,''),expires_at FROM devices WHERE id=?`, id).Scan(&storedHash, &status, &sealed, &errorCode, &expires)
if err != nil || subtle.ConstantTimeCompare(storedHash, s.store.digest("device-secret", secret)) != 1 {
http.Error(w, "invalid device transaction", http.StatusForbidden)
return
}
if s.store.now().Unix() >= expires {
http.Error(w, "device transaction expired", http.StatusGone)
return
}
switch status {
case "created", "authorizing":
_ = tx.Rollback()
writeJSON(w, http.StatusAccepted, map[string]any{"status": "pending", "retry_after": 2})
case "failed":
_ = tx.Rollback()
writeJSON(w, http.StatusForbidden, map[string]string{"status": "failed", "error": errorCode})
case "complete":
plain, err := s.store.open(id, "result", sealed)
if err != nil {
http.Error(w, "login result unavailable", http.StatusInternalServerError)
return
}
result, err := tx.Exec(`UPDATE devices SET result_cipher=NULL,status='consumed' WHERE id=? AND status='complete'`, id)
count, affectedErr := rowsAffected(result)
if err != nil || affectedErr != nil || count != 1 {
_ = tx.Rollback()
clear(plain)
http.Error(w, "login result already consumed", http.StatusGone)
return
}
if err := tx.Commit(); err != nil {
clear(plain)
http.Error(w, "login result unavailable", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
_, _ = w.Write(plain)
clear(plain)
default:
_ = tx.Rollback()
http.Error(w, "device transaction consumed", http.StatusGone)
}
}
func (l *requestLimiter) allow(key string, now time.Time, duration time.Duration, maximum int) bool {
l.mu.Lock()
defer l.mu.Unlock()
if l.windows == nil {
l.windows = make(map[string]limitWindow)
}
if l.lastSweep.IsZero() || now.Sub(l.lastSweep) >= time.Minute {
for candidate, window := range l.windows {
if now.Sub(window.started) >= duration {
delete(l.windows, candidate)
}
}
l.lastSweep = now
}
window, exists := l.windows[key]
if !exists || now.Sub(window.started) >= duration {
if !exists && len(l.windows) >= 4096 {
return false
}
l.windows[key] = limitWindow{started: now, count: 1}
return true
}
if window.count >= maximum {
return false
}
window.count++
l.windows[key] = window
return true
}
func remoteIP(remoteAddr string) string {
host, _, err := net.SplitHostPort(remoteAddr)
if err == nil && host != "" {
return host
}
return remoteAddr
}
func cleanupExpired(tx *sql.Tx, now int64) error {
statements := []struct {
query string
args []any
}{
{`DELETE FROM devices WHERE expires_at<=?`, []any{now}},
{`DELETE FROM access_tokens WHERE expires_at<=? OR family_id IN (SELECT id FROM families WHERE expires_at<=?)`, []any{now, now}},
// Used refresh rows remain until their family expires so their reuse can
// still revoke every credential in that family.
{`DELETE FROM refresh_tokens WHERE family_id IN (SELECT id FROM families WHERE expires_at<=?)`, []any{now}},
{`DELETE FROM families WHERE expires_at<=?`, []any{now}},
}
for _, statement := range statements {
if _, err := tx.Exec(statement.query, statement.args...); err != nil {
return err
}
}
return nil
}
func (s *Service) refresh(w http.ResponseWriter, r *http.Request) {
var request struct {
RefreshToken string `json:"refresh_token"`
}
if !decodeJSON(w, r, &request) || request.RefreshToken == "" {
return
}
now := s.store.now()
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
defer tx.Rollback()
var familyID, accountID, provider, accountStatus string
var tokenExpires, familyExpires int64
var usedAt, revokedAt sql.NullInt64
err = tx.QueryRow(`SELECT r.family_id,f.account_id,f.provider,a.status,r.expires_at,f.expires_at,r.used_at,COALESCE(r.revoked_at,f.revoked_at) FROM refresh_tokens r JOIN families f ON f.id=r.family_id JOIN accounts a ON a.id=f.account_id WHERE r.token_hash=?`, s.store.digest("refresh-token", request.RefreshToken)).Scan(&familyID, &accountID, &provider, &accountStatus, &tokenExpires, &familyExpires, &usedAt, &revokedAt)
if err != nil || revokedAt.Valid || accountStatus != "active" || now.Unix() >= tokenExpires || now.Unix() >= familyExpires {
http.Error(w, "refresh token invalid", http.StatusUnauthorized)
return
}
if usedAt.Valid {
if _, err := tx.Exec(`UPDATE families SET revoked_at=? WHERE id=? AND revoked_at IS NULL`, now.Unix(), familyID); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if err := tx.Commit(); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
http.Error(w, "refresh token replayed", http.StatusUnauthorized)
return
}
newAccess, err := randomToken(32)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
newRefresh, err := randomToken(32)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
updated, err := tx.Exec(`UPDATE refresh_tokens SET used_at=? WHERE token_hash=? AND used_at IS NULL AND revoked_at IS NULL`, now.Unix(), s.store.digest("refresh-token", request.RefreshToken))
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if count, err := rowsAffected(updated); err != nil || count != 1 {
http.Error(w, "refresh token invalid", http.StatusUnauthorized)
return
}
if _, err = tx.Exec(`DELETE FROM access_tokens WHERE family_id=?`, familyID); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if _, err = tx.Exec(`INSERT INTO access_tokens(token_hash,family_id,account_id,created_at,expires_at) VALUES(?,?,?,?,?)`, s.store.digest("access-token", newAccess), familyID, accountID, now.Unix(), now.Add(s.config.AccessTTL).Unix()); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
refreshExpiry := min(familyExpires, now.Add(s.config.RefreshTTL).Unix())
if _, err = tx.Exec(`INSERT INTO refresh_tokens(token_hash,family_id,created_at,expires_at) VALUES(?,?,?,?)`, s.store.digest("refresh-token", newRefresh), familyID, now.Unix(), refreshExpiry); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if err = tx.Commit(); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
writeJSON(w, http.StatusOK, deviceResult{Provider: provider, AccessToken: newAccess, AccessExpiresIn: int64(s.config.AccessTTL.Seconds()), RefreshToken: newRefresh, RefreshExpiresIn: refreshExpiry - now.Unix()})
}
func (s *Service) revoke(w http.ResponseWriter, r *http.Request) {
authorization := r.Header.Get("Authorization")
if !strings.HasPrefix(authorization, "Bearer ") {
http.Error(w, "access token required", http.StatusUnauthorized)
return
}
token := strings.TrimPrefix(authorization, "Bearer ")
if token == "" {
http.Error(w, "access token required", http.StatusUnauthorized)
return
}
now := s.store.now().Unix()
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "revocation unavailable", http.StatusInternalServerError)
return
}
defer tx.Rollback()
var familyID, accountStatus string
var expires int64
var revoked sql.NullInt64
err = tx.QueryRow(`SELECT t.family_id,a.status,t.expires_at,COALESCE(t.revoked_at,f.revoked_at)
FROM access_tokens t JOIN families f ON f.id=t.family_id JOIN accounts a ON a.id=t.account_id
WHERE t.token_hash=?`, s.store.digest("access-token", token)).Scan(&familyID, &accountStatus, &expires, &revoked)
if err != nil || accountStatus != "active" || revoked.Valid || now >= expires {
http.Error(w, "access token invalid", http.StatusUnauthorized)
return
}
result, err := tx.Exec(`UPDATE families SET revoked_at=? WHERE id=? AND revoked_at IS NULL`, now, familyID)
if err != nil {
http.Error(w, "revocation unavailable", http.StatusInternalServerError)
return
}
if count, err := rowsAffected(result); err != nil || count != 1 {
http.Error(w, "access token invalid", http.StatusUnauthorized)
return
}
if err := tx.Commit(); err != nil {
http.Error(w, "revocation unavailable", http.StatusInternalServerError)
return
}
w.WriteHeader(http.StatusNoContent)
}
func (s *Service) ValidateAccess(token string) (string, error) {
if token == "" {
return "", ErrUnauthorized
}
var accountID, status string
var expires int64
var revoked sql.NullInt64
err := s.store.db.QueryRow(`SELECT t.account_id,a.status,t.expires_at,COALESCE(t.revoked_at,f.revoked_at) FROM access_tokens t JOIN families f ON f.id=t.family_id JOIN accounts a ON a.id=t.account_id WHERE t.token_hash=?`, s.store.digest("access-token", token)).Scan(&accountID, &status, &expires, &revoked)
if err != nil || status != "active" || revoked.Valid || s.store.now().Unix() >= expires {
return "", ErrUnauthorized
}
return accountID, nil
}
// AuthenticateLogin validates LoginUserRequest.access_token (field 2) before
// the game session is established.
func (s *Service) AuthenticateLogin(request []byte) (string, error) {
token, found, err := wire.Bytes(request, 2)
if err != nil || !found {
return "", ErrUnauthorized
}
return s.ValidateAccess(string(token))
}
func (s *Service) redirectURL(provider string) string {
return s.config.PublicURL + "/auth/" + provider + "/callback"
}
func providerScope(provider string) string {
if provider == "discord" {
return "identify"
}
return "openid"
}
func providerAuthorizeURL(provider string) string {
if provider == "discord" {
return "https://discord.com/oauth2/authorize"
}
return "https://accounts.google.com/o/oauth2/v2/auth"
}
func providerTokenURL(provider string) string {
if provider == "discord" {
return "https://discord.com/api/v10/oauth2/token"
}
return "https://oauth2.googleapis.com/token"
}
func providerUserURL(provider string) string {
return "https://discord.com/api/v10/users/@me"
}
func rowsAffected(result sql.Result) (int64, error) {
if result == nil {
return 0, errors.New("auth: missing SQL result")
}
value, err := result.RowsAffected()
if err != nil {
return 0, fmt.Errorf("auth: count affected rows: %w", err)
}
return value, nil
}
+504
View File
@@ -0,0 +1,504 @@
package auth
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
"net/url"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"bd2server/internal/server/authconfig"
)
const testNowUnix = int64(1_800_000_000)
func testService(t *testing.T) (*Service, *Store) {
t.Helper()
master := bytes.Repeat([]byte{0x42}, 32)
store, err := Open(filepath.Join(t.TempDir(), "auth.db"), master)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = store.Close() })
store.now = func() time.Time { return time.Unix(testNowUnix, 0) }
public, _ := url.Parse("https://login.example.test")
runtime := authconfig.Runtime{
Config: authconfig.Config{
Mode: "oauth",
PublicURL: "https://login.example.test",
Providers: map[string]authconfig.ProviderConfig{
"discord": {ClientID: "discord-client", ClientSecretEnv: "DISCORD_SECRET"},
"google": {ClientID: "google-client", ClientSecretEnv: "GOOGLE_SECRET"},
},
},
PublicURLParsed: public,
MasterKey: master,
ProviderSecrets: map[string]string{"discord": "discord-secret", "google": "google-secret"},
AccessTTL: 15 * time.Minute,
RefreshTTL: 30 * 24 * time.Hour,
DeviceTTL: 10 * time.Minute,
}
service, err := New(runtime, store)
if err != nil {
t.Fatal(err)
}
if service.config.MasterKey != nil {
t.Fatal("service retained the authentication master key")
}
for i, value := range master {
if value != 0 {
t.Fatalf("master key byte %d was not cleared", i)
}
}
return service, store
}
func insertAuthorizingDevice(t *testing.T, store *Store, id, provider string) {
t.Helper()
_, err := store.db.Exec(`INSERT INTO devices(id,client_hash,secret_hash,start_hash,provider,status,created_at,expires_at)
VALUES(?,?,?,?,?,'authorizing',?,?)`, id, store.digest("client-ip", "192.0.2.1"), store.digest("device-secret", "device-secret"), store.digest("start-ticket", "start-ticket"), provider, testNowUnix, testNowUnix+600)
if err != nil {
t.Fatal(err)
}
}
func completeAndPoll(t *testing.T, service *Service, store *Store, deviceID, provider, issuer, subject string) deviceResult {
t.Helper()
insertAuthorizingDevice(t, store, deviceID, provider)
if err := service.completeDevice(deviceID, provider, providerIdentity{issuer: issuer, subject: subject}); err != nil {
t.Fatal(err)
}
request := httptest.NewRequest(http.MethodPost, "/auth/device/"+deviceID+"/poll", nil)
request.Header.Set("Authorization", "Device device-secret")
response := httptest.NewRecorder()
service.Handler().ServeHTTP(response, request)
if response.Code != http.StatusOK {
t.Fatalf("poll status=%d body=%q", response.Code, response.Body.String())
}
var result deviceResult
if err := json.Unmarshal(response.Body.Bytes(), &result); err != nil {
t.Fatal(err)
}
return result
}
func postJSON(handler http.Handler, path string, value any) *httptest.ResponseRecorder {
body, _ := json.Marshal(value)
request := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(body))
request.Header.Set("Content-Type", "application/json")
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
return response
}
func TestCompleteDeviceRequiresAuthorizingTransition(t *testing.T) {
service, store := testService(t)
insertAuthorizingDevice(t, store, "already-consumed", "discord")
if _, err := store.db.Exec(`UPDATE devices SET status='consumed' WHERE id='already-consumed'`); err != nil {
t.Fatal(err)
}
err := service.completeDevice("already-consumed", "discord", providerIdentity{issuer: "https://discord.com", subject: "123456789"})
if !errors.Is(err, ErrConsumed) {
t.Fatalf("completeDevice error=%v, want ErrConsumed", err)
}
for _, table := range []string{"accounts", "identities", "families", "access_tokens", "refresh_tokens"} {
var count int
if err := store.db.QueryRow(`SELECT COUNT(*) FROM ` + table).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("%s has %d rows after rejected completion", table, count)
}
}
}
func TestPollConsumesEncryptedResultExactlyOnce(t *testing.T) {
service, store := testService(t)
result := completeAndPoll(t, service, store, "poll-once", "discord", "https://discord.com", "123456789")
if result.AccessToken == "" || result.RefreshToken == "" {
t.Fatal("poll omitted issued tokens")
}
request := httptest.NewRequest(http.MethodPost, "/auth/device/poll-once/poll", nil)
request.Header.Set("Authorization", "Device device-secret")
response := httptest.NewRecorder()
service.Handler().ServeHTTP(response, request)
if response.Code != http.StatusGone {
t.Fatalf("second poll status=%d body=%q", response.Code, response.Body.String())
}
var status string
var cipher []byte
if err := store.db.QueryRow(`SELECT status,COALESCE(result_cipher,X'') FROM devices WHERE id='poll-once'`).Scan(&status, &cipher); err != nil {
t.Fatal(err)
}
if status != "consumed" || len(cipher) != 0 {
t.Fatalf("device status=%q result bytes=%d", status, len(cipher))
}
}
func TestRefreshRotationReplayRevokesFamily(t *testing.T) {
service, store := testService(t)
first := completeAndPoll(t, service, store, "refresh-device", "discord", "https://discord.com", "123456789")
handler := service.Handler()
response := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": first.RefreshToken})
if response.Code != http.StatusOK {
t.Fatalf("refresh status=%d body=%q", response.Code, response.Body.String())
}
var rotated deviceResult
if err := json.Unmarshal(response.Body.Bytes(), &rotated); err != nil {
t.Fatal(err)
}
if rotated.RefreshToken == "" || rotated.RefreshToken == first.RefreshToken || rotated.AccessToken == first.AccessToken {
t.Fatal("refresh did not rotate both credentials")
}
if _, err := service.ValidateAccess(first.AccessToken); err == nil {
t.Fatal("old access token survived refresh rotation")
}
if _, err := service.ValidateAccess(rotated.AccessToken); err != nil {
t.Fatalf("new access token rejected: %v", err)
}
replay := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": first.RefreshToken})
if replay.Code != http.StatusUnauthorized {
t.Fatalf("replay status=%d body=%q", replay.Code, replay.Body.String())
}
if _, err := service.ValidateAccess(rotated.AccessToken); err == nil {
t.Fatal("refresh replay did not revoke the token family")
}
next := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": rotated.RefreshToken})
if next.Code != http.StatusUnauthorized {
t.Fatalf("family refresh after replay status=%d", next.Code)
}
}
func TestRevokeInvalidatesAccessAndRefreshFamily(t *testing.T) {
service, store := testService(t)
tokens := completeAndPoll(t, service, store, "revoke-device", "discord", "https://discord.com", "123456789")
request := httptest.NewRequest(http.MethodPost, "/auth/session/revoke", nil)
request.Header.Set("Authorization", "Bearer "+tokens.AccessToken)
response := httptest.NewRecorder()
service.Handler().ServeHTTP(response, request)
if response.Code != http.StatusNoContent {
t.Fatalf("revoke status=%d body=%q", response.Code, response.Body.String())
}
if _, err := service.ValidateAccess(tokens.AccessToken); err == nil {
t.Fatal("revoked access token remained valid")
}
refresh := postJSON(service.Handler(), "/auth/session/refresh", map[string]string{"refresh_token": tokens.RefreshToken})
if refresh.Code != http.StatusUnauthorized {
t.Fatalf("revoked refresh status=%d", refresh.Code)
}
}
func TestSensitiveAuthenticationMaterialIsNotStoredInPlaintext(t *testing.T) {
service, store := testService(t)
handler := service.Handler()
created := postJSON(handler, "/auth/device", map[string]string{"provider": "discord"})
if created.Code != http.StatusCreated {
t.Fatalf("create status=%d body=%q", created.Code, created.Body.String())
}
var device struct {
ID string `json:"transaction_id"`
Secret string `json:"device_secret"`
StartURL string `json:"start_url"`
}
if err := json.Unmarshal(created.Body.Bytes(), &device); err != nil {
t.Fatal(err)
}
startURL, _ := url.Parse(device.StartURL)
ticket := startURL.Query().Get("ticket")
start := httptest.NewRequest(http.MethodGet, startURL.RequestURI(), nil)
started := httptest.NewRecorder()
handler.ServeHTTP(started, start)
if started.Code != http.StatusFound {
t.Fatalf("start status=%d body=%q", started.Code, started.Body.String())
}
authorize, _ := url.Parse(started.Header().Get("Location"))
state := authorize.Query().Get("state")
var secretHash, startHash, stateHash, verifierCipher, nonceCipher []byte
if err := store.db.QueryRow(`SELECT secret_hash,start_hash,state_hash,verifier_cipher,nonce_cipher FROM devices WHERE id=?`, device.ID).
Scan(&secretHash, &startHash, &stateHash, &verifierCipher, &nonceCipher); err != nil {
t.Fatal(err)
}
verifier, err := store.open(device.ID, "pkce", verifierCipher)
if err != nil {
t.Fatal(err)
}
nonce, err := store.open(device.ID, "nonce", nonceCipher)
if err != nil {
t.Fatal(err)
}
for name, pair := range map[string]struct{ stored, raw []byte }{
"device secret": {secretHash, []byte(device.Secret)},
"start ticket": {startHash, []byte(ticket)},
"oauth state": {stateHash, []byte(state)},
"pkce verifier": {verifierCipher, verifier},
"oidc nonce": {nonceCipher, nonce},
} {
if bytes.Equal(pair.stored, pair.raw) || bytes.Contains(pair.stored, pair.raw) {
t.Fatalf("%s was stored in plaintext", name)
}
}
const providerSubject = "987654321012345678"
if err := service.completeDevice(device.ID, "discord", providerIdentity{issuer: "https://discord.com", subject: providerSubject}); err != nil {
t.Fatal(err)
}
var subjectHash, sealedResult []byte
if err := store.db.QueryRow(`SELECT subject_hash FROM identities`).Scan(&subjectHash); err != nil {
t.Fatal(err)
}
if err := store.db.QueryRow(`SELECT result_cipher FROM devices WHERE id=?`, device.ID).Scan(&sealedResult); err != nil {
t.Fatal(err)
}
plainResult, err := store.open(device.ID, "result", sealedResult)
if err != nil {
t.Fatal(err)
}
var issued deviceResult
if err := json.Unmarshal(plainResult, &issued); err != nil {
t.Fatal(err)
}
if bytes.Contains(subjectHash, []byte(providerSubject)) {
t.Fatal("provider subject was stored in plaintext")
}
for name, raw := range map[string]string{"access token": issued.AccessToken, "refresh token": issued.RefreshToken} {
if bytes.Contains(sealedResult, []byte(raw)) {
t.Fatalf("pending %s was stored outside AES-GCM ciphertext", name)
}
var count int
table := "access_tokens"
if name == "refresh token" {
table = "refresh_tokens"
}
if err := store.db.QueryRow(`SELECT COUNT(*) FROM `+table+` WHERE token_hash=?`, []byte(raw)).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("%s was stored in plaintext", name)
}
}
}
func TestJSONLimitsAndSecurityHeaders(t *testing.T) {
service, _ := testService(t)
handler := service.Handler()
for name, body := range map[string]struct {
body string
want int
}{
"trailing": {`{"provider":"discord"}{}`, http.StatusBadRequest},
"oversize": {`{"provider":"discord","padding":"` + strings.Repeat("x", 17<<10) + `"}`, http.StatusRequestEntityTooLarge},
} {
t.Run(name, func(t *testing.T) {
request := httptest.NewRequest(http.MethodPost, "/auth/device", strings.NewReader(body.body))
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
if response.Code != body.want {
t.Fatalf("status=%d body=%q", response.Code, response.Body.String())
}
for header, want := range map[string]string{
"Cache-Control": "no-store",
"Referrer-Policy": "no-referrer",
"X-Content-Type-Options": "nosniff",
"X-Frame-Options": "DENY",
} {
if got := response.Header().Get(header); got != want {
t.Fatalf("%s=%q want %q", header, got, want)
}
}
})
}
}
func TestCreateDeviceLimitsPendingTransactionsPerClient(t *testing.T) {
service, store := testService(t)
handler := service.Handler()
for i := 0; i < 5; i++ {
response := postJSON(handler, "/auth/device", map[string]string{"provider": "discord"})
if response.Code != http.StatusCreated {
t.Fatalf("create %d status=%d body=%q", i, response.Code, response.Body.String())
}
}
response := postJSON(handler, "/auth/device", map[string]string{"provider": "discord"})
if response.Code != http.StatusTooManyRequests {
t.Fatalf("pending limit status=%d body=%q", response.Code, response.Body.String())
}
var count int
if err := store.db.QueryRow(`SELECT COUNT(*) FROM devices`).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 5 {
t.Fatalf("device count=%d want 5", count)
}
}
func TestDecodeProviderJSONRejectsOversizeAndTrailingValues(t *testing.T) {
var target map[string]any
if err := decodeProviderJSON(strings.NewReader(`{"id":"1"}{}`), &target); err == nil {
t.Fatal("accepted provider response with trailing JSON")
}
oversize := `{"padding":"` + strings.Repeat("x", 1<<20) + `"}`
if err := decodeProviderJSON(strings.NewReader(oversize), &target); err == nil {
t.Fatal("accepted oversized provider response")
}
}
func TestRequestLimiterIsBoundedAndExpiresWindows(t *testing.T) {
limiter := requestLimiter{windows: make(map[string]limitWindow)}
now := time.Unix(testNowUnix, 0)
for i := 0; i < 4096; i++ {
if !limiter.allow(strconv.Itoa(i), now, time.Minute, 1) {
t.Fatalf("rejected window %d before capacity", i)
}
}
if limiter.allow("overflow", now, time.Minute, 1) {
t.Fatal("accepted a limiter key beyond its bounded capacity")
}
if !limiter.allow("after-expiry", now.Add(time.Minute), time.Minute, 1) {
t.Fatal("did not clean expired limiter windows")
}
}
func TestCleanupRetainsUsedRefreshForReplayUntilFamilyExpiry(t *testing.T) {
service, store := testService(t)
tokens := completeAndPoll(t, service, store, "cleanup-device", "discord", "https://discord.com", "123456789")
if _, err := store.db.Exec(`UPDATE refresh_tokens SET used_at=? WHERE token_hash=?`, testNowUnix, store.digest("refresh-token", tokens.RefreshToken)); err != nil {
t.Fatal(err)
}
tx, err := store.db.Begin()
if err != nil {
t.Fatal(err)
}
if err := cleanupExpired(tx, testNowUnix+int64((29*24*time.Hour).Seconds())); err != nil {
_ = tx.Rollback()
t.Fatal(err)
}
if err := tx.Commit(); err != nil {
t.Fatal(err)
}
var count int
if err := store.db.QueryRow(`SELECT COUNT(*) FROM refresh_tokens WHERE used_at IS NOT NULL`).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 1 {
t.Fatal("used refresh token was removed before family expiry")
}
tx, err = store.db.Begin()
if err != nil {
t.Fatal(err)
}
if err := cleanupExpired(tx, testNowUnix+int64((31*24*time.Hour).Seconds())); err != nil {
_ = tx.Rollback()
t.Fatal(err)
}
if err := tx.Commit(); err != nil {
t.Fatal(err)
}
if err := store.db.QueryRow(`SELECT COUNT(*) FROM refresh_tokens`).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatal("expired family refresh token was not cleaned")
}
}
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { return f(request) }
func jsonResponse(status int, body string) *http.Response {
return &http.Response{StatusCode: status, Body: io.NopCloser(strings.NewReader(body)), Header: make(http.Header)}
}
func TestProviderIdentityVerification(t *testing.T) {
service, _ := testService(t)
service.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) {
switch request.URL.Host + request.URL.Path {
case "discord.com/api/v10/oauth2/token":
return jsonResponse(http.StatusOK, `{"access_token":"provider-access"}`), nil
case "discord.com/api/v10/users/@me":
if request.Header.Get("Authorization") != "Bearer provider-access" {
t.Fatal("Discord bearer token missing")
}
return jsonResponse(http.StatusOK, `{"id":"123456789"}`), nil
case "oauth2.googleapis.com/token":
return jsonResponse(http.StatusOK, `{"access_token":"provider-access","id_token":"signed-id-token"}`), nil
case "oauth2.googleapis.com/tokeninfo":
return jsonResponse(http.StatusOK, `{"iss":"https://accounts.google.com","aud":"google-client","sub":"google-subject","nonce":"expected-nonce","exp":"1900000000"}`), nil
default:
t.Fatalf("unexpected provider request %s", request.URL)
return nil, nil
}
})}
discord, err := service.exchangeIdentity(context.Background(), "discord", "code", "verifier", "nonce")
if err != nil || discord.issuer != "https://discord.com" || discord.subject != "123456789" {
t.Fatalf("Discord identity=%+v err=%v", discord, err)
}
google, err := service.exchangeIdentity(context.Background(), "google", "code", "verifier", "expected-nonce")
if err != nil || google.issuer != "https://accounts.google.com" || google.subject != "google-subject" {
t.Fatalf("Google identity=%+v err=%v", google, err)
}
if _, err := service.exchangeIdentity(context.Background(), "google", "code", "verifier", "wrong-nonce"); err == nil {
t.Fatal("Google identity accepted the wrong OIDC nonce")
}
}
func TestProviderScopesUseLeastPrivilege(t *testing.T) {
if got := providerScope("discord"); got != "identify" {
t.Fatalf("Discord scope=%q, want identify", got)
}
if got := providerScope("google"); got != "openid" {
t.Fatalf("Google scope=%q, want openid", got)
}
}
func TestProviderErrorConsumesAuthorizationStateWithoutExchange(t *testing.T) {
service, store := testService(t)
insertAuthorizingDevice(t, store, "cancelled-device", "discord")
state := "cancelled-oauth-state"
if _, err := store.db.Exec(`UPDATE devices SET state_hash=? WHERE id='cancelled-device'`, store.digest("oauth-state", state)); err != nil {
t.Fatal(err)
}
service.client = &http.Client{Transport: roundTripFunc(func(request *http.Request) (*http.Response, error) {
t.Fatal("provider error callback attempted a token exchange")
return nil, nil
})}
request := httptest.NewRequest(http.MethodGet, "/auth/discord/callback?error=access_denied&state="+url.QueryEscape(state), nil)
response := httptest.NewRecorder()
service.Handler().ServeHTTP(response, request)
if response.Code != http.StatusBadRequest {
t.Fatalf("status=%d body=%q", response.Code, response.Body.String())
}
var status, errorCode string
var stateHash, verifier, nonce []byte
if err := store.db.QueryRow(`SELECT status,error_code,COALESCE(state_hash,X''),COALESCE(verifier_cipher,X''),COALESCE(nonce_cipher,X'') FROM devices WHERE id='cancelled-device'`).
Scan(&status, &errorCode, &stateHash, &verifier, &nonce); err != nil {
t.Fatal(err)
}
if status != "failed" || errorCode != "provider_cancelled" || len(stateHash) != 0 || len(verifier) != 0 || len(nonce) != 0 {
t.Fatalf("cancelled device status=%q error=%q state=%d verifier=%d nonce=%d", status, errorCode, len(stateHash), len(verifier), len(nonce))
}
}
func TestCompleteDeviceRejectsMalformedProviderIdentity(t *testing.T) {
service, store := testService(t)
insertAuthorizingDevice(t, store, "bad-identity", "discord")
if err := service.completeDevice("bad-identity", "discord", providerIdentity{issuer: "https://discord.com", subject: "not-a-snowflake"}); err == nil {
t.Fatal("accepted malformed Discord identity")
}
var count int
if err := store.db.QueryRow(`SELECT COUNT(*) FROM identities`).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatal("malformed identity was persisted")
}
}
+166
View File
@@ -0,0 +1,166 @@
package auth
import (
"context"
"crypto/aes"
"crypto/cipher"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/base64"
"errors"
"fmt"
"os"
"path/filepath"
"time"
_ "modernc.org/sqlite"
)
const schemaVersion = 1
var (
ErrUnauthorized = errors.New("auth: unauthorized")
ErrExpired = errors.New("auth: expired")
ErrPending = errors.New("auth: pending")
ErrConsumed = errors.New("auth: consumed")
ErrNotAllowed = errors.New("auth: identity is not allowed on this single-owner server")
)
type Store struct {
db *sql.DB
hashKey []byte
aead cipher.AEAD
now func() time.Time
}
func Open(path string, masterKey []byte) (*Store, error) {
if len(masterKey) != 32 {
return nil, errors.New("auth: master key must contain 32 bytes")
}
defer clear(masterKey)
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
return nil, fmt.Errorf("auth: create database directory: %w", err)
}
db, err := sql.Open("sqlite", filepath.Clean(path))
if err != nil {
return nil, fmt.Errorf("auth: open database: %w", err)
}
db.SetMaxOpenConns(1)
opened := false
defer func() {
if !opened {
_ = db.Close()
}
}()
if _, err = db.Exec(`PRAGMA foreign_keys=ON; PRAGMA journal_mode=WAL; PRAGMA synchronous=FULL; PRAGMA busy_timeout=5000`); err != nil {
return nil, fmt.Errorf("auth: configure database: %w", err)
}
if err = os.Chmod(filepath.Clean(path), 0o600); err != nil {
return nil, fmt.Errorf("auth: restrict database permissions: %w", err)
}
tx, err := db.BeginTx(context.Background(), nil)
if err != nil {
return nil, err
}
defer tx.Rollback()
statements := []string{
`CREATE TABLE IF NOT EXISTS metadata (key TEXT PRIMARY KEY NOT NULL, value TEXT NOT NULL) WITHOUT ROWID`,
`CREATE TABLE IF NOT EXISTS accounts (id TEXT PRIMARY KEY NOT NULL, status TEXT NOT NULL, created_at INTEGER NOT NULL, last_login_at INTEGER NOT NULL) WITHOUT ROWID`,
`CREATE TABLE IF NOT EXISTS identities (provider TEXT NOT NULL, issuer TEXT NOT NULL, subject_hash BLOB NOT NULL, account_id TEXT NOT NULL REFERENCES accounts(id), created_at INTEGER NOT NULL, last_login_at INTEGER NOT NULL, PRIMARY KEY(issuer,subject_hash)) WITHOUT ROWID`,
`CREATE TABLE IF NOT EXISTS devices (id TEXT PRIMARY KEY NOT NULL, client_hash BLOB NOT NULL, secret_hash BLOB NOT NULL, start_hash BLOB NOT NULL, provider TEXT NOT NULL, state_hash BLOB, verifier_cipher BLOB, nonce_cipher BLOB, result_cipher BLOB, status TEXT NOT NULL, error_code TEXT, created_at INTEGER NOT NULL, expires_at INTEGER NOT NULL) WITHOUT ROWID`,
`CREATE INDEX IF NOT EXISTS devices_client_pending ON devices(client_hash,status,expires_at)`,
`CREATE UNIQUE INDEX IF NOT EXISTS devices_state ON devices(state_hash) WHERE state_hash IS NOT NULL`,
`CREATE TABLE IF NOT EXISTS families (id TEXT PRIMARY KEY NOT NULL, account_id TEXT NOT NULL REFERENCES accounts(id), provider TEXT NOT NULL, created_at INTEGER NOT NULL, expires_at INTEGER NOT NULL, revoked_at INTEGER) WITHOUT ROWID`,
`CREATE TABLE IF NOT EXISTS refresh_tokens (token_hash BLOB PRIMARY KEY NOT NULL, family_id TEXT NOT NULL REFERENCES families(id), created_at INTEGER NOT NULL, expires_at INTEGER NOT NULL, used_at INTEGER, revoked_at INTEGER) WITHOUT ROWID`,
`CREATE TABLE IF NOT EXISTS access_tokens (token_hash BLOB PRIMARY KEY NOT NULL, family_id TEXT NOT NULL REFERENCES families(id), account_id TEXT NOT NULL REFERENCES accounts(id), created_at INTEGER NOT NULL, expires_at INTEGER NOT NULL, revoked_at INTEGER) WITHOUT ROWID`,
}
for _, statement := range statements {
if _, err = tx.Exec(statement); err != nil {
return nil, fmt.Errorf("auth: create schema: %w", err)
}
}
var version int
err = tx.QueryRow(`SELECT CAST(value AS INTEGER) FROM metadata WHERE key='schema_version'`).Scan(&version)
if errors.Is(err, sql.ErrNoRows) {
if _, err = tx.Exec(`INSERT INTO metadata(key,value) VALUES('schema_version',?)`, schemaVersion); err != nil {
return nil, err
}
version = schemaVersion
} else if err != nil {
return nil, err
}
if version != schemaVersion {
return nil, fmt.Errorf("auth: schema version %d, want %d", version, schemaVersion)
}
if err = tx.Commit(); err != nil {
return nil, err
}
encryptionKey := derive(masterKey, "auth-encryption")
hashKey := derive(masterKey, "auth-token-hmac")
block, err := aes.NewCipher(encryptionKey)
clear(encryptionKey)
if err != nil {
return nil, err
}
aead, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
opened = true
return &Store{db: db, hashKey: hashKey, aead: aead, now: time.Now}, nil
}
func (s *Store) Close() error {
clear(s.hashKey)
return s.db.Close()
}
func derive(master []byte, purpose string) []byte {
mac := hmac.New(sha256.New, master)
_, _ = mac.Write([]byte("bd2/" + purpose + "/v1"))
return mac.Sum(nil)
}
func (s *Store) digest(purpose, raw string) []byte {
mac := hmac.New(sha256.New, s.hashKey)
_, _ = mac.Write([]byte(purpose))
_, _ = mac.Write([]byte{'\x00'})
_, _ = mac.Write([]byte(raw))
return mac.Sum(nil)
}
func (s *Store) identityDigest(issuer, subject string) []byte {
mac := hmac.New(sha256.New, s.hashKey)
_, _ = mac.Write([]byte("identity\x00"))
_, _ = mac.Write([]byte(issuer))
_, _ = mac.Write([]byte{'\x00'})
_, _ = mac.Write([]byte(subject))
return mac.Sum(nil)
}
func randomToken(bytes int) (string, error) {
value := make([]byte, bytes)
if _, err := rand.Read(value); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(value), nil
}
func (s *Store) seal(id, field string, plain []byte) ([]byte, error) {
nonce := make([]byte, s.aead.NonceSize())
if _, err := rand.Read(nonce); err != nil {
return nil, err
}
aad := []byte("bd2/auth/v1/" + id + "/" + field)
return s.aead.Seal(nonce, nonce, plain, aad), nil
}
func (s *Store) open(id, field string, sealed []byte) ([]byte, error) {
if len(sealed) < s.aead.NonceSize() {
return nil, errors.New("auth: invalid ciphertext")
}
nonce, ciphertext := sealed[:s.aead.NonceSize()], sealed[s.aead.NonceSize():]
return s.aead.Open(nil, nonce, ciphertext, []byte("bd2/auth/v1/"+id+"/"+field))
}
+48
View File
@@ -0,0 +1,48 @@
package auth
import (
"bytes"
"path/filepath"
"testing"
)
func TestStoreRequiresAndClearsExactMasterKey(t *testing.T) {
if _, err := Open(filepath.Join(t.TempDir(), "short.db"), make([]byte, 31)); err == nil {
t.Fatal("store accepted a non-256-bit master key")
}
key := bytes.Repeat([]byte{0x7a}, 32)
store, err := Open(filepath.Join(t.TempDir(), "auth.db"), key)
if err != nil {
t.Fatal(err)
}
defer store.Close()
for index, value := range key {
if value != 0 {
t.Fatalf("master key byte %d was retained by the caller buffer", index)
}
}
}
func TestStorePersistsExplicitSchemaVersionAndRejectsUnknownVersion(t *testing.T) {
path := filepath.Join(t.TempDir(), "auth.db")
store, err := Open(path, bytes.Repeat([]byte{0x35}, 32))
if err != nil {
t.Fatal(err)
}
var version int
if err := store.db.QueryRow(`SELECT CAST(value AS INTEGER) FROM metadata WHERE key='schema_version'`).Scan(&version); err != nil {
t.Fatal(err)
}
if version != schemaVersion {
t.Fatalf("schema_version=%d, want %d", version, schemaVersion)
}
if _, err := store.db.Exec(`UPDATE metadata SET value='999' WHERE key='schema_version'`); err != nil {
t.Fatal(err)
}
if err := store.Close(); err != nil {
t.Fatal(err)
}
if _, err := Open(path, bytes.Repeat([]byte{0x35}, 32)); err == nil {
t.Fatal("store accepted an unknown schema version")
}
}
+175
View File
@@ -0,0 +1,175 @@
// Package authconfig loads the server-authoritative authentication policy.
package authconfig
import (
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net/url"
"os"
"path/filepath"
"strings"
"time"
)
const FileName = "authentication.json"
// TODO: Replace the executable-adjacent file with a small authenticated
// configuration API while keeping this policy server-authoritative.
type Config struct {
Mode string `json:"mode"`
PublicURL string `json:"public_url,omitempty"`
MasterKeyEnv string `json:"master_key_env,omitempty"`
Providers map[string]ProviderConfig `json:"providers,omitempty"`
Session SessionConfig `json:"session,omitempty"`
}
type ProviderConfig struct {
ClientID string `json:"client_id"`
ClientSecretEnv string `json:"client_secret_env"`
}
type SessionConfig struct {
AccessTTL string `json:"access_ttl,omitempty"`
RefreshTTL string `json:"refresh_ttl,omitempty"`
DeviceTransactionTTL string `json:"device_transaction_ttl,omitempty"`
}
type Runtime struct {
Config
PublicURLParsed *url.URL
MasterKey []byte
ProviderSecrets map[string]string
AccessTTL time.Duration
RefreshTTL time.Duration
DeviceTTL time.Duration
}
type Public struct {
Mode string `json:"mode"`
Providers []string `json:"providers"`
}
func Load(path string) (Config, error) {
path, err := filepath.Abs(filepath.Clean(path))
if err != nil {
return Config{}, fmt.Errorf("authconfig: resolve path: %w", err)
}
file, err := os.Open(path)
if err != nil {
return Config{}, fmt.Errorf("authconfig: open %s: %w", path, err)
}
defer file.Close()
var config Config
decoder := json.NewDecoder(file)
decoder.DisallowUnknownFields()
if err := decoder.Decode(&config); err != nil {
return Config{}, fmt.Errorf("authconfig: decode %s: %w", path, err)
}
var trailing any
if err := decoder.Decode(&trailing); err == nil {
return Config{}, fmt.Errorf("authconfig: trailing JSON in %s", path)
} else if !errors.Is(err, io.EOF) {
return Config{}, fmt.Errorf("authconfig: trailing data in %s: %w", path, err)
}
if err := config.Validate(); err != nil {
return Config{}, fmt.Errorf("authconfig: %s: %w", path, err)
}
return config, nil
}
func BesideExecutable() (string, error) {
executable, err := os.Executable()
if err != nil {
return "", fmt.Errorf("authconfig: resolve executable: %w", err)
}
return filepath.Join(filepath.Dir(executable), FileName), nil
}
func (c Config) Validate() error {
if c.Mode != "local" && c.Mode != "oauth" {
return errors.New("mode must be local or oauth")
}
if c.Mode == "local" {
if c.PublicURL != "" || c.MasterKeyEnv != "" || len(c.Providers) != 0 {
return errors.New("local mode must not configure OAuth")
}
return validateTTLs(c.Session)
}
if c.MasterKeyEnv == "" || c.PublicURL == "" || len(c.Providers) == 0 {
return errors.New("oauth mode requires public_url, master_key_env, and providers")
}
publicURL, err := url.Parse(c.PublicURL)
if err != nil || publicURL.Host == "" || publicURL.User != nil || publicURL.RawQuery != "" || publicURL.Fragment != "" || publicURL.Path != "" {
return errors.New("public_url must be an absolute origin without path, query, fragment, or user info")
}
localhost := publicURL.Hostname() == "127.0.0.1" || publicURL.Hostname() == "localhost" || publicURL.Hostname() == "::1"
if publicURL.Scheme != "https" && !(localhost && publicURL.Scheme == "http") {
return errors.New("public_url must use HTTPS except on localhost")
}
for name, provider := range c.Providers {
if name != "discord" && name != "google" {
return fmt.Errorf("unsupported provider %q", name)
}
if provider.ClientID == "" || provider.ClientSecretEnv == "" {
return fmt.Errorf("provider %q requires client_id and client_secret_env", name)
}
}
return validateTTLs(c.Session)
}
func validateTTLs(session SessionConfig) error {
for name, raw := range map[string]string{"access_ttl": session.AccessTTL, "refresh_ttl": session.RefreshTTL, "device_transaction_ttl": session.DeviceTransactionTTL} {
if raw != "" {
if value, err := time.ParseDuration(raw); err != nil || value <= 0 {
return fmt.Errorf("%s must be a positive Go duration", name)
}
}
}
return nil
}
func (c Config) ResolveEnvironment() (Runtime, error) {
if err := c.Validate(); err != nil {
return Runtime{}, err
}
runtime := Runtime{Config: c, ProviderSecrets: make(map[string]string), AccessTTL: 15 * time.Minute, RefreshTTL: 30 * 24 * time.Hour, DeviceTTL: 10 * time.Minute}
if c.Session.AccessTTL != "" {
runtime.AccessTTL, _ = time.ParseDuration(c.Session.AccessTTL)
}
if c.Session.RefreshTTL != "" {
runtime.RefreshTTL, _ = time.ParseDuration(c.Session.RefreshTTL)
}
if c.Session.DeviceTransactionTTL != "" {
runtime.DeviceTTL, _ = time.ParseDuration(c.Session.DeviceTransactionTTL)
}
if c.Mode == "local" {
return runtime, nil
}
runtime.PublicURLParsed, _ = url.Parse(c.PublicURL)
key, err := base64.StdEncoding.DecodeString(os.Getenv(c.MasterKeyEnv))
if err != nil || len(key) != 32 {
return Runtime{}, fmt.Errorf("authconfig: %s must contain a base64-encoded 32-byte key", c.MasterKeyEnv)
}
runtime.MasterKey = key
for name, provider := range c.Providers {
secret := os.Getenv(provider.ClientSecretEnv)
if strings.TrimSpace(secret) == "" {
return Runtime{}, fmt.Errorf("authconfig: provider %s secret environment %s is empty", name, provider.ClientSecretEnv)
}
runtime.ProviderSecrets[name] = secret
}
return runtime, nil
}
func (c Config) Public() Public {
view := Public{Mode: c.Mode, Providers: []string{}}
for _, name := range []string{"discord", "google"} {
if _, ok := c.Providers[name]; ok {
view.Providers = append(view.Providers, name)
}
}
return view
}
@@ -0,0 +1,99 @@
package authconfig
import (
"encoding/base64"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
)
func TestLoad(t *testing.T) {
path := filepath.Join(t.TempDir(), FileName)
if err := os.WriteFile(path, []byte(`{"mode":"oauth","public_url":"https://example.com","master_key_env":"MASTER","providers":{"discord":{"client_id":"d","client_secret_env":"DS"},"google":{"client_id":"g","client_secret_env":"GS"}}}`), 0o600); err != nil {
t.Fatal(err)
}
config, err := Load(path)
if err != nil {
t.Fatal(err)
}
if config.Mode != "oauth" || len(config.Providers) != 2 || config.Providers["discord"].ClientID != "d" || config.Providers["google"].ClientID != "g" {
t.Fatalf("unexpected config: %+v", config)
}
}
func TestValidateRejectsUnsafePolicies(t *testing.T) {
for name, config := range map[string]Config{
"unknown mode": {Mode: "disabled"},
"local providers": {Mode: "local", Providers: map[string]ProviderConfig{"discord": {ClientID: "d", ClientSecretEnv: "DS"}}},
"empty oauth": {Mode: "oauth"},
"unknown provider": {Mode: "oauth", PublicURL: "https://example.com", MasterKeyEnv: "MASTER", Providers: map[string]ProviderConfig{"github": {ClientID: "g", ClientSecretEnv: "GS"}}},
} {
t.Run(name, func(t *testing.T) {
if err := config.Validate(); err == nil {
t.Fatal("accepted invalid authentication policy")
}
})
}
}
func TestResolveEnvironmentKeepsSecretsOutOfPublicView(t *testing.T) {
master := make([]byte, 32)
for i := range master {
master[i] = byte(i + 1)
}
t.Setenv("AUTH_MASTER", base64.StdEncoding.EncodeToString(master))
t.Setenv("DISCORD_SECRET", "private-discord-secret")
config := Config{
Mode: "oauth",
PublicURL: "https://example.com",
MasterKeyEnv: "AUTH_MASTER",
Providers: map[string]ProviderConfig{
"discord": {ClientID: "public-client-id", ClientSecretEnv: "DISCORD_SECRET"},
},
}
runtime, err := config.ResolveEnvironment()
if err != nil {
t.Fatal(err)
}
if len(runtime.MasterKey) != 32 || runtime.ProviderSecrets["discord"] != "private-discord-secret" {
t.Fatal("runtime did not resolve authentication secrets")
}
publicJSON, err := json.Marshal(config.Public())
if err != nil {
t.Fatal(err)
}
for _, forbidden := range []string{"AUTH_MASTER", "DISCORD_SECRET", "private-discord-secret", "public-client-id"} {
if strings.Contains(string(publicJSON), forbidden) {
t.Fatalf("public authentication view leaked %q: %s", forbidden, publicJSON)
}
}
}
func TestValidateRejectsUnsafePublicURLsAndTTLs(t *testing.T) {
base := Config{
Mode: "oauth",
PublicURL: "https://example.com",
MasterKeyEnv: "MASTER",
Providers: map[string]ProviderConfig{
"discord": {ClientID: "d", ClientSecretEnv: "DS"},
},
}
for name, mutate := range map[string]func(*Config){
"http public": func(c *Config) { c.PublicURL = "http://example.com" },
"path": func(c *Config) { c.PublicURL = "https://example.com/auth" },
"query": func(c *Config) { c.PublicURL = "https://example.com?x=y" },
"missing secret": func(c *Config) { c.Providers["discord"] = ProviderConfig{ClientID: "d"} },
"invalid ttl": func(c *Config) { c.Session.AccessTTL = "0s" },
} {
t.Run(name, func(t *testing.T) {
candidate := base
candidate.Providers = map[string]ProviderConfig{"discord": base.Providers["discord"]}
mutate(&candidate)
if err := candidate.Validate(); err == nil {
t.Fatal("accepted unsafe authentication configuration")
}
})
}
}
@@ -9,9 +9,9 @@ import (
"log/slog"
"sync"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/wire"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/wire"
)
type Service struct {
@@ -3,10 +3,10 @@ package battle
import (
"testing"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/wire"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/wire"
)
func request(seq uint64) []byte { return wire.AppendVarint(nil, 1, seq) }
@@ -125,7 +125,7 @@ func TestBattleEnterUsesSamePictorialSnapshotAsAllCharRefresh(t *testing.T) {
func TestBattleVictoryLocksPackAtEnterForRewardsAndIdentity(t *testing.T) {
storage := stateio.NewMemory()
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.34.13"})
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.35.10"})
if err != nil {
t.Fatal(err)
}
@@ -8,7 +8,7 @@ import (
"strings"
"time"
"bd2server/internal/wire"
"bd2server/internal/server/wire"
)
type Config struct {
@@ -23,8 +23,9 @@ type Config struct {
func (c Config) Validate() error {
for _, pair := range [][2]string{{"game server", c.BaseURL}, {"CDN", c.CDNURL}} {
u, err := url.Parse(pair[1])
if err != nil || u.Scheme != "http" || u.Host == "" {
return fmt.Errorf("%s needs a valid local HTTP URL: %q", pair[0], pair[1])
if err != nil || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" ||
(u.Scheme != "http" && u.Scheme != "https") {
return fmt.Errorf("%s needs a valid HTTP(S) URL: %q", pair[0], pair[1])
}
}
if !strings.HasSuffix(c.BaseURL, "/") {
@@ -4,8 +4,8 @@ import (
"testing"
"time"
"bd2server/internal/versionconfig"
"bd2server/internal/wire"
"bd2server/internal/server/versionconfig"
"bd2server/internal/server/wire"
)
func TestMaintenance(t *testing.T) {
@@ -49,6 +49,16 @@ func TestServerInfoNoOfficialEndpoints(t *testing.T) {
}
}
func TestConfigAcceptsHTTPSPublicOrigin(t *testing.T) {
cfg := Config{
BaseURL: "https://bd2.example.com/game/", CDNURL: "https://bd2.example.com/assets/ServerData",
Version: "client", BundleVer: "bundle",
}
if err := cfg.Validate(); err != nil {
t.Fatalf("HTTPS self-hosted server config rejected: %v", err)
}
}
func TestServerInfoIncludesLocalGameData(t *testing.T) {
c := Config{
BaseURL: "http://127.0.0.1:8080/game/",
@@ -1,5 +1,6 @@
// Package dbcrypt implements the page cipher shared by Intro and GameData
// SQLite files in Brown Dust II 2.34.13.
// Package dbcrypt implements the page cipher used by versioned GameData
// SQLite archives. Client resource patching owns a separate copy so the pure
// server dependency graph never imports internal/client.
package dbcrypt
import (
@@ -11,10 +11,10 @@ import (
"sort"
"sync"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/versionconfig"
"bd2server/internal/wire"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/versionconfig"
"bd2server/internal/server/wire"
)
type DeckEntry struct {
@@ -90,7 +90,7 @@ func LoadSeed(path string) (Seed, error) {
return s, nil
}
func (s Seed) validate() error {
if s.Version != versionconfig.Protocol() {
if s.Version != versionconfig.State() {
return errors.New("deck: wrong seed version")
}
return validField(s.FieldDeck)
@@ -143,7 +143,7 @@ func NewStore(seed Seed) (*Store, error) {
if e := seed.validate(); e != nil {
return nil, e
}
return &Store{state: state{Version: versionconfig.Protocol(), FieldDeck: append([]FieldEntry(nil), seed.FieldDeck...), FieldCharControlDeckType: seed.FieldCharControlDeckType, AutoReviveCatalyst: seed.AutoReviveCatalyst, Waypoints: map[uint64]uint64{}, Costumes: map[uint64]uint64{}, Packs: map[uint64]uint64{}}, presets: map[uint64]Preset{}, presetSlots: presetBaseCount, costumeSettings: map[uint64]CostumeSetting{}, replies: map[string]deckReply{}}, nil
return &Store{state: state{Version: versionconfig.State(), FieldDeck: append([]FieldEntry(nil), seed.FieldDeck...), FieldCharControlDeckType: seed.FieldCharControlDeckType, AutoReviveCatalyst: seed.AutoReviveCatalyst, Waypoints: map[uint64]uint64{}, Costumes: map[uint64]uint64{}, Packs: map[uint64]uint64{}}, presets: map[uint64]Preset{}, presetSlots: presetBaseCount, costumeSettings: map[uint64]CostumeSetting{}, replies: map[string]deckReply{}}, nil
}
func OpenStore(storage stateio.Store, seed Seed) (*Store, error) {
s, e := NewStore(seed)
@@ -172,7 +172,7 @@ func OpenStore(storage stateio.Store, seed Seed) (*Store, error) {
if e = json.Unmarshal(b, &loaded); e != nil {
return nil, fmt.Errorf("deck: malformed state: %w", e)
}
if loaded.Version != versionconfig.Protocol() || (len(loaded.Deck) != 0 && validDeck(loaded.Deck) != nil) || validField(loaded.FieldDeck) != nil || loaded.Waypoints == nil || loaded.Costumes == nil || loaded.Packs == nil {
if loaded.Version != versionconfig.State() || (len(loaded.Deck) != 0 && validDeck(loaded.Deck) != nil) || validField(loaded.FieldDeck) != nil || loaded.Waypoints == nil || loaded.Costumes == nil || loaded.Packs == nil {
return nil, errors.New("deck: invalid saved state")
}
s.state = loaded
@@ -1,17 +1,17 @@
package deck
import (
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/versionconfig"
"bd2server/internal/wire"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/versionconfig"
"bd2server/internal/server/wire"
"path/filepath"
"testing"
)
func seeded(t *testing.T) *Store {
t.Helper()
x, e := LoadSeed(filepath.Join("..", "..", "seed", "v2_34_13", "decks.json"))
x, e := LoadSeed(filepath.Join("..", "..", "..", "seed", "v2_35_10", "decks.json"))
if e != nil {
t.Fatal(e)
}
@@ -38,7 +38,7 @@ func triple(a, b, c uint64) []byte {
func attachFormationOwnership(t *testing.T, store *Store) {
t.Helper()
starter := &player.Starter{
Version: versionconfig.Protocol(),
Version: versionconfig.State(),
Characters: []player.Character{
{InvenIndex: 101, ID: 350, Level: 1},
{InvenIndex: 102, ID: 351, Level: 1},
@@ -103,7 +103,7 @@ func TestFieldDeckSeedAndSave(t *testing.T) {
}
}
func TestDeckPersistenceAndCommands(t *testing.T) {
seed, e := LoadSeed(filepath.Join("..", "..", "seed", "v2_34_13", "decks.json"))
seed, e := LoadSeed(filepath.Join("..", "..", "..", "seed", "v2_35_10", "decks.json"))
if e != nil {
t.Fatal(e)
}
@@ -346,7 +346,7 @@ func TestDeckCharAutoReviveUsesCurrentFormationWithoutInventingRevives(t *testin
}
func TestDeckCharAutoRevivePreservesExplicitZeroCatalyst(t *testing.T) {
seed, err := LoadSeed(filepath.Join("..", "..", "seed", "v2_34_13", "decks.json"))
seed, err := LoadSeed(filepath.Join("..", "..", "..", "seed", "v2_35_10", "decks.json"))
if err != nil {
t.Fatal(err)
}
@@ -11,9 +11,9 @@ import (
"unicode"
"unicode/utf8"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/wire"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/wire"
)
const (
@@ -5,9 +5,9 @@ import (
"sort"
"testing"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/wire"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/wire"
)
type presetFixture struct {
@@ -24,11 +24,11 @@ type presetFixture struct {
func newPresetFixture(t *testing.T) *presetFixture {
t.Helper()
storage := stateio.NewMemory()
seed, err := LoadSeed("../../seed/v2_34_13/decks.json")
seed, err := LoadSeed("../../../seed/v2_35_10/decks.json")
if err != nil {
t.Fatal(err)
}
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.34.13"})
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.35.10"})
if err != nil {
t.Fatal(err)
}
@@ -1,6 +1,6 @@
package feature
import "bd2server/internal/wire"
import "bd2server/internal/server/wire"
// initialResponses are locally constructed, typed defaults for the 2.34.13
// new-player account. These do not reuse recorded response bytes. Stateful
@@ -5,14 +5,14 @@ import (
"path/filepath"
"testing"
"bd2server/internal/fixture"
"bd2server/internal/wire"
"bd2server/internal/server/fixture"
"bd2server/internal/server/wire"
)
// TestCaptureCompatibility is an optional development-time protocol audit.
// Normal server execution never opens the capture.
func TestCaptureCompatibility(t *testing.T) {
root := filepath.Join("..", "..", "..", "data", "capture", "2.34.13", "20260920-003254")
root := filepath.Join("..", "..", "..", "..", "data", "capture", "2.34.13", "20260920-003254")
set, err := fixture.Load(root)
if err != nil {
t.Skipf("optional capture unavailable: %v", err)
@@ -9,7 +9,7 @@ import (
"errors"
"fmt"
"bd2server/internal/wire"
"bd2server/internal/server/wire"
)
// ErrInvalidRequest means a known endpoint was sent a malformed request.
@@ -4,7 +4,7 @@ import (
"errors"
"testing"
"bd2server/internal/wire"
"bd2server/internal/server/wire"
)
func TestHandleAuditedEmptyResponses(t *testing.T) {
@@ -5,11 +5,11 @@ import (
"path/filepath"
"testing"
"bd2server/internal/wire"
"bd2server/internal/server/wire"
)
func TestDebugBattleStartResponse(t *testing.T) {
set, err := Load(filepath.Join("..", "..", "..", "data", "capture", "2.34.13", "20260920-003254"))
set, err := Load(filepath.Join("..", "..", "..", "..", "data", "capture", "2.34.13", "20260920-003254"))
if err != nil {
t.Skip(err)
}
@@ -17,7 +17,7 @@ import (
"strconv"
"strings"
"bd2server/internal/cryptox"
"bd2server/internal/server/cryptox"
)
const CaptureVersion = "2.34.13"
@@ -8,13 +8,13 @@ import (
"path/filepath"
"testing"
"bd2server/internal/cryptox"
"bd2server/internal/wire"
"bd2server/internal/server/cryptox"
"bd2server/internal/server/wire"
)
func captureRoot(t *testing.T) string {
t.Helper()
root, err := filepath.Abs(filepath.Join("..", "..", "..", "data", "capture", "2.34.13", "20260920-003254"))
root, err := filepath.Abs(filepath.Join("..", "..", "..", "..", "data", "capture", "2.34.13", "20260920-003254"))
if err != nil {
t.Fatal(err)
}
@@ -9,9 +9,9 @@ import (
"sync"
"time"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/wire"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/wire"
)
const infiniteGrant = "cash-product:1100001:9100033"
@@ -5,10 +5,10 @@ import (
"testing"
"time"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/wire"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/wire"
)
func TestTicketOnlyEquipmentDrawUsesGameDataAndNoScheduleAccounting(t *testing.T) {
@@ -38,7 +38,7 @@ func TestTicketOnlyEquipmentDrawUsesGameDataAndNoScheduleAccounting(t *testing.T
if err != nil {
t.Fatal(err)
}
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.34.13"})
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.35.10"})
if err != nil {
t.Fatal(err)
}
@@ -1106,7 +1106,7 @@ func TestMoonriseSelectionCashProductAndOneTimeTicketDraw(t *testing.T) {
if err != nil {
t.Fatal(err)
}
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.34.13"})
inventory, err := player.OpenInventory(storage, &player.Starter{Version: "2.35.10"})
if err != nil {
t.Fatal(err)
}
@@ -4,11 +4,11 @@ import (
"path/filepath"
"testing"
"bd2server/internal/account"
"bd2server/internal/accountstate"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/wire"
"bd2server/internal/server/account"
"bd2server/internal/server/accountstate"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/wire"
)
func TestLoginPurchaseCountsRestoredFromSQLiteGrant(t *testing.T) {
@@ -101,7 +101,7 @@ func TestLoginPurchaseCountsRestoredFromSQLiteGrant(t *testing.T) {
stale := wire.AppendVarint(nil, 1, 999)
userTemplate := wire.AppendVarint(nil, 1, 42)
userTemplate = wire.AppendBytes(userTemplate, 26, stale)
loginSeed := &account.LoginSeed{Version: account.ProtocolVersion(), PacketCode: 11, UserInfo: userTemplate}
loginSeed := &account.LoginSeed{Version: account.StateVersion(), PacketCode: 11, UserInfo: userTemplate}
if err := loginSeed.AttachPurchaseCounts(service); err != nil {
t.Fatal(err)
}
@@ -3,10 +3,10 @@ package gacha
import (
"testing"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/wire"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/wire"
)
func TestGachaPointExchangeGrantsUpgradesOverflowsAndRetries(t *testing.T) {
@@ -15,7 +15,7 @@ import (
// source file from acquiring a runtime dependency on the fixture reader or an
// on-disk capture path as gacha evolves.
func TestServerRuntimeHasNoCaptureFixtureDependency(t *testing.T) {
root := filepath.Join("..", "..")
root := filepath.Join("..", "..", "..")
for _, directory := range []string{filepath.Join(root, "cmd"), filepath.Join(root, "internal")} {
err := filepath.WalkDir(directory, func(path string, entry fs.DirEntry, walkErr error) error {
if walkErr != nil {
@@ -6,14 +6,14 @@ import (
"path/filepath"
"testing"
"bd2server/internal/gamedata"
"bd2server/internal/player"
"bd2server/internal/stateio"
"bd2server/internal/wire"
"bd2server/internal/server/gamedata"
"bd2server/internal/server/player"
"bd2server/internal/server/stateio"
"bd2server/internal/server/wire"
)
func TestVersionedScheduleMatchesOfficial23510GachaInfo(t *testing.T) {
seed, err := LoadScheduleSeed(filepath.Join("..", "..", "seed", "v2_34_13", "gacha_schedule.json"), "2.35.10")
seed, err := LoadScheduleSeed(filepath.Join("..", "..", "..", "seed", "v2_35_10", "gacha_schedule.json"), "2.35.10")
if err != nil {
t.Fatal(err)
}
@@ -53,13 +53,13 @@ func TestScheduleSeedStrictValidation(t *testing.T) {
if _, err := LoadScheduleSeed(path, "2.35.10"); err == nil {
t.Fatal("unknown schedule field was accepted")
}
if _, err := LoadScheduleSeed(filepath.Join("..", "..", "seed", "v2_34_13", "gacha_schedule.json"), "2.34.13"); err == nil {
if _, err := LoadScheduleSeed(filepath.Join("..", "..", "..", "seed", "v2_35_10", "gacha_schedule.json"), "0.0.0"); err == nil {
t.Fatal("wrong client version was accepted")
}
}
func TestGachaInfoUsesInjectedScheduleAndEmptyAccountHasNoPreview(t *testing.T) {
seed, err := LoadScheduleSeed(filepath.Join("..", "..", "seed", "v2_34_13", "gacha_schedule.json"), "2.35.10")
seed, err := LoadScheduleSeed(filepath.Join("..", "..", "..", "seed", "v2_35_10", "gacha_schedule.json"), "2.35.10")
if err != nil {
t.Fatal(err)
}
@@ -7,7 +7,7 @@ import (
"os"
"path/filepath"
"bd2server/internal/wire"
"bd2server/internal/server/wire"
_ "modernc.org/sqlite"
)
@@ -36,7 +36,7 @@ type CharAwakeCharacterStage struct {
MaximumLevel uint64
}
// CharAwakeDesign is immutable 2.34.13 design data. Account progress remains
// CharAwakeDesign is immutable 2.35.10 design data. Account progress remains
// in CollectionStore and is indexed by UniqueCharId, as CharAwakeDBInfo is.
type CharAwakeDesign struct {
Characters map[uint64]CharAwakeCharacter
@@ -9,11 +9,11 @@ import (
"path/filepath"
"strings"
"bd2server/internal/dbcrypt"
"bd2server/internal/server/dbcrypt"
)
// DatabaseName maps a logical client DB name to its GameData archive entry.
// Version 1 is the current 2.34.13 DB schema generation.
// Version 1 is the current 2.35.10 DB schema generation.
func DatabaseName(logical string) (string, error) {
if logical == "" || strings.ContainsAny(logical, `/\\.`) {
return "", fmt.Errorf("gamedata: invalid logical database name %q", logical)
@@ -32,7 +32,7 @@ func ReadDatabase(root, version, logical string) ([]byte, error) {
return readEntry(root, version, name, logical)
}
// questDatabaseEntry is the member in the 2.34.13 GameData archive that
// questDatabaseEntry is the member in the 2.35.10 GameData archive that
// contains the shared QuestTable* SQLite database, including QuestTable21.
// Quest data is not stored in an individual per-pack database.
const questDatabaseEntry = "9F251C63BC72551C681EE75D328FA090D56E444B"
@@ -7,7 +7,7 @@ import (
"path/filepath"
"testing"
"bd2server/internal/dbcrypt"
"bd2server/internal/server/dbcrypt"
)
// These tests cover the production archive/decryption contract. Interactive

Some files were not shown because too many files have changed in this diff Show More