Files
bd2/go/internal/server/session/server.go
T

448 lines
13 KiB
Go

// Package session owns authentication, encryption, and request routing for a
// local account. It deliberately has no capture/fixture dependency.
package session
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log/slog"
"strings"
"sync"
"time"
"bd2server/internal/server/cryptox"
"bd2server/internal/server/progress"
"bd2server/internal/server/protocol"
"bd2server/internal/server/stateio"
"bd2server/internal/server/transport"
"bd2server/internal/server/wire"
)
const (
gameSessionTTL = 24 * time.Hour
maxGameSessions = 1024
maxCookieHeaderLen = 8 << 10
)
var errSessionRequired = errors.New("session login required")
type LoginService interface {
Login(request, sessionKey []byte) ([]byte, error)
}
type LoginAuthenticator interface {
AuthenticateLogin(request []byte) (accountID string, err error)
}
// Handler is implemented by domain services. ok=false means the endpoint is
// not owned by that service; unknown endpoints fail closed.
type Handler interface {
Handle(path string, request []byte) (packetCode int, response []byte, ok bool, err error)
}
// SessionAware handlers use a login-scoped opaque ID when protobuf request
// sequences participate in durable idempotency keys.
type SessionAware interface {
BeginSession(id string)
}
type gameSession struct {
key []byte
accountID string
id string
expiresAt time.Time
lastUsed time.Time
}
type Server struct {
mu sync.Mutex
sessions map[[sha256.Size]byte]*gameSession
latestSessionToken [sha256.Size]byte
latestSessionSet bool
login LoginService
handlers []Handler
progress *progress.Store
stateTx stateio.TransactionalStore
auth LoginAuthenticator
now func() time.Time
sessionTTL time.Duration
maxSessions int
}
func (s *Server) AttachLoginAuthenticator(authenticator LoginAuthenticator) error {
if authenticator == nil {
return errors.New("session login authenticator is nil")
}
s.mu.Lock()
defer s.mu.Unlock()
s.auth = authenticator
return nil
}
// AttachStateStore wraps each authenticated request (the complete batch for
// BatchRequest) in one durable account database transaction.
func (s *Server) AttachStateStore(store stateio.TransactionalStore) error {
if store == nil {
return errors.New("session state store is nil")
}
s.mu.Lock()
defer s.mu.Unlock()
s.stateTx = store
return nil
}
func NewServer(login LoginService, handlers ...Handler) (*Server, error) {
return NewServerWithProgress(login, progress.NewStore(), handlers...)
}
// NewServerWithProgress uses a caller-owned player progress store; the serve
// command supplies a file-backed store so changes survive server restarts.
func NewServerWithProgress(login LoginService, player *progress.Store, handlers ...Handler) (*Server, error) {
if login == nil {
return nil, errors.New("session login service is nil")
}
if player == nil {
return nil, errors.New("session progress store is nil")
}
return &Server{
sessions: make(map[[sha256.Size]byte]*gameSession), login: login,
handlers: append([]Handler(nil), handlers...), progress: player,
now: time.Now, sessionTTL: gameSessionTTL, maxSessions: maxGameSessions,
}, nil
}
func (s *Server) DispatchRaw(path string, body []byte, cookie string) (transport.RawReply, error) {
s.mu.Lock()
defer s.mu.Unlock()
if s.stateTx != nil {
if err := s.stateTx.Check(); err != nil {
return transport.RawReply{}, fmt.Errorf("account state unavailable: %w", err)
}
}
if path == "/LoginUser" {
now := s.now()
s.pruneSessions(now)
request, err := cryptox.DecryptBase64Payload(string(body), cryptox.Key())
if err != nil {
return transport.RawReply{}, fmt.Errorf("LoginUser decrypt: %w", err)
}
accountID := "local-owner"
if s.auth != nil {
accountID, err = s.auth.AuthenticateLogin(request)
if err != nil {
return transport.RawReply{}, fmt.Errorf("%w: %v", transport.ErrAccessCredentialInvalid, err)
}
if accountID == "" {
return transport.RawReply{}, errors.New("LoginUser authentication returned an empty account ID")
}
}
game, token, err := newGameSession(accountID, now, s.sessionTTL)
if err != nil {
return transport.RawReply{}, err
}
proto, err := s.login.Login(request, game.key)
if err != nil {
clear(game.key)
return transport.RawReply{}, fmt.Errorf("LoginUser: %w", err)
}
encoded, err := protocol.Encode(3, proto, cryptox.Key(), now.UnixMilli())
if err != nil {
clear(game.key)
return transport.RawReply{}, err
}
s.makeSessionRoom()
tokenKey := sessionTokenKey(token)
s.sessions[tokenKey] = game
s.latestSessionToken = tokenKey
s.latestSessionSet = true
s.activate(game)
return transport.RawReply{Body: encoded, Cookie: token}, nil
}
game, err := s.authorize(cookie)
if err != nil {
return transport.RawReply{}, err
}
s.activate(game)
if path == "/BatchRequest" {
return s.withStateTransaction(func() (transport.RawReply, error) {
return s.handleBatch(body, game.key)
})
}
request, err := cryptox.DecryptBase64Payload(string(body), game.key)
if err != nil {
return transport.RawReply{}, fmt.Errorf("%s decrypt: %w", path, err)
}
return s.withStateTransaction(func() (transport.RawReply, error) {
code, response, err := s.dispatch(path, request)
if err != nil {
return transport.RawReply{}, err
}
encoded, err := protocol.Encode(code, response, game.key, s.now().UnixMilli())
return transport.RawReply{Body: encoded}, err
})
}
func newGameSession(accountID string, now time.Time, ttl time.Duration) (*gameSession, string, error) {
keyBytes := make([]byte, 16)
if _, err := rand.Read(keyBytes); err != nil {
return nil, "", fmt.Errorf("create game session key: %w", err)
}
tokenBytes := make([]byte, 24)
if _, err := rand.Read(tokenBytes); err != nil {
return nil, "", fmt.Errorf("create game session token: %w", err)
}
idBytes := make([]byte, 12)
if _, err := rand.Read(idBytes); err != nil {
return nil, "", fmt.Errorf("create login request identity: %w", err)
}
return &gameSession{
key: []byte(hex.EncodeToString(keyBytes)),
accountID: accountID,
id: hex.EncodeToString(idBytes),
expiresAt: now.Add(ttl),
lastUsed: now,
}, hex.EncodeToString(tokenBytes) + "|1", nil
}
func sessionTokenKey(token string) [sha256.Size]byte { return sha256.Sum256([]byte(token)) }
func (s *Server) pruneSessions(now time.Time) {
for token, game := range s.sessions {
if !now.Before(game.expiresAt) {
s.deleteSession(token, game)
}
}
}
func (s *Server) makeSessionRoom() {
for len(s.sessions) >= s.maxSessions {
var oldestToken [sha256.Size]byte
var oldest *gameSession
for token, game := range s.sessions {
if oldest == nil || game.lastUsed.Before(oldest.lastUsed) {
oldestToken, oldest = token, game
}
}
if oldest == nil {
return
}
s.deleteSession(oldestToken, oldest)
}
}
func (s *Server) deleteSession(token [sha256.Size]byte, game *gameSession) {
delete(s.sessions, token)
for i := range game.key {
game.key[i] = 0
}
game.accountID = ""
game.id = ""
if s.latestSessionSet && token == s.latestSessionToken {
s.latestSessionSet = false
}
}
func (s *Server) activate(game *gameSession) {
for _, handler := range s.handlers {
if aware, ok := handler.(SessionAware); ok {
aware.BeginSession(game.id)
}
}
}
func (s *Server) withStateTransaction(run func() (transport.RawReply, error)) (reply transport.RawReply, err error) {
if s.stateTx == nil {
return run()
}
operation, err := s.stateTx.BeginOperation()
if err != nil {
return transport.RawReply{}, fmt.Errorf("begin account transaction: %w", err)
}
finished := false
defer func() {
if finished {
return
}
rollbackErr := operation.Rollback()
if recovered := recover(); recovered != nil {
panic(recovered)
}
if rollbackErr != nil {
err = errors.Join(err, rollbackErr)
}
}()
reply, err = run()
if err != nil {
rollbackErr := operation.Rollback()
finished = true
if rollbackErr != nil {
return transport.RawReply{}, errors.Join(err, rollbackErr)
}
return transport.RawReply{}, err
}
if err := operation.Commit(); err != nil {
finished = true
return transport.RawReply{}, fmt.Errorf("commit account transaction: %w", err)
}
finished = true
return reply, nil
}
func (s *Server) handleBatch(body, key []byte) (transport.RawReply, error) {
batchStarted := time.Now()
requests, decoded, err := protocol.DecodeBatchRequest(body, key)
if err != nil {
return transport.RawReply{}, err
}
items := make([]protocol.BatchResponse, 0, len(requests))
for i, request := range requests {
itemStarted := time.Now()
code, response, err := s.dispatch(request.Path, decoded[i])
if err != nil {
return transport.RawReply{}, fmt.Errorf("batch %s: %w", request.Path, err)
}
itemElapsed := time.Since(itemStarted)
if formationTimingPath(request.Path) {
slog.Info("formation batch item handled", "index", i, "path", request.Path, "duration_ms", float64(itemElapsed.Microseconds())/1000, "response_bytes", len(response))
} else if itemElapsed >= 100*time.Millisecond {
slog.Warn("slow batch item", "index", i, "path", request.Path, "duration_ms", float64(itemElapsed.Microseconds())/1000)
}
raw, err := protocol.Encode(code, response, key, time.Now().UnixMilli())
if err != nil {
return transport.RawReply{}, err
}
var envelope protocol.Envelope
if err := json.Unmarshal(raw, &envelope); err != nil {
return transport.RawReply{}, err
}
items = append(items, protocol.BatchResponse{Path: request.Path, ResponseData: envelope})
}
encoded, err := json.Marshal(items)
batchElapsed := time.Since(batchStarted)
if batchElapsed >= time.Second {
slog.Warn("slow batch request", "items", len(requests), "duration_ms", float64(batchElapsed.Microseconds())/1000, "responseBytes", len(encoded))
}
return transport.RawReply{Body: encoded}, err
}
func formationTimingPath(path string) bool {
return strings.HasPrefix(path, "/Preset") ||
strings.HasPrefix(path, "/Deck") ||
strings.HasPrefix(path, "/FieldDeck") ||
path == "/EquipBatchUse"
}
func (s *Server) dispatch(path string, request []byte) (int, []byte, error) {
if _, found, err := wire.Varint(request, 1); err != nil || !found {
return 0, nil, fmt.Errorf("%s has no request sequence", path)
}
switch path {
case "/SaveUserPosition":
if err := s.progress.SaveUserPosition(request); err != nil {
return 0, nil, fmt.Errorf("%s: %w", path, err)
}
if saved, found := s.progress.Position(); found {
slog.Info("field position saved", "pack", saved.PackID, "map", saved.Position.MapID)
}
return 7, nil, nil
case "/TutorialClear":
if err := s.progress.ClearTutorial(request); err != nil {
return 0, nil, fmt.Errorf("%s: %w", path, err)
}
return 102, nil, nil
case "/QuestUpdate":
questID, err := s.progress.UpdateQuest(request)
if err != nil {
return 0, nil, fmt.Errorf("%s: %w", path, err)
}
response := wire.AppendVarint(nil, 1, uint64(questID))
return 19, wire.AppendBytes(response, 2, nil), nil
}
for _, handler := range s.handlers {
code, response, ok, err := handler.Handle(path, request)
if err != nil {
return 0, nil, err
}
if ok {
return code, response, nil
}
}
return 0, nil, fmt.Errorf("%w: %s", transport.ErrNotImplemented, path)
}
func (s *Server) authorize(cookie string) (*gameSession, error) {
token, err := parseSessionCookie(cookie)
if err != nil {
if errors.Is(err, errSessionRequired) {
return nil, fmt.Errorf("%w: %v", transport.ErrGameSessionExpired, err)
}
return nil, err
}
now := s.now()
s.pruneSessions(now)
game, ok := s.sessions[sessionTokenKey(token)]
if !ok {
return nil, fmt.Errorf("%w: cookie no longer names a live session", transport.ErrGameSessionExpired)
}
game.lastUsed = now
return game, nil
}
func parseSessionCookie(cookie string) (string, error) {
if cookie == "" {
return "", errSessionRequired
}
if len(cookie) > maxCookieHeaderLen {
return "", errors.New("game session cookie header is too large")
}
var token string
seen := false
for _, value := range strings.Split(cookie, ";") {
name, candidate, found := strings.Cut(strings.TrimSpace(value), "=")
if !found || name != "s" {
continue
}
if seen {
return "", errors.New("duplicate game session cookie")
}
seen = true
token = candidate
}
if !seen {
return "", errSessionRequired
}
if len(token) != 50 || token[48:] != "|1" {
return "", errors.New("invalid game session cookie")
}
for _, char := range token[:48] {
if !(char >= '0' && char <= '9' || char >= 'a' && char <= 'f') {
return "", errors.New("invalid game session cookie")
}
}
return token, nil
}
// KeyForTest returns the most recently created session key when called without
// a token. Supplying a raw cookie value selects that client's isolated key.
func (s *Server) KeyForTest(token ...string) []byte {
s.mu.Lock()
defer s.mu.Unlock()
selected := s.latestSessionToken
if len(token) != 0 {
selected = sessionTokenKey(token[0])
} else if !s.latestSessionSet {
return nil
}
game := s.sessions[selected]
if game == nil {
return nil
}
return append([]byte(nil), game.key...)
}
func (s *Server) ProgressForTest() *progress.Store { return s.progress }