feat(all): split client tooling and add OAuth server login
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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]"
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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())
|
||||
}
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
|
||||
"bd2server/internal/accountstate"
|
||||
"bd2server/internal/server/accountstate"
|
||||
)
|
||||
|
||||
func stateCommand(args []string) error {
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"bd2server/internal/accountstate"
|
||||
"bd2server/internal/server/accountstate"
|
||||
)
|
||||
|
||||
func stateProblemsError(prefix string, problems []accountstate.Problem) error {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
-1
@@ -1,6 +1,6 @@
|
||||
//go:build !windows
|
||||
|
||||
package clientplugin
|
||||
package config
|
||||
|
||||
import "os"
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
//go:build windows
|
||||
|
||||
package clientplugin
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
+1
-1
@@ -4,7 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"bd2server/internal/stateio"
|
||||
"bd2server/internal/server/stateio"
|
||||
)
|
||||
|
||||
var _ stateio.AtomicEntryStore = (*Repository)(nil)
|
||||
+1
-1
@@ -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)
|
||||
+1
-1
@@ -13,7 +13,7 @@ import (
|
||||
"strconv"
|
||||
"sync"
|
||||
|
||||
"bd2server/internal/stateio"
|
||||
"bd2server/internal/server/stateio"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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, "/") {
|
||||
+12
-2
@@ -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
|
||||
+3
-3
@@ -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) {
|
||||
+2
-2
@@ -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)
|
||||
}
|
||||
+6
-6
@@ -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)
|
||||
}
|
||||
+4
-4
@@ -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) {
|
||||
+1
-1
@@ -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"
|
||||
+1
-1
@@ -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
Reference in New Issue
Block a user