diff --git a/.gitignore b/.gitignore index ab7b4e9..4c3e12a 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,7 @@ /server /verstak-sync-server *.exe +/build/ # Data directory server-data/ diff --git a/README.md b/README.md index 53028c8..c16f8a8 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ This server provides synchronization between devices running Verstak2. It handle - Device registration and authentication - Vault-scoped, ordered operation-log relay with server sequence numbers -- Optional blob endpoints, not used by the current bounded Desktop file sync +- Scoped content-addressed Blob transport for binary and large file content - User management with email confirmation ## Quick Start @@ -18,10 +18,12 @@ This server provides synchronization between devices running Verstak2. It handle ./scripts/build.sh # Run -./build/bin/verstak-sync-server --port 47732 --data ./server-data +./build/bin/verstak-sync-server --data ./server-data # First run with admin user -./build/bin/verstak-sync-server --admin-user admin --admin-pass secret +printf '%s\n' 'choose-a-long-password' > /tmp/verstak-admin-password +chmod 600 /tmp/verstak-admin-password +./build/bin/verstak-sync-server --admin-user admin --admin-pass-file /tmp/verstak-admin-password ``` ## Release packages @@ -51,16 +53,43 @@ annotated tag when necessary, then creates or updates the GitHub Release. | Flag | Default | Description | |------|---------|-------------| -| `--port` | 47732 | HTTP port | +| `--listen` | `127.0.0.1:47732` | HTTP address; an administrator must explicitly expose another interface | +| `--port` | — | Deprecated compatibility shortcut; always binds loopback | | `--data` | ./server-data | Data directory | | `--admin-user` | | Create admin user (first run) | -| `--admin-pass` | | Admin password (first run) | +| `--admin-pass-file` | | Read the initial admin password from a protected file | +| `--admin-pass-stdin` | | Read the initial admin password from stdin | + +The server has explicit header/read/write/idle timeouts and a 16 KiB header +limit. It handles SIGINT/SIGTERM with a 20-second graceful shutdown and closes +SQLite afterwards. Release builds publish `version` and `build_commit` through +the health response; neither logs nor health contain credentials. + +`config.yml` can set `listen`, `public_url`, `trusted_proxies`, and limits: + +```yaml +listen: 127.0.0.1:47732 +public_url: https://sync.example.test +trusted_proxies: [127.0.0.1, ::1] +limits: + max_json_body: 2097152 + max_push_operations: 100 + max_payload_json: 262144 + max_pull_page: 100 + max_blob_bytes: 268435456 + max_vault_blob_bytes: 4294967296 + max_user_blob_bytes: 8589934592 +retention: + idempotency_hours: 24 + audit_days: 90 + temp_upload_hours: 24 +``` Production installs use: - binary: `/opt/verstak-sync-server/verstak-sync-server`; - data directory: `/var/lib/verstak-sync-server`; -- port environment file: `/etc/verstak-server/env`; +- listen-address environment file: `/etc/verstak-server/env`; - service: `verstak-server`. Install from a built binary: @@ -69,9 +98,9 @@ Install from a built binary: ./scripts/build.sh sudo ./scripts/install.sh \ --bin ./build/bin/verstak-sync-server \ - --port 47732 \ + --listen 127.0.0.1:47732 \ --admin-user admin \ - --admin-pass 'change-this-password' + --admin-pass-file /root/verstak-admin-password ``` The install script creates a locked-down system user, initializes the data @@ -95,7 +124,7 @@ curl http://127.0.0.1:47732/api/v1/health Change the listen port: ```bash -echo 'VERSTAK_PORT=47733' | sudo tee /etc/verstak-server/env +echo 'VERSTAK_LISTEN=127.0.0.1:47733' | sudo tee /etc/verstak-server/env sudo systemctl restart verstak-server ``` @@ -139,7 +168,7 @@ sudo tar --xattrs --acls -xzf verstak-sync-backup-YYYYMMDD-HHMMSS.tar.gz -C /var sudo chown -R verstak:verstak /var/lib/verstak-sync-server sudo chmod 750 /var/lib/verstak-sync-server sudo systemctl start verstak-server -curl http://127.0.0.1:${VERSTAK_PORT:-47732}/api/v1/health +curl http://127.0.0.1:47732/api/v1/health ``` After restore, connected desktop clients keep their existing device tokens. @@ -176,7 +205,7 @@ Desktop sync client: User API: - `POST /api/v1/auth/register` - Register a user -- `GET /api/v1/auth/confirm?token=...` - Confirm email +- `GET /api/v1/auth/confirm?token=...` - Display a confirmation form; `POST` performs confirmation - `POST /api/v1/auth/login` - User login - `POST /api/v1/auth/forgot` - Request password reset - `POST /api/v1/auth/reset` - Reset password @@ -200,11 +229,67 @@ does not merge files, resolve conflicts, or create replacement names. The desktop pairing payload may supply an existing `vault_id` to add a new empty local vault to that remote scope. The server treats that value only as a scope selector: reconciliation, conflict detection, snapshots, and durable -workspace identity remain Desktop-core responsibilities. Current Desktop core -uses bounded inline file payloads (text or base64 up to 8 MB) and workspace -operations (`create`, `rename`, `trash`, `restore`). Blob transport, quotas, -pull pagination, and operation retention are a later milestone; the existing -blob endpoints must not be interpreted as enabled Desktop large-file sync. +workspace identity remain Desktop-core responsibilities. Small text can remain +inline; binary and large files are uploaded first and their operations carry a +`blob` `{sha256,size}` reference. Blob bytes are physically deduplicated but a +`user_id`/`vault_id` reference is mandatory. Knowing another scope's SHA-256 +never grants download access. Revoked devices and blocked users lose sync and +blob access immediately. + +`POST /api/v1/sync/pull` accepts `since_sequence` and optional `page_limit`. +The response has ordered `ops`, `page_last_sequence`, `server_sequence`, and +`has_more`. Clients must persist a cursor only after applying each operation. +Push is capped by the configured JSON/body/field limits and returns stable +JSON `{error,code}` errors (including `request_too_large`, `rate_limited`, and +`quota_exceeded`); desktop and UI localize codes rather than server text. + +New device tokens, sessions, and email/reset tokens are stored only as SHA-256 +hashes. The plaintext device token is returned once by pairing. Older plaintext +API keys are marked `legacy_api_key=1` during migration and are never created +again; rotate/re-pair them during normal deployment. Admin key endpoints show +only a prefix/suffix hint. Sessions are database-backed, expire after 24 hours, +rotate on login, and use HttpOnly/Lax cookies plus a SameSite/CSRF companion +cookie. Mutating browser endpoints require a matching CSRF token and do not use +GET for destructive work. + +### Reverse proxy and TLS + +TLS terminates at nginx or Caddy; the server has no built-in TLS. It ignores +`Forwarded`, `X-Forwarded-For`, and `X-Forwarded-Proto` unless the TCP peer is +listed in `trusted_proxies`. For a local nginx/Caddy proxy use +`trusted_proxies: [127.0.0.1, ::1]` and set `public_url` to the HTTPS URL. + +```nginx +location / { + proxy_pass http://127.0.0.1:47732; + proxy_set_header Host $host; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; +} +``` + +Never bind `0.0.0.0` merely to make a proxy work. If an external listener is +intentional, restrict it with a firewall and configure the real proxy CIDR. + +### Operations, retention, and privacy + +`GET /api/v1/health`, `/livez`, and `/readyz` report status, version, build, +uptime, database reachability, blob writability, schema version, and server +time without paths or secrets. The internal stats service exposes user/device, +vault, operation, database/blob size, and last-sync counters for a future +admin panel. + +Retention cleans expired sessions/email tokens, bounded idempotency records, +old audit entries, stale upload temp files, and in-memory rate buckets. It does +**not** delete sync operations or referenced blobs: without a materialized +checkpoint and verified recovery protocol, pruning would prevent a new device +from restoring a vault. That checkpoint/operation-retention design is a future +milestone. + +The server is optional: Desktop remains local-first. The relay can see metadata +and file bytes needed to serve operations/blobs; files are not end-to-end +encrypted in this milestone. Secrets, plugin settings, Todo, Journal, +Activity, and Browser Inbox are not synchronized here. New device enrollment requires a non-empty `vault_id`. The `legacy:` prefix is reserved for server-side migration of older records and cannot be selected by diff --git a/build/bin/verstak-sync-server b/build/bin/verstak-sync-server deleted file mode 100755 index 5f12b15..0000000 Binary files a/build/bin/verstak-sync-server and /dev/null differ diff --git a/cmd/server/main.go b/cmd/server/main.go index 907f1dc..d27ad3b 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -1,21 +1,36 @@ package main import ( + "context" + "errors" "flag" "fmt" "log" + "net" + "net/http" "os" + "os/signal" "path/filepath" + "strings" + "syscall" + "time" "github.com/verstak/verstak-sync-server/internal/server" ) func main() { - port := flag.Int("port", 47732, "HTTP port") dataDir := flag.String("data", "./server-data", "Data directory (db, blobs, config)") + listen := flag.String("listen", "", "HTTP listen address (default 127.0.0.1:47732)") + port := flag.Int("port", 0, "Deprecated compatibility override for the loopback port") adminUser := flag.String("admin-user", "", "Create admin user (first run)") - adminPass := flag.String("admin-pass", "", "Admin password (first run)") + adminPassFile := flag.String("admin-pass-file", "", "Read initial admin password from a 0600 file") + adminPassStdin := flag.Bool("admin-pass-stdin", false, "Read initial admin password from stdin") + showVersion := flag.Bool("version", false, "Print build version and exit") flag.Parse() + if *showVersion { + fmt.Printf("verstak-sync-server %s (%s)\n", server.Version, server.BuildCommit) + return + } absData, err := filepath.Abs(*dataDir) if err != nil { @@ -31,11 +46,37 @@ func main() { log.Fatalf("config: %v", err) } - if *adminUser != "" && *adminPass != "" { - if err := cfg.SetAdmin(*adminUser, *adminPass); err != nil { + if envListen := strings.TrimSpace(os.Getenv("VERSTAK_LISTEN")); envListen != "" { + cfg.Listen = envListen + } + if *listen != "" { + cfg.Listen = *listen + } + if *port != 0 && *listen == "" { + cfg.Listen = fmt.Sprintf("127.0.0.1:%d", *port) + } + if publicURL := strings.TrimSpace(os.Getenv("VERSTAK_PUBLIC_URL")); publicURL != "" { + cfg.PublicURL = publicURL + } + if trusted := strings.TrimSpace(os.Getenv("VERSTAK_TRUSTED_PROXIES")); trusted != "" { + cfg.TrustedProxies = strings.Split(trusted, ",") + } + if err := cfg.Normalize(); err != nil { + log.Fatalf("config: %v", err) + } + + adminPass, err := initialAdminPassword(*adminPassFile, *adminPassStdin) + if err != nil { + log.Fatalf("admin password: %v", err) + } + if (*adminUser == "") != (adminPass == "") { + log.Fatal("admin-user and one admin password source must be supplied together") + } + if *adminUser != "" { + if err := cfg.SetAdmin(*adminUser, adminPass); err != nil { log.Fatalf("set admin: %v", err) } - fmt.Printf("Admin user %q created.\n", *adminUser) + log.Printf("initial admin user %q configured", *adminUser) } dbPath := filepath.Join(absData, "server.db") @@ -43,13 +84,64 @@ func main() { if err != nil { log.Fatalf("server: %v", err) } - defer srv.Close() - srv.SetupRoutes() + if err := srv.CleanupRetention(time.Now().UTC()); err != nil { + _ = srv.Close() + log.Fatalf("retention cleanup: %v", err) + } + addr := cfg.ListenAddress() + listener, err := net.Listen("tcp", addr) + if err != nil { + _ = srv.Close() + log.Fatalf("listen %s: %v", addr, err) + } + httpServer := srv.HTTPServer(addr) + serveDone := make(chan error, 1) + go func() { serveDone <- httpServer.Serve(listener) }() + log.Printf("Verstak Sync Server %s (%s) listening on %s", server.Version, server.BuildCommit, addr) - addr := fmt.Sprintf(":%d", *port) - log.Printf("Verstak Sync Server starting on %s (data: %s)", addr, absData) - if err := srv.ListenAndServe(addr); err != nil { - log.Fatalf("serve: %v", err) + signals := make(chan os.Signal, 1) + signal.Notify(signals, syscall.SIGINT, syscall.SIGTERM) + select { + case sig := <-signals: + log.Printf("received %s; shutting down", sig) + shutdownCtx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + err = httpServer.Shutdown(shutdownCtx) + cancel() + if err != nil { + log.Printf("graceful shutdown: %v", err) + } + case err = <-serveDone: + if err != nil && !errors.Is(err, http.ErrServerClosed) { + _ = srv.Close() + log.Fatalf("serve: %v", err) + } + } + if err := srv.Close(); err != nil { + log.Printf("close database: %v", err) } } + +func initialAdminPassword(path string, useStdin bool) (string, error) { + if path != "" && useStdin { + return "", fmt.Errorf("choose either --admin-pass-file or --admin-pass-stdin") + } + if path == "" && !useStdin { + return "", nil + } + var data []byte + var err error + if path != "" { + data, err = os.ReadFile(path) + } else { + data, err = os.ReadFile("/dev/stdin") + } + if err != nil { + return "", err + } + password := strings.TrimSpace(string(data)) + if password == "" { + return "", fmt.Errorf("password source is empty") + } + return password, nil +} diff --git a/internal/server/blobs.go b/internal/server/blobs.go new file mode 100644 index 0000000..51bc604 --- /dev/null +++ b/internal/server/blobs.go @@ -0,0 +1,277 @@ +package server + +import ( + "crypto/sha256" + "database/sql" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "os" + "path/filepath" + "strings" + "time" +) + +func validSHA256(value string) bool { + if len(value) != sha256.Size*2 || strings.ToLower(value) != value { + return false + } + _, err := hex.DecodeString(value) + return err == nil +} + +func blobPath(root, hash string) string { + return filepath.Join(root, hash[:2], hash[2:4], hash) +} + +func (s *Server) handleBlobUpload(w http.ResponseWriter, r *http.Request, scope authenticatedDevice) { + // Reserve a little multipart framing headroom while the file itself is + // measured independently. The stream is never materialized in memory. + r.Body = http.MaxBytesReader(w, r.Body, s.cfg.Limits.MaxBlobBytes+(1<<20)) + reader, err := r.MultipartReader() + if err != nil { + jsonErrCode(w, http.StatusBadRequest, "invalid_multipart", "invalid multipart request") + return + } + part, err := reader.NextPart() + if errors.Is(err, io.EOF) { + jsonErrCode(w, http.StatusBadRequest, "invalid_multipart", "file field is required") + return + } + if err != nil { + jsonErrCode(w, http.StatusBadRequest, "invalid_multipart", "invalid multipart request") + return + } + defer part.Close() + if part.FormName() != "file" || part.FileName() == "" { + jsonErrCode(w, http.StatusBadRequest, "invalid_multipart", "exactly one file field is required") + return + } + tmp, err := os.CreateTemp(s.blobsDir, ".upload-*") + if err != nil { + jsonInternalError(w, err) + return + } + tmpName := tmp.Name() + cleanupTmp := true + defer func() { + if cleanupTmp { + _ = os.Remove(tmpName) + } + }() + if err := tmp.Chmod(0640); err != nil { + _ = tmp.Close() + jsonInternalError(w, err) + return + } + hash := sha256.New() + written, err := io.Copy(io.MultiWriter(tmp, hash), io.LimitReader(part, s.cfg.Limits.MaxBlobBytes+1)) + if err != nil { + _ = tmp.Close() + var tooLarge *http.MaxBytesError + if errors.As(err, &tooLarge) { + jsonErrCode(w, http.StatusRequestEntityTooLarge, "blob_too_large", "blob is too large") + } else { + jsonInternalError(w, err) + } + return + } + if written > s.cfg.Limits.MaxBlobBytes { + _ = tmp.Close() + jsonErrCode(w, http.StatusRequestEntityTooLarge, "blob_too_large", "blob is too large") + return + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + jsonInternalError(w, err) + return + } + if err := tmp.Close(); err != nil { + jsonInternalError(w, err) + return + } + // Do not silently accept a second part: this endpoint has one unambiguous + // file contract and no form fields that could hide extra request data. + if extra, err := reader.NextPart(); err != io.EOF { + if extra != nil { + _ = extra.Close() + } + jsonErrCode(w, http.StatusBadRequest, "invalid_multipart", "exactly one file field is required") + return + } + sha := hex.EncodeToString(hash.Sum(nil)) + if err := s.attachBlob(scope.UserID, scope.VaultID, sha, written, tmpName); err != nil { + if errors.Is(err, errBlobTooLarge) { + jsonErrCode(w, http.StatusRequestEntityTooLarge, "quota_exceeded", "blob quota exceeded") + return + } + jsonInternalError(w, err) + return + } + jsonOK(w, map[string]interface{}{"sha256": sha, "size": written}) +} + +var errBlobTooLarge = errors.New("blob quota exceeded") + +// attachBlob checks logical quotas before the rename, then atomically makes +// the content reachable and commits the scope reference. A DB error removes a +// newly-created physical file so rejected requests leave no blob behind. +func (s *Server) attachBlob(userID, vaultID, sha string, size int64, tmpName string) error { + tx, err := s.db.Begin() + if err != nil { + return err + } + defer tx.Rollback() + var existingSize int64 + err = tx.QueryRow(`SELECT size FROM server_blob_refs WHERE user_id=? AND vault_id=? AND sha256=?`, userID, vaultID, sha).Scan(&existingSize) + if err == nil { + if existingSize != size { + return fmt.Errorf("existing blob reference has a different size") + } + if _, err := tx.Exec(`UPDATE server_blob_refs SET last_accessed=? WHERE user_id=? AND vault_id=? AND sha256=?`, time.Now().UTC().Format(time.RFC3339), userID, vaultID, sha); err != nil { + return err + } + if err := tx.Commit(); err != nil { + return err + } + return nil + } + if !errors.Is(err, sql.ErrNoRows) { + return err + } + var vaultUsed, userUsed int64 + if err := tx.QueryRow(`SELECT COALESCE(SUM(size), 0) FROM server_blob_refs WHERE user_id=? AND vault_id=?`, userID, vaultID).Scan(&vaultUsed); err != nil { + return err + } + if err := tx.QueryRow(`SELECT COALESCE(SUM(size), 0) FROM server_blob_refs WHERE user_id=?`, userID).Scan(&userUsed); err != nil { + return err + } + if vaultUsed+size > s.cfg.Limits.MaxVaultBlobBytes || userUsed+size > s.cfg.Limits.MaxUserBlobBytes { + return errBlobTooLarge + } + destination := blobPath(s.blobsDir, sha) + if err := os.MkdirAll(filepath.Dir(destination), 0750); err != nil { + return err + } + createdPhysical := false + if _, err := os.Stat(destination); errors.Is(err, os.ErrNotExist) { + if err := os.Rename(tmpName, destination); err != nil { + return err + } + createdPhysical = true + } else if err != nil { + return err + } + now := time.Now().UTC().Format(time.RFC3339) + if _, err := tx.Exec(`INSERT OR IGNORE INTO server_blobs (sha256, size, created_at) VALUES (?, ?, ?)`, sha, size, now); err != nil { + if createdPhysical { + _ = os.Remove(destination) + } + return err + } + if _, err := tx.Exec(`INSERT INTO server_blob_refs (user_id, vault_id, sha256, size, created_at, last_accessed) VALUES (?, ?, ?, ?, ?, ?)`, userID, vaultID, sha, size, now, now); err != nil { + if createdPhysical { + _ = os.Remove(destination) + } + return err + } + if err := tx.Commit(); err != nil { + if createdPhysical { + _ = os.Remove(destination) + } + return err + } + return nil +} + +func (s *Server) handleBlobDownload(w http.ResponseWriter, r *http.Request, scope authenticatedDevice, sha string) { + if !validSHA256(sha) { + jsonErrCode(w, http.StatusBadRequest, "invalid_blob_hash", "invalid SHA-256") + return + } + var size int64 + err := s.db.QueryRow(`SELECT size FROM server_blob_refs WHERE user_id=? AND vault_id=? AND sha256=?`, scope.UserID, scope.VaultID, sha).Scan(&size) + if errors.Is(err, sql.ErrNoRows) { + jsonErrCode(w, http.StatusNotFound, "blob_not_found", "blob not found") + return + } + if err != nil { + jsonInternalError(w, err) + return + } + file, err := os.Open(blobPath(s.blobsDir, sha)) + if errors.Is(err, os.ErrNotExist) { + jsonErrCode(w, http.StatusNotFound, "blob_not_found", "blob not found") + return + } + if err != nil { + jsonInternalError(w, err) + return + } + defer file.Close() + info, err := file.Stat() + if err != nil { + jsonInternalError(w, err) + return + } + if !info.Mode().IsRegular() || info.Size() != size { + jsonErrCode(w, http.StatusInternalServerError, "blob_unavailable", "blob is unavailable") + return + } + if _, err := s.db.Exec(`UPDATE server_blob_refs SET last_accessed=? WHERE user_id=? AND vault_id=? AND sha256=?`, time.Now().UTC().Format(time.RFC3339), scope.UserID, scope.VaultID, sha); err != nil { + jsonInternalError(w, err) + return + } + w.Header().Set("Content-Type", "application/octet-stream") + w.Header().Set("Content-Length", fmt.Sprintf("%d", size)) + w.Header().Set("Content-Disposition", "attachment; filename=\""+sha+"\"") + // ServeContent streams from the file and provides safe Range support. + http.ServeContent(w, r, sha, info.ModTime(), file) +} + +// validateScopedBlobReferences ensures an accepted operation can later be +// restored by every authorized device in this scope. Legacy base64 binary +// payloads are deliberately rejected: binary content belongs in Blob storage, +// never in the operation log. +func validateScopedBlobReferences(tx *sql.Tx, scope authenticatedDevice, ops []syncPushOperation) (code, message string, err error) { + for _, op := range ops { + if op.EntityType != "file" || op.PayloadJSON == "" { + continue + } + var payload struct { + DataBase64 *string `json:"dataBase64"` + Blob *struct { + SHA256 string `json:"sha256"` + Size int64 `json:"size"` + } `json:"blob"` + } + if err := json.Unmarshal([]byte(op.PayloadJSON), &payload); err != nil { + return "invalid_payload", "operation payload_json must be valid JSON", nil + } + if payload.DataBase64 != nil { + return "inline_binary_forbidden", "binary file payload must reference a blob", nil + } + if payload.Blob == nil { + continue + } + if !validSHA256(payload.Blob.SHA256) || payload.Blob.Size < 0 { + return "invalid_blob_reference", "invalid blob reference", nil + } + var storedSize int64 + err := tx.QueryRow(`SELECT size FROM server_blob_refs WHERE user_id=? AND vault_id=? AND sha256=?`, + scope.UserID, scope.VaultID, payload.Blob.SHA256).Scan(&storedSize) + if errors.Is(err, sql.ErrNoRows) { + return "blob_not_owned", "blob must be uploaded before its operation", nil + } + if err != nil { + return "", "", err + } + if storedSize != payload.Blob.Size { + return "invalid_blob_reference", "blob size does not match uploaded content", nil + } + } + return "", "", nil +} diff --git a/internal/server/config.go b/internal/server/config.go index 4e7565e..0140988 100644 --- a/internal/server/config.go +++ b/internal/server/config.go @@ -2,8 +2,10 @@ package server import ( "fmt" + "net/netip" "os" "path/filepath" + "strings" "sync" "golang.org/x/crypto/bcrypt" @@ -16,28 +18,178 @@ type AdminUser struct { } type Config struct { - Port int `yaml:"port"` - Admin []AdminUser `yaml:"admin"` - mu sync.Mutex - path string + // Port remains readable for older config.yml files. New deployments use + // Listen so an administrator must deliberately expose a non-loopback + // listener. + Port int `yaml:"port,omitempty"` + Listen string `yaml:"listen,omitempty"` + TrustedProxies []string `yaml:"trusted_proxies,omitempty"` + PublicURL string `yaml:"public_url,omitempty"` + DevelopmentTokenLogging bool `yaml:"development_token_logging,omitempty"` + Limits Limits `yaml:"limits,omitempty"` + Retention Retention `yaml:"retention,omitempty"` + Admin []AdminUser `yaml:"admin"` + mu sync.Mutex + path string + trustedProxyPrefixes []netip.Prefix +} + +// Retention controls data that has no role in reconstructing a vault. Sync +// operations and referenced blobs are deliberately absent: pruning either +// requires a checkpoint protocol or risks making a new device unrecoverable. +type Retention struct { + IdempotencyHours int `yaml:"idempotency_hours"` + AuditDays int `yaml:"audit_days"` + TempUploadHours int `yaml:"temp_upload_hours"` +} + +// Limits protect the process and operation log independently of a client +// supplied value. They are deliberately conservative defaults for a +// self-hosted service and can be adjusted in config.yml. +type Limits struct { + MaxJSONBody int64 `yaml:"max_json_body"` + MaxPushOperations int `yaml:"max_push_operations"` + MaxPayloadJSON int `yaml:"max_payload_json"` + MaxPullPage int `yaml:"max_pull_page"` + MaxBlobBytes int64 `yaml:"max_blob_bytes"` + MaxVaultBlobBytes int64 `yaml:"max_vault_blob_bytes"` + MaxUserBlobBytes int64 `yaml:"max_user_blob_bytes"` +} + +func defaultLimits() Limits { + return Limits{ + MaxJSONBody: 2 << 20, + MaxPushOperations: 100, + MaxPayloadJSON: 256 << 10, + MaxPullPage: 100, + MaxBlobBytes: 256 << 20, + MaxVaultBlobBytes: 4 << 30, + MaxUserBlobBytes: 8 << 30, + } +} + +// DefaultConfig is the safe baseline used by a new server and tests. +func DefaultConfig() *Config { + return &Config{ + Port: 47732, + Listen: "127.0.0.1:47732", + Limits: defaultLimits(), + Retention: Retention{IdempotencyHours: 24, AuditDays: 90, TempUploadHours: 24}, + } } func LoadConfig(dataDir string) (*Config, error) { path := filepath.Join(dataDir, "config.yml") - cfg := &Config{ - Port: 47732, - Admin: nil, - path: path, - } + cfg := DefaultConfig() + cfg.path = path data, err := os.ReadFile(path) if err == nil { if err := yaml.Unmarshal(data, cfg); err != nil { return nil, fmt.Errorf("parse config: %w", err) } } + if err != nil && !os.IsNotExist(err) { + return nil, fmt.Errorf("read config: %w", err) + } + if err := cfg.normalize(); err != nil { + return nil, err + } return cfg, nil } +// ListenAddress returns a loopback address unless an administrator explicitly +// configured another address. A legacy port-only config is also loopback. +func (c *Config) ListenAddress() string { + if c == nil { + return "127.0.0.1:47732" + } + if listen := strings.TrimSpace(c.Listen); listen != "" { + return listen + } + port := c.Port + if port == 0 { + port = 47732 + } + return fmt.Sprintf("127.0.0.1:%d", port) +} + +func (c *Config) normalize() error { + if c.Port == 0 { + c.Port = 47732 + } + if strings.TrimSpace(c.Listen) == "" { + c.Listen = fmt.Sprintf("127.0.0.1:%d", c.Port) + } + defaults := defaultLimits() + if c.Limits.MaxJSONBody <= 0 { + c.Limits.MaxJSONBody = defaults.MaxJSONBody + } + if c.Limits.MaxPushOperations <= 0 { + c.Limits.MaxPushOperations = defaults.MaxPushOperations + } + if c.Limits.MaxPayloadJSON <= 0 { + c.Limits.MaxPayloadJSON = defaults.MaxPayloadJSON + } + if c.Limits.MaxPullPage <= 0 { + c.Limits.MaxPullPage = defaults.MaxPullPage + } + if c.Limits.MaxBlobBytes <= 0 { + c.Limits.MaxBlobBytes = defaults.MaxBlobBytes + } + if c.Limits.MaxVaultBlobBytes <= 0 { + c.Limits.MaxVaultBlobBytes = defaults.MaxVaultBlobBytes + } + if c.Limits.MaxUserBlobBytes <= 0 { + c.Limits.MaxUserBlobBytes = defaults.MaxUserBlobBytes + } + if c.Retention.IdempotencyHours <= 0 { + c.Retention.IdempotencyHours = 24 + } + if c.Retention.AuditDays <= 0 { + c.Retention.AuditDays = 90 + } + if c.Retention.TempUploadHours <= 0 { + c.Retention.TempUploadHours = 24 + } + prefixes := make([]netip.Prefix, 0, len(c.TrustedProxies)) + for _, raw := range c.TrustedProxies { + raw = strings.TrimSpace(raw) + if raw == "" { + continue + } + prefix, err := netip.ParsePrefix(raw) + if err != nil { + addr, addrErr := netip.ParseAddr(raw) + if addrErr != nil { + return fmt.Errorf("trusted proxy %q is neither an IP address nor CIDR: %w", raw, err) + } + bits := 32 + if addr.Is6() { + bits = 128 + } + prefix = netip.PrefixFrom(addr, bits) + } + prefixes = append(prefixes, prefix.Masked()) + } + c.trustedProxyPrefixes = prefixes + return nil +} + +// Normalize applies safe defaults and validates proxy configuration after +// command-line and environment overrides. +func (c *Config) Normalize() error { + return c.normalize() +} + +func (c *Config) isTrustedProxy(addr netip.Addr) bool { + for _, prefix := range c.trustedProxyPrefixes { + if prefix.Contains(addr) { + return true + } + } + return false +} + func (c *Config) Save() error { c.mu.Lock() defer c.mu.Unlock() @@ -80,5 +232,29 @@ func (c *Config) saveLocked() error { if err != nil { return err } - return os.WriteFile(c.path, data, 0640) + if err := os.MkdirAll(filepath.Dir(c.path), 0750); err != nil { + return err + } + tmp, err := os.CreateTemp(filepath.Dir(c.path), ".config-*") + if err != nil { + return err + } + tmpName := tmp.Name() + defer os.Remove(tmpName) + if err := tmp.Chmod(0640); err != nil { + _ = tmp.Close() + return err + } + if _, err := tmp.Write(data); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + return err + } + if err := tmp.Close(); err != nil { + return err + } + return os.Rename(tmpName, c.path) } diff --git a/internal/server/csrf.go b/internal/server/csrf.go new file mode 100644 index 0000000..eae22bc --- /dev/null +++ b/internal/server/csrf.go @@ -0,0 +1,76 @@ +package server + +import ( + "crypto/subtle" + "net/http" + "net/url" + "strings" +) + +func (s *Server) requireAdminMutation(w http.ResponseWriter, r *http.Request) bool { + session, ok := s.requireSession(w, r, sessionScopeAdmin) + if !ok { + return false + } + return s.verifyCSRF(w, r, session) +} + +func (s *Server) requireUserMutation(w http.ResponseWriter, r *http.Request) bool { + session, ok := s.requireSession(w, r, sessionScopeUser) + if !ok { + return false + } + return s.verifyCSRF(w, r, session) +} + +func (s *Server) verifyCSRF(w http.ResponseWriter, r *http.Request, session webSession) bool { + cookie, err := r.Cookie("csrf_token") + if err != nil || cookie.Value == "" { + jsonErrCode(w, http.StatusForbidden, "csrf_invalid", "CSRF token is required") + return false + } + candidate := r.Header.Get("X-CSRF-Token") + if candidate == "" { + // Form parsing is only needed for regular HTML form posts. JSON callers + // must send the header, avoiding any interference with their decoder. + if strings.HasPrefix(r.Header.Get("Content-Type"), "application/x-www-form-urlencoded") || strings.HasPrefix(r.Header.Get("Content-Type"), "multipart/form-data") { + if err := r.ParseForm(); err != nil { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", "invalid form") + return false + } + candidate = r.FormValue("csrf_token") + } + } + if candidate == "" || subtle.ConstantTimeCompare([]byte(cookie.Value), []byte(candidate)) != 1 || subtle.ConstantTimeCompare([]byte(sha256Hex(candidate)), []byte(session.CSRFHash)) != 1 { + jsonErrCode(w, http.StatusForbidden, "csrf_invalid", "CSRF token is invalid") + return false + } + if !s.sameOrigin(r) { + jsonErrCode(w, http.StatusForbidden, "csrf_invalid", "request origin is not allowed") + return false + } + return true +} + +func (s *Server) sameOrigin(r *http.Request) bool { + origin := strings.TrimSpace(r.Header.Get("Origin")) + if origin == "" { + // Older same-origin form submissions can omit Origin. The CSRF token is + // still mandatory, so accepting this remains protected. + return true + } + originURL, err := url.Parse(origin) + if err != nil || originURL.Host == "" { + return false + } + expected := strings.TrimSpace(s.cfg.PublicURL) + if expected != "" { + expectedURL, err := url.Parse(expected) + return err == nil && strings.EqualFold(originURL.Scheme, expectedURL.Scheme) && strings.EqualFold(originURL.Host, expectedURL.Host) + } + scheme := "http" + if s.requestIsHTTPS(r) { + scheme = "https" + } + return strings.EqualFold(originURL.Scheme, scheme) && strings.EqualFold(originURL.Host, r.Host) +} diff --git a/internal/server/devices.go b/internal/server/devices.go new file mode 100644 index 0000000..b6de1b8 --- /dev/null +++ b/internal/server/devices.go @@ -0,0 +1,41 @@ +package server + +import ( + "fmt" +) + +func validatePairRequest(login, deviceName, clientVersion, vaultID string) error { + for _, field := range []struct { + name string + value string + max int + }{ + {"login", login, maxLoginLength}, + {"device name", deviceName, maxDeviceNameLength}, + {"client version", clientVersion, maxClientVersionLength}, + {"vault id", vaultID, maxVaultIDLength}, + } { + if err := validateStringLength(field.name, field.value, field.max); err != nil { + return err + } + } + if vaultID == "" { + return fmt.Errorf("vault_id required") + } + return nil +} + +func (s *Server) revokeDevice(deviceID, when string) error { + tx, err := s.db.Begin() + if err != nil { + return err + } + defer tx.Rollback() + if _, err := tx.Exec("UPDATE server_devices SET revoked_at=? WHERE id=?", when, deviceID); err != nil { + return err + } + if _, err := tx.Exec("DELETE FROM server_sessions WHERE subject_id=? AND scope='device'", deviceID); err != nil { + return err + } + return tx.Commit() +} diff --git a/internal/server/email_tokens.go b/internal/server/email_tokens.go new file mode 100644 index 0000000..1e19fcb --- /dev/null +++ b/internal/server/email_tokens.go @@ -0,0 +1,26 @@ +package server + +import ( + "database/sql" + "time" +) + +type sqlExecutor interface { + Exec(query string, args ...interface{}) (sql.Result, error) +} + +func issueEmailToken(exec sqlExecutor, userID, purpose string, lifetime time.Duration) (string, error) { + token, err := randomSecret(24) + if err != nil { + return "", err + } + now := time.Now().UTC() + _, err = exec.Exec(`INSERT INTO server_email_tokens (token_hash, user_id, purpose, expires_at, created_at) + VALUES (?, ?, ?, ?, ?)`, sha256Hex(token), userID, purpose, now.Add(lifetime).Format(time.RFC3339), now.Format(time.RFC3339)) + if err != nil { + return "", err + } + return token, nil +} + +func emailTokenHash(token string) string { return sha256Hex(token) } diff --git a/internal/server/handlers_admin.go b/internal/server/handlers_admin.go index 517eb6e..92873f2 100644 --- a/internal/server/handlers_admin.go +++ b/internal/server/handlers_admin.go @@ -35,15 +35,19 @@ button{background:#4ecca3;color:#1a1a2e;border:none;padding:0.5rem 1rem;border-r } user := r.FormValue("username") pass := r.FormValue("password") + if !s.allowRate(w, r, "login", user) { + return + } if !s.cfg.CheckAdmin(user, pass) { http.Error(w, "401 Unauthorized", 401) return } - tok := s.tokens.Create() - http.SetCookie(w, &http.Cookie{ - Name: "admin_session", Value: tok, Path: "/admin", - HttpOnly: true, SameSite: http.SameSiteLaxMode, MaxAge: 86400, - }) + tok, csrf, err := s.createSession(sessionScopeAdmin, user) + if err != nil { + jsonInternalError(w, err) + return + } + s.setSessionCookies(w, r, sessionScopeAdmin, tok, csrf) http.Redirect(w, r, "/admin/dashboard", http.StatusFound) default: http.Error(w, "method not allowed", 405) @@ -51,6 +55,10 @@ button{background:#4ecca3;color:#1a1a2e;border:none;padding:0.5rem 1rem;border-r } func (s *Server) handleAdminDashboard(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + methodNotAllowed(w, http.MethodGet) + return + } if !s.requireAdminCookie(w, r) { return } @@ -77,6 +85,10 @@ a{color:#4ecca3}
} func (s *Server) handleAdminUsers(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + methodNotAllowed(w, http.MethodGet) + return + } if !s.requireAdminCookie(w, r) { return } @@ -123,6 +135,10 @@ td{padding:0.5rem;border-bottom:1px solid #0f3460}a{color:#4ecca3}= 0 { - ip = ip[:idx] - } - if !s.pairLimit.allow(ip) { - s.auditLog("rate_limit_exceeded", "", "", ip, "pair rate limit exceeded") - jsonErr(w, 429, "too many attempts") - return - } + ip := s.clientIP(r) var req struct { Login string `json:"login"` Password string `json:"password"` @@ -52,14 +56,16 @@ func (s *Server) handleClientPair(w http.ResponseWriter, r *http.Request) { ClientVersion string `json:"client_version"` VaultID string `json:"vault_id"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "bad json") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Login == "" || req.Password == "" { jsonErr(w, 400, "login and password required") return } + if !s.allowRate(w, r, "pair", req.Login) { + return + } req.VaultID = strings.TrimSpace(req.VaultID) if req.VaultID == "" { jsonErr(w, 400, "vault_id required") @@ -72,12 +78,16 @@ func (s *Server) handleClientPair(w http.ResponseWriter, r *http.Request) { if req.DeviceName == "" { req.DeviceName = "unknown" } + if err := validatePairRequest(req.Login, req.DeviceName, req.ClientVersion, req.VaultID); err != nil { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", err.Error()) + return + } var userID, hash string var confirmed, blocked int err := s.db.QueryRow("SELECT id, password_hash, confirmed, blocked FROM server_users WHERE username=? OR email=?", req.Login, strings.ToLower(req.Login)).Scan(&userID, &hash, &confirmed, &blocked) if err != nil { - s.auditLog("device_auth_failed", "", "", ip, "pair: user not found: "+req.Login) + s.auditLog("device_auth_failed", "", "", ip, "pair: user not found") jsonErr(w, 401, "invalid credentials") return } @@ -102,20 +112,33 @@ func (s *Server) handleClientPair(w http.ResponseWriter, r *http.Request) { token, prefix, suffix := genDeviceToken() tokenHash := sha256Hex(token) now := time.Now().UTC().Format(time.RFC3339) - apiKey := make([]byte, 20) - rand.Read(apiKey) - _, err = s.db.Exec(`INSERT INTO server_devices - (id, name, api_key, token_hash, token_prefix, token_suffix, user_id, vault_id, client_version, last_ip, last_seen, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - deviceID, req.DeviceName, hex.EncodeToString(apiKey), tokenHash, prefix, suffix, + tx, err := s.db.Begin() + if err != nil { + jsonInternalError(w, err) + return + } + defer tx.Rollback() + _, err = tx.Exec(`INSERT INTO server_devices + (id, name, api_key, token_hash, token_prefix, token_suffix, legacy_api_key, user_id, vault_id, client_version, last_ip, last_seen, created_at) + VALUES (?, ?, ?, ?, ?, ?, 0, ?, ?, ?, ?, ?, ?)`, + deviceID, req.DeviceName, "disabled:"+deviceID, tokenHash, prefix, suffix, userID, req.VaultID, req.ClientVersion, ip, now, now) if err != nil { jsonInternalError(w, err) return } - s.db.Exec("INSERT OR IGNORE INTO server_user_devices (user_id, device_id) VALUES (?, ?)", userID, deviceID) - s.db.Exec("UPDATE server_users SET last_seen=? WHERE id=?", now, userID) - s.pairLimit.reset(ip) + if _, err := tx.Exec("INSERT OR IGNORE INTO server_user_devices (user_id, device_id) VALUES (?, ?)", userID, deviceID); err != nil { + jsonInternalError(w, err) + return + } + if _, err := tx.Exec("UPDATE server_users SET last_seen=? WHERE id=?", now, userID); err != nil { + jsonInternalError(w, err) + return + } + if err := tx.Commit(); err != nil { + jsonInternalError(w, err) + return + } s.auditLog("device_paired", userID, deviceID, ip, "device paired: "+req.DeviceName) jsonOK(w, map[string]interface{}{ "user_id": userID, @@ -135,14 +158,16 @@ func (s *Server) handleAuthTest(w http.ResponseWriter, r *http.Request) { Username string `json:"username"` Password string `json:"password"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "bad json") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Username == "" || req.Password == "" { jsonErr(w, 400, "username and password required") return } + if !s.allowRate(w, r, "auth-test", req.Username) { + return + } var hash string var confirmed, blocked int err := s.db.QueryRow("SELECT password_hash, confirmed, blocked FROM server_users WHERE username=? OR email=?", @@ -171,21 +196,16 @@ func (s *Server) handleClientRevoke(w http.ResponseWriter, r *http.Request) { jsonErr(w, 405, "POST required") return } - tok := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") - if tok == "" { - jsonErr(w, 401, "token required") - return - } - hash := sha256Hex(tok) - var deviceID, userID string - err := s.db.QueryRow("SELECT id, user_id FROM server_devices WHERE token_hash=?", hash).Scan(&deviceID, &userID) - if err != nil { - jsonErr(w, 401, "invalid token") + device, ok := s.authenticateDevice(w, r) + if !ok { return } now := time.Now().UTC().Format(time.RFC3339) - s.db.Exec("UPDATE server_devices SET revoked_at=? WHERE id=?", now, deviceID) - s.auditLog("device_revoked", userID, deviceID, r.RemoteAddr, "device revoked by user") + if err := s.revokeDevice(device.DeviceID, now); err != nil { + jsonInternalError(w, err) + return + } + s.auditLog("device_revoked", device.UserID, device.DeviceID, s.clientIP(r), "device revoked by user") jsonOK(w, map[string]string{"status": "revoked"}) } @@ -194,32 +214,27 @@ func (s *Server) handleClientRevokeDevice(w http.ResponseWriter, r *http.Request jsonErr(w, 405, "POST required") return } - tok := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") - if tok == "" { - jsonErr(w, 401, "token required") - return - } - hash := sha256Hex(tok) - var curUserID string - err := s.db.QueryRow("SELECT user_id FROM server_devices WHERE token_hash=?", hash).Scan(&curUserID) - if err != nil || curUserID == "" { - jsonErr(w, 401, "invalid token") + device, ok := s.authenticateDevice(w, r) + if !ok || device.UserID == "" { return } + curUserID := device.UserID var req struct { DeviceID string `json:"device_id"` Password string `json:"password"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.DeviceID == "" || req.Password == "" { jsonErr(w, 400, "device_id and password required") return } + if !s.allowRate(w, r, "auth-test", curUserID) { + return + } var pwHash string - err = s.db.QueryRow("SELECT password_hash FROM server_users WHERE id=?", curUserID).Scan(&pwHash) + err := s.db.QueryRow("SELECT password_hash FROM server_users WHERE id=?", curUserID).Scan(&pwHash) if err != nil { jsonErr(w, 403, "access denied") return @@ -239,21 +254,26 @@ func (s *Server) handleClientRevokeDevice(w http.ResponseWriter, r *http.Request return } now := time.Now().UTC().Format(time.RFC3339) - s.db.Exec("UPDATE server_devices SET revoked_at=? WHERE id=?", now, req.DeviceID) - s.auditLog("device_revoked", curUserID, req.DeviceID, r.RemoteAddr, "device revoked via API") + if err := s.revokeDevice(req.DeviceID, now); err != nil { + jsonInternalError(w, err) + return + } + s.auditLog("device_revoked", curUserID, req.DeviceID, s.clientIP(r), "device revoked via API") jsonOK(w, map[string]string{"status": "revoked"}) } func (s *Server) handleClientMe(w http.ResponseWriter, r *http.Request) { - tok := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") - if tok == "" { - jsonErr(w, 401, "token required") + if r.Method != http.MethodGet { + methodNotAllowed(w, http.MethodGet) + return + } + device, ok := s.authenticateDevice(w, r) + if !ok { return } - hash := sha256Hex(tok) var deviceID, userID, name, clientVer, lastSeen, revokedAt, createdAt string err := s.db.QueryRow(`SELECT d.id, d.user_id, d.name, COALESCE(d.client_version,''), COALESCE(d.last_seen,''), COALESCE(d.revoked_at,''), d.created_at - FROM server_devices d WHERE d.token_hash=?`, hash). + FROM server_devices d WHERE d.id=?`, device.DeviceID). Scan(&deviceID, &userID, &name, &clientVer, &lastSeen, &revokedAt, &createdAt) if err != nil { jsonErr(w, 401, "invalid token") @@ -284,8 +304,7 @@ func (s *Server) handleDeviceRegister(w http.ResponseWriter, r *http.Request) { Password string `json:"password"` VaultID string `json:"vault_id"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Name == "" { @@ -296,6 +315,13 @@ func (s *Server) handleDeviceRegister(w http.ResponseWriter, r *http.Request) { jsonErr(w, 401, "username and password required") return } + if !s.allowRate(w, r, "device-register", req.Username) { + return + } + if err := validatePairRequest(req.Username, req.Name, "", req.VaultID); err != nil { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", err.Error()) + return + } req.VaultID = strings.TrimSpace(req.VaultID) if req.VaultID == "" { jsonErr(w, 400, "vault_id required") @@ -325,61 +351,86 @@ func (s *Server) handleDeviceRegister(w http.ResponseWriter, r *http.Request) { jsonErr(w, 401, "invalid credentials") return } - b := make([]byte, 20) + b := make([]byte, 12) rand.Read(b) - apiKey := hex.EncodeToString(b) - deviceID := apiKey[:12] + deviceID := "dev_" + hex.EncodeToString(b) + token, prefix, suffix := genDeviceToken() now := time.Now().UTC().Format(time.RFC3339) - _, err = s.db.Exec( - "INSERT INTO server_devices (id, name, api_key, user_id, vault_id, last_seen, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)", - deviceID, req.Name, apiKey, userID, req.VaultID, now, now, + tx, err := s.db.Begin() + if err != nil { + jsonInternalError(w, err) + return + } + defer tx.Rollback() + _, err = tx.Exec( + "INSERT INTO server_devices (id, name, api_key, token_hash, token_prefix, token_suffix, legacy_api_key, user_id, vault_id, last_seen, created_at) VALUES (?, ?, ?, ?, ?, ?, 0, ?, ?, ?, ?)", + deviceID, req.Name, "disabled:"+deviceID, sha256Hex(token), prefix, suffix, userID, req.VaultID, now, now, ) if err != nil { jsonInternalError(w, err) return } - s.db.Exec("INSERT OR IGNORE INTO server_user_devices (user_id, device_id) VALUES (?, ?)", userID, deviceID) + if _, err := tx.Exec("INSERT OR IGNORE INTO server_user_devices (user_id, device_id) VALUES (?, ?)", userID, deviceID); err != nil { + jsonInternalError(w, err) + return + } + if err := tx.Commit(); err != nil { + jsonInternalError(w, err) + return + } jsonOK(w, map[string]interface{}{ - "device_id": deviceID, - "api_key": apiKey, + "device_id": deviceID, + "device_token": token, }) } func (s *Server) handleSyncPush(w http.ResponseWriter, r *http.Request) { + if r.Method != "POST" { + methodNotAllowed(w, "POST") + return + } scope, ok := s.requireSyncScope(w, r) if !ok { return } - if r.Method != "POST" { - jsonErr(w, 405, "POST required") + var req syncPushRequest + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } - var req struct { - DeviceID string `json:"device_id"` - IdempotencyKey string `json:"idempotency_key"` - Ops []struct { - OpID string `json:"op_id"` - EntityType string `json:"entity_type"` - EntityID string `json:"entity_id"` - OpType string `json:"op_type"` - PayloadJSON string `json:"payload_json"` - ClientSequence int `json:"client_sequence"` - LastSeenServerSeq int `json:"last_seen_server_seq"` - CreatedAt string `json:"created_at"` - } `json:"ops"` + if code, message := s.validateSyncPush(req); code != "" { + status := http.StatusBadRequest + if code == "too_many_operations" || code == "payload_too_large" { + status = http.StatusRequestEntityTooLarge + } + jsonErrCode(w, status, code, message) + return } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + + tx, err := s.db.Begin() + if err != nil { + jsonInternalError(w, err) + return + } + defer tx.Rollback() + if code, message, err := validateScopedBlobReferences(tx, scope, req.Ops); err != nil { + jsonInternalError(w, err) + return + } else if code != "" { + jsonErrCode(w, http.StatusBadRequest, code, message) return } if req.IdempotencyKey != "" { var cachedJSON string - err := s.db.QueryRow(`SELECT response_json FROM server_idempotency_keys + err := tx.QueryRow(`SELECT response_json FROM server_idempotency_keys WHERE user_id=? AND vault_id=? AND idempotency_key=?`, scope.UserID, scope.VaultID, req.IdempotencyKey).Scan(&cachedJSON) if err == nil { w.Header().Set("Content-Type", "application/json") - w.Write([]byte(cachedJSON)) + _, _ = w.Write([]byte(cachedJSON)) + return + } + if err != nil && !errors.Is(err, sql.ErrNoRows) { + jsonInternalError(w, err) return } } @@ -387,55 +438,82 @@ func (s *Server) handleSyncPush(w http.ResponseWriter, r *http.Request) { var accepted []string var conflicts []map[string]interface{} for _, op := range req.Ops { - if op.OpID == "" || op.EntityType == "" || op.EntityID == "" || op.OpType == "" { - continue - } if op.LastSeenServerSeq > 0 { - conflictRows, err := s.db.Query(` + conflictRows, err := tx.Query(` SELECT op_id, device_id, op_type, server_sequence FROM server_ops WHERE user_id=? AND vault_id=? AND entity_type=? AND entity_id=? AND device_id!=? AND server_sequence > ? AND op_type != 'delete' ORDER BY server_sequence`, scope.UserID, scope.VaultID, op.EntityType, op.EntityID, scope.DeviceID, op.LastSeenServerSeq) - if err == nil { - for conflictRows.Next() { - var cOpID, cDevID, cOpType string - var cSeq int - conflictRows.Scan(&cOpID, &cDevID, &cOpType, &cSeq) - conflicts = append(conflicts, map[string]interface{}{ - "op_id": cOpID, - "device_id": cDevID, - "op_type": cOpType, - "server_sequence": cSeq, - "entity_type": op.EntityType, - "entity_id": op.EntityID, - }) + if err != nil { + jsonInternalError(w, err) + return + } + for conflictRows.Next() { + var cOpID, cDevID, cOpType string + var cSeq int + if err := conflictRows.Scan(&cOpID, &cDevID, &cOpType, &cSeq); err != nil { + _ = conflictRows.Close() + jsonInternalError(w, err) + return } - conflictRows.Close() + conflicts = append(conflicts, map[string]interface{}{ + "op_id": cOpID, + "device_id": cDevID, + "op_type": cOpType, + "server_sequence": cSeq, + "entity_type": op.EntityType, + "entity_id": op.EntityID, + }) + } + if err := conflictRows.Err(); err != nil { + _ = conflictRows.Close() + jsonInternalError(w, err) + return + } + if err := conflictRows.Close(); err != nil { + jsonInternalError(w, err) + return } } - res, err := s.db.Exec( + res, err := tx.Exec( `INSERT OR IGNORE INTO server_ops (op_id, server_sequence, user_id, vault_id, device_id, entity_type, entity_id, op_type, payload_json, idempotency_key, client_sequence, last_seen_server_seq, created_at, pushed_at) VALUES (?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, op.OpID, scope.UserID, scope.VaultID, scope.DeviceID, op.EntityType, op.EntityID, op.OpType, op.PayloadJSON, req.IdempotencyKey, op.ClientSequence, op.LastSeenServerSeq, op.CreatedAt, now, ) if err != nil { - continue + jsonInternalError(w, err) + return + } + n, err := res.RowsAffected() + if err != nil { + jsonInternalError(w, err) + return } - n, _ := res.RowsAffected() if n == 0 { continue } - seqRes, err := s.db.Exec("INSERT INTO server_revisions (op_id, device_id) VALUES (?, ?)", op.OpID, scope.DeviceID) + seqRes, err := tx.Exec("INSERT INTO server_revisions (op_id, device_id) VALUES (?, ?)", op.OpID, scope.DeviceID) if err != nil { - continue + jsonInternalError(w, err) + return + } + seq, err := seqRes.LastInsertId() + if err != nil { + jsonInternalError(w, err) + return + } + if _, err := tx.Exec("UPDATE server_ops SET server_sequence=? WHERE op_id=?", seq, op.OpID); err != nil { + jsonInternalError(w, err) + return } - seq, _ := seqRes.LastInsertId() - s.db.Exec("UPDATE server_ops SET server_sequence=? WHERE op_id=?", seq, op.OpID) if op.OpType == "delete" { - s.db.Exec(`INSERT OR REPLACE INTO server_tombstones + if _, err := tx.Exec(`INSERT OR REPLACE INTO server_tombstones (user_id, vault_id, entity_type, entity_id, op_id, deleted_at) VALUES (?, ?, ?, ?, ?, ?)`, - scope.UserID, scope.VaultID, op.EntityType, op.EntityID, op.OpID, now) + scope.UserID, scope.VaultID, op.EntityType, op.EntityID, op.OpID, now); err != nil { + jsonInternalError(w, err) + return + } } accepted = append(accepted, op.OpID) } @@ -445,39 +523,54 @@ func (s *Server) handleSyncPush(w http.ResponseWriter, r *http.Request) { "conflicts": conflicts, } if req.IdempotencyKey != "" { - if respJSON, err := json.Marshal(resp); err == nil { - s.db.Exec(`INSERT OR IGNORE INTO server_idempotency_keys - (user_id, vault_id, idempotency_key, response_json, created_at) VALUES (?, ?, ?, ?, ?)`, - scope.UserID, scope.VaultID, req.IdempotencyKey, string(respJSON), now) + respJSON, err := json.Marshal(resp) + if err != nil { + jsonInternalError(w, err) + return } + if _, err := tx.Exec(`INSERT INTO server_idempotency_keys + (user_id, vault_id, idempotency_key, response_json, created_at) VALUES (?, ?, ?, ?, ?)`, + scope.UserID, scope.VaultID, req.IdempotencyKey, string(respJSON), now); err != nil { + jsonInternalError(w, err) + return + } + } + if err := tx.Commit(); err != nil { + jsonInternalError(w, err) + return } jsonOK(w, resp) } func (s *Server) handleSyncPull(w http.ResponseWriter, r *http.Request) { + if r.Method != "POST" { + methodNotAllowed(w, "POST") + return + } scope, ok := s.requireSyncScope(w, r) if !ok { return } - if r.Method != "POST" { - jsonErr(w, 405, "POST required") + var req syncPullRequest + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } - var req struct { - SinceSequence int `json:"since_sequence"` - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if req.SinceSequence < 0 || req.PageLimit < 0 { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", "sequence and page_limit must be non-negative") return } + pageLimit := s.pullPageLimit(req.PageLimit) var serverSeq int - s.db.QueryRow(`SELECT COALESCE(MAX(server_sequence), 0) FROM server_ops - WHERE user_id=? AND vault_id=?`, scope.UserID, scope.VaultID).Scan(&serverSeq) + if err := s.db.QueryRow(`SELECT COALESCE(MAX(server_sequence), 0) FROM server_ops + WHERE user_id=? AND vault_id=?`, scope.UserID, scope.VaultID).Scan(&serverSeq); err != nil { + jsonInternalError(w, err) + return + } rows, err := s.db.Query(` SELECT op_id, server_sequence, device_id, entity_type, entity_id, op_type, payload_json, created_at FROM server_ops WHERE user_id=? AND vault_id=? AND server_sequence > ? AND server_sequence IS NOT NULL - ORDER BY server_sequence`, scope.UserID, scope.VaultID, req.SinceSequence) + ORDER BY server_sequence LIMIT ?`, scope.UserID, scope.VaultID, req.SinceSequence, pageLimit+1) if err != nil { jsonInternalError(w, err) return @@ -493,81 +586,47 @@ func (s *Server) handleSyncPull(w http.ResponseWriter, r *http.Request) { PayloadJSON string `json:"payload_json"` CreatedAt string `json:"created_at"` } - ops := []opDTO{} + ops := make([]opDTO, 0, pageLimit) for rows.Next() { var o opDTO if err := rows.Scan(&o.OpID, &o.ServerSequence, &o.DeviceID, &o.EntityType, &o.EntityID, &o.OpType, &o.PayloadJSON, &o.CreatedAt); err != nil { - continue + jsonInternalError(w, err) + return } ops = append(ops, o) } + if err := rows.Err(); err != nil { + jsonInternalError(w, err) + return + } + hasMore := len(ops) > pageLimit + if hasMore { + ops = ops[:pageLimit] + } + pageLastSequence := req.SinceSequence + if len(ops) > 0 { + pageLastSequence = ops[len(ops)-1].ServerSequence + } jsonOK(w, map[string]interface{}{ - "server_sequence": serverSeq, - "ops": ops, + "server_sequence": serverSeq, + "page_last_sequence": pageLastSequence, + "has_more": hasMore, + "ops": ops, }) } func (s *Server) handleBlobs(w http.ResponseWriter, r *http.Request) { - _, _, ok := s.requireAuth(w, r) + scope, ok := s.requireSyncScope(w, r) if !ok { return } switch r.Method { case "POST": - if err := r.ParseMultipartForm(200 << 20); err != nil { - jsonErr(w, 400, "invalid multipart request") - return - } - file, header, err := r.FormFile("file") - if err != nil { - jsonErr(w, 400, "file field required") - return - } - defer file.Close() - data, err := io.ReadAll(file) - if err != nil { - jsonErr(w, 500, "read error") - return - } - hash := sha256Hex(string(data)) - blobDir := filepath.Join(s.blobsDir, hash[:2], hash[2:4]) - if err := os.MkdirAll(blobDir, 0750); err != nil { - jsonErr(w, 500, "mkdir error") - return - } - blobPath := filepath.Join(blobDir, hash) - if err := os.WriteFile(blobPath, data, 0640); err != nil { - jsonErr(w, 500, "write error") - return - } - _ = header - now := time.Now().UTC().Format(time.RFC3339) - s.db.Exec("INSERT OR IGNORE INTO server_blobs (sha256, size, created_at) VALUES (?, ?, ?)", - hash, len(data), now) - jsonOK(w, map[string]interface{}{ - "sha256": hash, - "size": len(data), - }) + s.handleBlobUpload(w, r, scope) case "GET": shaHex := strings.TrimPrefix(r.URL.Path, "/api/v1/blobs/") - if len(shaHex) != 64 { - jsonErr(w, 400, "invalid SHA-256") - return - } - blobPath := filepath.Join(s.blobsDir, shaHex[:2], shaHex[2:4], shaHex) - if _, err := os.Stat(blobPath); os.IsNotExist(err) { - jsonErr(w, 404, "blob not found") - return - } - data, err := os.ReadFile(blobPath) - if err != nil { - jsonErr(w, 500, "read error") - return - } - w.Header().Set("Content-Type", "application/octet-stream") - w.Header().Set("Content-Disposition", "attachment; filename=\""+shaHex+"\"") - w.Write(data) + s.handleBlobDownload(w, r, scope, shaHex) default: - jsonErr(w, 405, "method not allowed") + methodNotAllowed(w, "GET", "POST") } } diff --git a/internal/server/handlers_auth.go b/internal/server/handlers_auth.go index 05957ea..7e93560 100644 --- a/internal/server/handlers_auth.go +++ b/internal/server/handlers_auth.go @@ -3,7 +3,7 @@ package server import ( "crypto/rand" "encoding/hex" - "encoding/json" + "html" "log" "net/http" "strings" @@ -13,8 +13,8 @@ import ( ) func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) { - if r.Method != "POST" { - jsonErr(w, 405, "POST required") + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodPost) return } var req struct { @@ -22,14 +22,16 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) { Email string `json:"email"` Password string `json:"password"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Username == "" || req.Email == "" || req.Password == "" { jsonErr(w, 400, "username, email and password required") return } + if !s.allowRate(w, r, "register", req.Email) { + return + } if err := validatePassword(req.Password); err != "" { jsonErr(w, 400, err) return @@ -47,7 +49,13 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) { id := make([]byte, 12) rand.Read(id) userID := hex.EncodeToString(id) - _, err = s.db.Exec( + tx, err := s.db.Begin() + if err != nil { + jsonInternalError(w, err) + return + } + defer tx.Rollback() + _, err = tx.Exec( "INSERT INTO server_users (id, username, email, password_hash, confirmed, created_at) VALUES (?, ?, ?, ?, 0, ?)", userID, req.Username, strings.ToLower(req.Email), string(hash), now, ) @@ -59,29 +67,54 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) { jsonInternalError(w, err) return } - tok := make([]byte, 24) - rand.Read(tok) - tokenStr := hex.EncodeToString(tok) - exp := time.Now().Add(48 * time.Hour).UTC().Format(time.RFC3339) - s.db.Exec("INSERT INTO server_email_tokens (token, user_id, purpose, expires_at, created_at) VALUES (?, ?, 'confirm', ?, ?)", - tokenStr, userID, exp, now) - log.Printf("register: confirmation token=%s for user %s", tokenStr, req.Username) + if _, err := issueEmailToken(tx, userID, "confirm", 48*time.Hour); err != nil { + jsonInternalError(w, err) + return + } + if err := tx.Commit(); err != nil { + jsonInternalError(w, err) + return + } jsonOK(w, map[string]string{"status": "confirmation_sent"}) } func (s *Server) handleConfirm(w http.ResponseWriter, r *http.Request) { - if r.Method != "GET" { - jsonErr(w, 405, "GET required") + if r.Method == http.MethodGet { + tokenStr := r.URL.Query().Get("token") + if tokenStr == "" { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", "token required") + return + } + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = w.Write([]byte(``)) + return + } + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodGet, http.MethodPost) + return + } + tokenStr := "" + if strings.HasPrefix(r.Header.Get("Content-Type"), "application/json") { + var req struct { + Token string `json:"token"` + } + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { + return + } + tokenStr = req.Token + } else if err := r.ParseForm(); err == nil { + tokenStr = r.FormValue("token") + } else { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", "invalid form") return } - tokenStr := r.URL.Query().Get("token") if tokenStr == "" { jsonErr(w, 400, "token required") return } var userID, expiresAt string - err := s.db.QueryRow("SELECT user_id, expires_at FROM server_email_tokens WHERE token=? AND purpose='confirm'", - tokenStr).Scan(&userID, &expiresAt) + err := s.db.QueryRow("SELECT user_id, expires_at FROM server_email_tokens WHERE token_hash=? AND purpose='confirm'", + emailTokenHash(tokenStr)).Scan(&userID, &expiresAt) if err != nil { jsonErr(w, 400, "invalid or expired token") return @@ -91,29 +124,47 @@ func (s *Server) handleConfirm(w http.ResponseWriter, r *http.Request) { jsonErr(w, 400, "token expired") return } - s.db.Exec("UPDATE server_users SET confirmed=1 WHERE id=?", userID) + tx, err := s.db.Begin() + if err != nil { + jsonInternalError(w, err) + return + } + defer tx.Rollback() + if _, err := tx.Exec("UPDATE server_users SET confirmed=1 WHERE id=?", userID); err != nil { + jsonInternalError(w, err) + return + } + if _, err := tx.Exec("DELETE FROM server_email_tokens WHERE token_hash=?", emailTokenHash(tokenStr)); err != nil { + jsonInternalError(w, err) + return + } + if err := tx.Commit(); err != nil { + jsonInternalError(w, err) + return + } log.Printf("confirm: user %s confirmed email", userID) - s.db.Exec("DELETE FROM server_email_tokens WHERE token=?", tokenStr) jsonOK(w, map[string]string{"status": "confirmed"}) } func (s *Server) handleUserLogin(w http.ResponseWriter, r *http.Request) { - if r.Method != "POST" { - jsonErr(w, 405, "POST required") + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodPost) return } var req struct { Username string `json:"username"` Password string `json:"password"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Username == "" || req.Password == "" { jsonErr(w, 400, "username and password required") return } + if !s.allowRate(w, r, "login", req.Username) { + return + } var userID, hash string var confirmed, blocked int err := s.db.QueryRow("SELECT id, password_hash, confirmed, blocked FROM server_users WHERE username=? OR email=?", @@ -134,61 +185,68 @@ func (s *Server) handleUserLogin(w http.ResponseWriter, r *http.Request) { jsonErr(w, 401, "invalid credentials") return } - s.db.Exec("UPDATE server_users SET last_seen=? WHERE id=?", time.Now().UTC().Format(time.RFC3339), userID) - tok := s.userTokens.Create(userID) + if _, err := s.db.Exec("UPDATE server_users SET last_seen=? WHERE id=?", time.Now().UTC().Format(time.RFC3339), userID); err != nil { + jsonInternalError(w, err) + return + } + tok, _, err := s.createSession(sessionScopeUser, userID) + if err != nil { + jsonInternalError(w, err) + return + } jsonOK(w, map[string]string{"token": tok, "user_id": userID}) } func (s *Server) handleForgot(w http.ResponseWriter, r *http.Request) { - if r.Method != "POST" { - jsonErr(w, 405, "POST required") + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodPost) return } var req struct { Email string `json:"email"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Email == "" { jsonErr(w, 400, "email required") return } + if !s.allowRate(w, r, "forgot", req.Email) { + return + } var userID string err := s.db.QueryRow("SELECT id FROM server_users WHERE email=?", strings.ToLower(req.Email)).Scan(&userID) if err != nil { jsonOK(w, map[string]string{"status": "if email exists, reset link sent"}) return } - tok := make([]byte, 24) - rand.Read(tok) - tokenStr := hex.EncodeToString(tok) - exp := time.Now().Add(1 * time.Hour).UTC().Format(time.RFC3339) - now := time.Now().UTC().Format(time.RFC3339) - s.db.Exec("INSERT INTO server_email_tokens (token, user_id, purpose, expires_at, created_at) VALUES (?, ?, 'reset', ?, ?)", - tokenStr, userID, exp, now) - log.Printf("forgot: reset token=%s for user %s", tokenStr, userID) + if _, err := issueEmailToken(s.db, userID, "reset", time.Hour); err != nil { + jsonInternalError(w, err) + return + } jsonOK(w, map[string]string{"status": "if email exists, reset link sent"}) } func (s *Server) handleReset(w http.ResponseWriter, r *http.Request) { - if r.Method != "POST" { - jsonErr(w, 405, "POST required") + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodPost) return } var req struct { Token string `json:"token"` NewPassword string `json:"new_password"` } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - jsonErr(w, 400, "invalid JSON") + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { return } if req.Token == "" || req.NewPassword == "" { jsonErr(w, 400, "token and new_password required") return } + if !s.allowRate(w, r, "reset", "") { + return + } if err := validatePassword(req.NewPassword); err != "" { jsonErr(w, 400, err) return diff --git a/internal/server/handlers_web_user.go b/internal/server/handlers_web_user.go index 5629c57..47c5e9e 100644 --- a/internal/server/handlers_web_user.go +++ b/internal/server/handlers_web_user.go @@ -16,17 +16,11 @@ import ( ) func (s *Server) requireUserWeb(w http.ResponseWriter, r *http.Request) (string, bool) { - cookie, err := r.Cookie("user_session") - if err != nil { - http.Redirect(w, r, "/login", http.StatusFound) - return "", false - } - userID, ok := s.userTokens.Check(cookie.Value) + session, ok := s.requireSession(w, r, sessionScopeUser) if !ok { - http.Redirect(w, r, "/login", http.StatusFound) return "", false } - return userID, true + return session.SubjectID, true } func (s *Server) handleUserWebRegister(w http.ResponseWriter, r *http.Request) { @@ -51,6 +45,9 @@ func (s *Server) handleUserWebRegister(w http.ResponseWriter, r *http.Request) { w.Write([]byte(errorPageHTML(locale, t(locale, "common.error"), t(locale, "server.allFieldsRequired"), "/register"))) return } + if !s.allowRate(w, r, "register", email) { + return + } if err := validatePassword(password); err != "" { w.Header().Set("Content-Type", "text/html; charset=utf-8") w.WriteHeader(400) @@ -84,12 +81,11 @@ func (s *Server) handleUserWebRegister(w http.ResponseWriter, r *http.Request) { } return } - tok := make([]byte, 24) - rand.Read(tok) - tokenStr := hex.EncodeToString(tok) - exp := time.Now().Add(48 * time.Hour).UTC().Format(time.RFC3339) - s.db.Exec("INSERT INTO server_email_tokens (token, user_id, purpose, expires_at, created_at) VALUES (?, ?, 'confirm', ?, ?)", - tokenStr, userID, exp, now) + tokenStr, err := issueEmailToken(s.db, userID, "confirm", 48*time.Hour) + if err != nil { + jsonInternalError(w, err) + return + } host := s.smtpGet("smtp_host") if host != "" { srvURL := s.smtpGet("server_url") @@ -103,15 +99,11 @@ func (s *Server) handleUserWebRegister(w http.ResponseWriter, r *http.Request) { if err := s.smtpSend(email, t(locale, "server.emailConfirmSubject"), body); err != nil { log.Printf("register web: failed to send confirm email: %v", err) } - } else { - log.Printf("register web: SMTP not configured, confirmation token=%s for user %s", tokenStr, username) + } else if s.cfg.DevelopmentTokenLogging { + log.Printf("development confirmation token for user %s: %s", username, tokenStr) } w.Header().Set("Content-Type", "text/html; charset=utf-8") - regMsg := registrationOKHTML(locale) - if host == "" { - regMsg = registrationAutoHTML(locale) - } - w.Write([]byte(regMsg)) + w.Write([]byte(registrationOKHTML(locale))) default: jsonErr(w, 405, "method not allowed") } @@ -134,6 +126,9 @@ func (s *Server) handleUserWebForgot(w http.ResponseWriter, r *http.Request) { w.Write([]byte(errorPageHTML(locale, t(locale, "common.error"), t(locale, "server.needEmail"), "/forgot"))) return } + if !s.allowRate(w, r, "forgot", email) { + return + } var userID string err := s.db.QueryRow("SELECT id FROM server_users WHERE email=?", email).Scan(&userID) if err != nil { @@ -141,13 +136,11 @@ func (s *Server) handleUserWebForgot(w http.ResponseWriter, r *http.Request) { w.Write([]byte(forgotSentHTML(locale))) return } - tok := make([]byte, 24) - rand.Read(tok) - tokenStr := hex.EncodeToString(tok) - exp := time.Now().Add(1 * time.Hour).UTC().Format(time.RFC3339) - now := time.Now().UTC().Format(time.RFC3339) - s.db.Exec("INSERT INTO server_email_tokens (token, user_id, purpose, expires_at, created_at) VALUES (?, ?, 'reset', ?, ?)", - tokenStr, userID, exp, now) + tokenStr, err := issueEmailToken(s.db, userID, "reset", time.Hour) + if err != nil { + jsonInternalError(w, err) + return + } host := s.smtpGet("smtp_host") if host != "" { srvURL := s.smtpGet("server_url") @@ -159,8 +152,8 @@ func (s *Server) handleUserWebForgot(w http.ResponseWriter, r *http.Request) { if err := s.smtpSend(email, t(locale, "server.emailResetSubject"), body); err != nil { log.Printf("forgot web: failed to send reset email: %v", err) } - } else { - log.Printf("forgot web: SMTP not configured, reset token=%s for email %s", tokenStr, email) + } else if s.cfg.DevelopmentTokenLogging { + log.Printf("development reset token requested for %s: %s", email, tokenStr) } w.Header().Set("Content-Type", "text/html; charset=utf-8") w.Write([]byte(forgotSentHTML(locale))) @@ -179,8 +172,8 @@ func (s *Server) handleUserWebReset(w http.ResponseWriter, r *http.Request) { return } var userID, expiresAt string - err := s.db.QueryRow("SELECT user_id, expires_at FROM server_email_tokens WHERE token=? AND purpose='reset'", - token).Scan(&userID, &expiresAt) + err := s.db.QueryRow("SELECT user_id, expires_at FROM server_email_tokens WHERE token_hash=? AND purpose='reset'", + emailTokenHash(token)).Scan(&userID, &expiresAt) if err != nil { http.Redirect(w, r, "/forgot", http.StatusFound) return @@ -206,6 +199,9 @@ func (s *Server) handleUserWebReset(w http.ResponseWriter, r *http.Request) { w.Write([]byte(errorPageHTML(locale, t(locale, "common.error"), t(locale, "server.allFieldsRequired"), "/forgot"))) return } + if !s.allowRate(w, r, "reset", "") { + return + } if err := validatePassword(newPass); err != "" { w.Header().Set("Content-Type", "text/html; charset=utf-8") w.Write([]byte(errorPageHTML(locale, t(locale, "common.error"), string(err), "/reset?token="+url.QueryEscape(token)))) @@ -248,6 +244,9 @@ func (s *Server) handleUserWebLogin(w http.ResponseWriter, r *http.Request) { } username := r.FormValue("username") password := r.FormValue("password") + if !s.allowRate(w, r, "login", username) { + return + } var userID, hash string var confirmed, blocked int err := s.db.QueryRow("SELECT id, password_hash, confirmed, blocked FROM server_users WHERE username=? OR email=?", @@ -258,12 +257,12 @@ func (s *Server) handleUserWebLogin(w http.ResponseWriter, r *http.Request) { w.Write([]byte(errorPageHTML(locale, "401 Unauthorized", "401 Unauthorized", "/login"))) return } - tok := s.userTokens.Create(userID) - http.SetCookie(w, &http.Cookie{ - Name: "user_session", Value: tok, Path: "/", - HttpOnly: true, SameSite: http.SameSiteLaxMode, - MaxAge: 86400, - }) + tok, csrf, err := s.createSession(sessionScopeUser, userID) + if err != nil { + jsonInternalError(w, err) + return + } + s.setSessionCookies(w, r, sessionScopeUser, tok, csrf) http.Redirect(w, r, "/dashboard", http.StatusFound) default: jsonErr(w, 405, "method not allowed") @@ -329,18 +328,36 @@ func (s *Server) handleUserDashboard(w http.ResponseWriter, r *http.Request) { } } - w.Write([]byte(userDashboardHTML(locale, html.EscapeString(username), deviceRows))) + csrf := "" + if cookie, err := r.Cookie("csrf_token"); err == nil { + csrf = cookie.Value + } + w.Write([]byte(userDashboardHTML(locale, html.EscapeString(username), deviceRows, csrf))) } func (s *Server) handleUserWebLogout(w http.ResponseWriter, r *http.Request) { - http.SetCookie(w, &http.Cookie{ - Name: "user_session", Value: "", Path: "/", - HttpOnly: true, MaxAge: -1, - }) + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodPost) + return + } + if !s.requireUserMutation(w, r) { + return + } + if cookie, err := r.Cookie("user_session"); err == nil { + if err := s.deleteSession(cookie.Value); err != nil { + jsonInternalError(w, err) + return + } + } + s.clearSessionCookies(w, r, sessionScopeUser) http.Redirect(w, r, "/login", http.StatusFound) } func (s *Server) handleUserDevices(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + methodNotAllowed(w, http.MethodGet) + return + } userID, ok := s.requireUserWeb(w, r) if !ok { return @@ -374,3 +391,58 @@ func (s *Server) handleUserDevices(w http.ResponseWriter, r *http.Request) { } jsonOK(w, devices) } + +// handleUserWebDeviceAction is deliberately session/CSRF based. Browser UI +// never receives a desktop device bearer token. +func (s *Server) handleUserWebDeviceAction(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + methodNotAllowed(w, http.MethodPost) + return + } + session, ok := s.requireSession(w, r, sessionScopeUser) + if !ok || !s.verifyCSRF(w, r, session) { + return + } + path := strings.TrimPrefix(r.URL.Path, "/api/v1/user/devices/") + if !strings.HasSuffix(path, "/revoke") { + jsonErr(w, http.StatusNotFound, "not found") + return + } + deviceID := strings.TrimSuffix(strings.TrimSuffix(path, "/revoke"), "/") + var req struct { + Password string `json:"password"` + } + if !decodeJSONBody(w, r, &req, s.cfg.Limits.MaxJSONBody) { + return + } + if req.Password == "" || !s.allowRate(w, r, "auth-test", session.SubjectID) { + if req.Password == "" { + jsonErrCode(w, http.StatusBadRequest, "invalid_request", "password required") + } + return + } + var hash string + if err := s.db.QueryRow("SELECT password_hash FROM server_users WHERE id=?", session.SubjectID).Scan(&hash); err != nil { + jsonInternalError(w, err) + return + } + if bcrypt.CompareHashAndPassword([]byte(hash), []byte(req.Password)) != nil { + jsonErr(w, http.StatusForbidden, "wrong password") + return + } + var owner string + if err := s.db.QueryRow("SELECT user_id FROM server_devices WHERE id=?", deviceID).Scan(&owner); err != nil { + jsonErr(w, http.StatusNotFound, "device not found") + return + } + if owner != session.SubjectID { + jsonErr(w, http.StatusForbidden, "device does not belong to you") + return + } + if err := s.revokeDevice(deviceID, time.Now().UTC().Format(time.RFC3339)); err != nil { + jsonInternalError(w, err) + return + } + s.auditLog("device_revoked", session.SubjectID, deviceID, s.clientIP(r), "device revoked from web dashboard") + jsonOK(w, map[string]string{"status": "revoked"}) +} diff --git a/internal/server/hardening_test.go b/internal/server/hardening_test.go new file mode 100644 index 0000000..825391c --- /dev/null +++ b/internal/server/hardening_test.go @@ -0,0 +1,645 @@ +package server + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "mime/multipart" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" + "time" +) + +func newTestServer(t *testing.T) (*Server, error) { + t.Helper() + return newServerForTest(t, DefaultConfig()) +} + +func newServerForTest(t *testing.T, cfg *Config) (*Server, error) { + t.Helper() + dir := t.TempDir() + return NewServer(filepath.Join(dir, "server.db"), filepath.Join(dir, "data"), cfg) +} + +func serveJSON(t *testing.T, s *Server, method, path, token string, body interface{}) (int, map[string]interface{}) { + t.Helper() + data, err := json.Marshal(body) + if err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(method, path, bytes.NewReader(data)) + req.Header.Set("Content-Type", "application/json") + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + res := httptest.NewRecorder() + s.Handler().ServeHTTP(res, req) + result := make(map[string]interface{}) + if len(res.Body.Bytes()) > 0 { + if err := json.Unmarshal(res.Body.Bytes(), &result); err != nil { + t.Fatalf("decode response: %v (%s)", err, res.Body.String()) + } + } + return res.Code, result +} + +func insertScopedSyncDevice(t *testing.T, s *Server, deviceID, userID, vaultID, token string) { + t.Helper() + now := time.Now().UTC().Format(time.RFC3339) + if _, err := s.db.Exec(`INSERT INTO server_devices + (id, name, api_key, token_hash, user_id, vault_id, last_seen, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, deviceID, deviceID, "legacy-"+deviceID, sha256Hex(token), userID, vaultID, now, now); err != nil { + t.Fatalf("insert scoped device: %v", err) + } +} + +func uploadBlob(t *testing.T, s *Server, token string, data []byte) (int, map[string]interface{}) { + t.Helper() + var body bytes.Buffer + writer := multipart.NewWriter(&body) + part, err := writer.CreateFormFile("file", "blob.bin") + if err != nil { + t.Fatal(err) + } + if _, err := part.Write(data); err != nil { + t.Fatal(err) + } + if err := writer.Close(); err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/blobs/", &body) + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Content-Type", writer.FormDataContentType()) + res := httptest.NewRecorder() + s.Handler().ServeHTTP(res, req) + result := map[string]interface{}{} + if err := json.Unmarshal(res.Body.Bytes(), &result); err != nil { + t.Fatalf("decode blob response: %v (%s)", err, res.Body.String()) + } + return res.Code, result +} + +func TestDefaultListenAddressIsLoopback(t *testing.T) { + cfg := DefaultConfig() + if got, want := cfg.ListenAddress(), "127.0.0.1:47732"; got != want { + t.Fatalf("default listen address = %q, want %q", got, want) + } +} + +func TestHTTPServerUsesExplicitTimeouts(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + + httpServer := s.HTTPServer("127.0.0.1:0") + if httpServer.ReadHeaderTimeout <= 0 || httpServer.ReadTimeout <= 0 || httpServer.WriteTimeout <= 0 || httpServer.IdleTimeout <= 0 { + t.Fatalf("HTTP timeouts must all be set: %#v", httpServer) + } + if httpServer.MaxHeaderBytes <= 0 { + t.Fatalf("MaxHeaderBytes must be set: %#v", httpServer) + } +} + +func TestHTTPServerGracefulShutdown(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := s.HTTPServer(listener.Addr().String()) + server.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) }) + done := make(chan error, 1) + go func() { done <- server.Serve(listener) }() + + response, err := http.Get("http://" + listener.Addr().String()) + if err != nil { + t.Fatal(err) + } + _ = response.Body.Close() + shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := server.Shutdown(shutdownCtx); err != nil { + t.Fatalf("shutdown: %v", err) + } + if err := <-done; err != nil && err != http.ErrServerClosed { + t.Fatalf("serve returned %v", err) + } +} + +func TestClientIPIgnoresUntrustedForwardedHeaders(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + + r := httptest.NewRequest(http.MethodPost, "/api/client/pair", nil) + r.RemoteAddr = "198.51.100.20:4040" + r.Header.Set("X-Forwarded-For", "203.0.113.8") + if got, want := s.clientIP(r), "198.51.100.20"; got != want { + t.Fatalf("client IP = %q, want %q", got, want) + } +} + +func TestClientIPUsesTrustedProxyHeadersOnlyForTrustedPeer(t *testing.T) { + cfg := DefaultConfig() + cfg.TrustedProxies = []string{"127.0.0.1/32"} + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + + r := httptest.NewRequest(http.MethodPost, "/api/client/pair", nil) + r.RemoteAddr = "127.0.0.1:4040" + r.Header.Set("X-Forwarded-For", "203.0.113.8, 127.0.0.1") + if got, want := s.clientIP(r), "203.0.113.8"; got != want { + t.Fatalf("client IP = %q, want %q", got, want) + } +} + +func TestRateLimiterRecoversAfterWindow(t *testing.T) { + now := time.Now().UTC() + limiter := newRateLimiter(func() time.Time { return now }) + policy := RatePolicy{Limit: 2, Window: time.Minute} + if allowed, _ := limiter.Allow("198.51.100.1", policy); !allowed { + t.Fatal("first attempt unexpectedly limited") + } + if allowed, _ := limiter.Allow("198.51.100.1", policy); !allowed { + t.Fatal("second attempt unexpectedly limited") + } + if allowed, retryAfter := limiter.Allow("198.51.100.1", policy); allowed || retryAfter <= 0 { + t.Fatalf("third attempt = allowed:%t retry:%s, want limited with retry", allowed, retryAfter) + } + now = now.Add(time.Minute + time.Second) + if allowed, _ := limiter.Allow("198.51.100.1", policy); !allowed { + t.Fatal("attempt after window remained limited") + } +} + +func TestRateLimiterMemoryIsBounded(t *testing.T) { + limiter := newRateLimiter(nil) + policy := RatePolicy{Limit: 1, Window: time.Hour} + for i := 0; i < maxRateLimitBuckets+100; i++ { + if allowed, _ := limiter.Allow(strconvItoa(i), policy); !allowed { + t.Fatalf("new bucket %d unexpectedly limited", i) + } + } + if got := len(limiter.buckets); got > maxRateLimitBuckets { + t.Fatalf("rate buckets = %d, want <= %d", got, maxRateLimitBuckets) + } +} + +func TestSyncPushRejectsOversizedJSONBody(t *testing.T) { + cfg := DefaultConfig() + cfg.Limits.MaxJSONBody = 64 + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertSyncDevice(t, s, "device-a", "user-a", "token-a") + + overSizedJSON := append([]byte(`{"ops":[],"padding":"`), bytes.Repeat([]byte("x"), 65)...) + overSizedJSON = append(overSizedJSON, []byte(`"}`)...) + req := httptest.NewRequest(http.MethodPost, "/api/v1/sync/push", bytes.NewReader(overSizedJSON)) + req.Header.Set("Authorization", "Bearer token-a") + res := httptest.NewRecorder() + s.Handler().ServeHTTP(res, req) + if res.Code != http.StatusRequestEntityTooLarge { + t.Fatalf("status = %d, want %d: %s", res.Code, http.StatusRequestEntityTooLarge, res.Body.String()) + } + var body map[string]string + if err := json.Unmarshal(res.Body.Bytes(), &body); err != nil { + t.Fatal(err) + } + if body["code"] != "request_too_large" { + t.Fatalf("error body = %#v, want stable request_too_large code", body) + } +} + +func TestSyncPushRejectsOperationCountAboveLimit(t *testing.T) { + cfg := DefaultConfig() + cfg.Limits.MaxPushOperations = 1 + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertSyncDevice(t, s, "device-a", "user-a", "token-a") + + body := syncPushBody("device-a", "op-1", "") + body["ops"] = append(body["ops"].([]map[string]interface{}), map[string]interface{}{ + "op_id": "op-2", "entity_type": "file", "entity_id": "Docs/two.txt", "op_type": "create", + "payload_json": `{"path":"Docs/two.txt","content":"two"}`, "created_at": "2026-07-10T00:00:00Z", + }) + status, response := serveJSON(t, s, http.MethodPost, "/api/v1/sync/push", "token-a", body) + if status != http.StatusRequestEntityTooLarge || response["code"] != "too_many_operations" { + t.Fatalf("status=%d body=%#v, want 413 too_many_operations", status, response) + } +} + +func TestSyncPushRejectsTrailingJSONAndOversizedPayload(t *testing.T) { + cfg := DefaultConfig() + cfg.Limits.MaxPayloadJSON = 32 + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertSyncDevice(t, s, "device-a", "user-a", "token-a") + + data, err := json.Marshal(syncPushBody("device-a", "op-trailing", "")) + if err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, "/api/v1/sync/push", bytes.NewReader(append(data, []byte(` {}`)...))) + req.Header.Set("Authorization", "Bearer token-a") + res := httptest.NewRecorder() + s.Handler().ServeHTTP(res, req) + if res.Code != http.StatusBadRequest || !bytes.Contains(res.Body.Bytes(), []byte(`"trailing_json"`)) { + t.Fatalf("trailing JSON status=%d body=%s", res.Code, res.Body.String()) + } + + body := syncPushBody("device-a", "op-large-payload", "") + body["ops"].([]map[string]interface{})[0]["payload_json"] = `{"path":"Docs/large.txt","content":"this is intentionally longer than the configured payload bound"}` + status, response := serveJSON(t, s, http.MethodPost, "/api/v1/sync/push", "token-a", body) + if status != http.StatusRequestEntityTooLarge || response["code"] != "payload_too_large" { + t.Fatalf("payload limit status=%d body=%#v", status, response) + } +} + +func TestSyncPullPaginationHasNoGapsOrRepeats(t *testing.T) { + cfg := DefaultConfig() + cfg.Limits.MaxPullPage = 2 + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertSyncDevice(t, s, "device-a", "user-a", "token-a") + for _, opID := range []string{"op-1", "op-2", "op-3", "op-4", "op-5"} { + if status, response := serveJSON(t, s, http.MethodPost, "/api/v1/sync/push", "token-a", syncPushBody("device-a", opID, "")); status != http.StatusOK { + t.Fatalf("push %s status=%d body=%#v", opID, status, response) + } + } + + cursor := 0 + var sequences []int + for page := 0; page < 3; page++ { + status, response := serveJSON(t, s, http.MethodPost, "/api/v1/sync/pull", "token-a", map[string]int{"since_sequence": cursor, "page_limit": 2}) + if status != http.StatusOK { + t.Fatalf("pull page %d status=%d body=%#v", page, status, response) + } + for _, raw := range response["ops"].([]interface{}) { + sequences = append(sequences, int(raw.(map[string]interface{})["server_sequence"].(float64))) + } + cursor = int(response["page_last_sequence"].(float64)) + if !response["has_more"].(bool) { + break + } + } + if got, want := len(sequences), 5; got != want { + t.Fatalf("sequences=%v, want five ordered values", sequences) + } + for i, sequence := range sequences { + if sequence != i+1 { + t.Fatalf("sequences=%v, want [1 2 3 4 5]", sequences) + } + } +} + +func TestBlobOwnershipPreventsCrossVaultDownload(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertSyncUser(t, s, "user-b") + insertScopedSyncDevice(t, s, "device-a", "user-a", "vault-a", "token-a") + insertScopedSyncDevice(t, s, "device-b", "user-b", "vault-b", "token-b") + + status, uploaded := uploadBlob(t, s, "token-a", []byte("private blob")) + if status != http.StatusOK { + t.Fatalf("upload status=%d body=%#v", status, uploaded) + } + sha := uploaded["sha256"].(string) + + request := httptest.NewRequest(http.MethodGet, "/api/v1/blobs/"+sha, nil) + request.Header.Set("Authorization", "Bearer token-b") + response := httptest.NewRecorder() + s.Handler().ServeHTTP(response, request) + if response.Code != http.StatusNotFound { + t.Fatalf("cross-vault download status=%d, want 404: %s", response.Code, response.Body.String()) + } +} + +func TestBlobLimitAndQuotaRejectWithoutResidualFile(t *testing.T) { + cfg := DefaultConfig() + cfg.Limits.MaxBlobBytes = 8 + cfg.Limits.MaxVaultBlobBytes = 8 + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertScopedSyncDevice(t, s, "device-a", "user-a", "vault-a", "token-a") + + tooLarge := []byte("012345678") + status, body := uploadBlob(t, s, "token-a", tooLarge) + if status != http.StatusRequestEntityTooLarge || body["code"] != "blob_too_large" { + t.Fatalf("file limit status=%d body=%#v", status, body) + } + sum := sha256.Sum256(tooLarge) + path := blobPath(s.blobsDir, hex.EncodeToString(sum[:])) + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Fatalf("rejected oversized blob left physical file: %v", err) + } + + status, body = uploadBlob(t, s, "token-a", []byte("12345678")) + if status != http.StatusOK { + t.Fatalf("first quota upload status=%d body=%#v", status, body) + } + quotaCandidate := []byte("abcdefgh") + status, body = uploadBlob(t, s, "token-a", quotaCandidate) + if status != http.StatusRequestEntityTooLarge || body["code"] != "quota_exceeded" { + t.Fatalf("quota status=%d body=%#v", status, body) + } + sum = sha256.Sum256(quotaCandidate) + if _, err := os.Stat(blobPath(s.blobsDir, hex.EncodeToString(sum[:]))); !os.IsNotExist(err) { + t.Fatalf("quota-rejected blob left physical file: %v", err) + } +} + +func TestBlobUploadIsIdempotentWithinScope(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertScopedSyncDevice(t, s, "device-a", "user-a", "vault-a", "token-a") + + data := []byte("same content") + _, first := uploadBlob(t, s, "token-a", data) + _, second := uploadBlob(t, s, "token-a", data) + if first["sha256"] != second["sha256"] || first["size"] != second["size"] { + t.Fatalf("idempotent uploads differ: first=%#v second=%#v", first, second) + } + var refs int + if err := s.db.QueryRow("SELECT COUNT(*) FROM server_blob_refs WHERE user_id=? AND vault_id=?", "user-a", "vault-a").Scan(&refs); err != nil { + t.Fatal(err) + } + if refs != 1 { + t.Fatalf("blob refs = %d, want one", refs) + } +} + +func TestRevokedDeviceCannotUseBlobEndpoints(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertScopedSyncDevice(t, s, "device-a", "user-a", "vault-a", "token-a") + status, uploaded := uploadBlob(t, s, "token-a", []byte("before revoke")) + if status != http.StatusOK { + t.Fatalf("upload status=%d body=%#v", status, uploaded) + } + if err := s.revokeDevice("device-a", time.Now().UTC().Format(time.RFC3339)); err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodGet, "/api/v1/blobs/"+uploaded["sha256"].(string), nil) + request.Header.Set("Authorization", "Bearer token-a") + response := httptest.NewRecorder() + s.Handler().ServeHTTP(response, request) + if response.Code != http.StatusUnauthorized { + t.Fatalf("revoked blob download status=%d, want 401", response.Code) + } +} + +func TestAdminKeysNeverReturnPlaintextCredential(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertScopedSyncDevice(t, s, "device-a", "user-a", "vault-a", "secret-device-token") + token, _, err := s.createSession(sessionScopeAdmin, "admin") + if err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodGet, "/admin/api/keys", nil) + request.AddCookie(&http.Cookie{Name: "admin_session", Value: token}) + response := httptest.NewRecorder() + s.Handler().ServeHTTP(response, request) + if response.Code != http.StatusOK || bytes.Contains(response.Body.Bytes(), []byte("secret-device-token")) || bytes.Contains(response.Body.Bytes(), []byte("api_key")) { + t.Fatalf("admin keys leaked a credential: status=%d body=%s", response.Code, response.Body.String()) + } +} + +func TestHealthReportsDegradedDatabase(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + s.SetupRoutes() + if err := s.db.Close(); err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodGet, "/readyz", nil) + response := httptest.NewRecorder() + s.Handler().ServeHTTP(response, request) + if response.Code != http.StatusServiceUnavailable || !bytes.Contains(response.Body.Bytes(), []byte(`"database_reachable":false`)) { + t.Fatalf("readiness after db close: status=%d body=%s", response.Code, response.Body.String()) + } +} + +func TestRetentionDoesNotPruneOperations(t *testing.T) { + cfg := DefaultConfig() + cfg.Retention.IdempotencyHours = 1 + cfg.Retention.AuditDays = 1 + s, err := newServerForTest(t, cfg) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertSyncDevice(t, s, "device-a", "user-a", "token-a") + if status, body := serveJSON(t, s, http.MethodPost, "/api/v1/sync/push", "token-a", syncPushBody("device-a", "op-retained", "")); status != http.StatusOK { + t.Fatalf("push status=%d body=%#v", status, body) + } + if err := s.CleanupRetention(time.Now().UTC().Add(48 * time.Hour)); err != nil { + t.Fatal(err) + } + var ops int + if err := s.db.QueryRow("SELECT COUNT(*) FROM server_ops").Scan(&ops); err != nil { + t.Fatal(err) + } + if ops != 1 { + t.Fatalf("retention removed sync operations: %d", ops) + } +} + +func TestWebSessionSurvivesServerRestartAndLogoutInvalidatesIt(t *testing.T) { + dir := t.TempDir() + dbPath := filepath.Join(dir, "server.db") + dataDir := filepath.Join(dir, "data") + s, err := NewServer(dbPath, dataDir, DefaultConfig()) + if err != nil { + t.Fatal(err) + } + token, csrf, err := s.createSession(sessionScopeUser, "user-a") + if err != nil { + t.Fatal(err) + } + if err := s.Close(); err != nil { + t.Fatal(err) + } + + restarted, err := NewServer(dbPath, dataDir, DefaultConfig()) + if err != nil { + t.Fatal(err) + } + defer restarted.Close() + restarted.SetupRoutes() + req := httptest.NewRequest(http.MethodPost, "/logout", nil) + req.AddCookie(&http.Cookie{Name: "user_session", Value: token}) + req.AddCookie(&http.Cookie{Name: "csrf_token", Value: csrf}) + req.Header.Set("X-CSRF-Token", csrf) + res := httptest.NewRecorder() + restarted.Handler().ServeHTTP(res, req) + if res.Code != http.StatusFound { + t.Fatalf("logout status=%d body=%s", res.Code, res.Body.String()) + } + if _, ok := restarted.loadSession(token, sessionScopeUser); ok { + t.Fatal("logout did not invalidate server-side session") + } +} + +func TestAdminMutationRejectsMissingCSRFToken(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + token, csrf, err := s.createSession(sessionScopeAdmin, "admin") + if err != nil { + t.Fatal(err) + } + missing := httptest.NewRequest(http.MethodDelete, "/admin/api/keys/missing", nil) + missing.AddCookie(&http.Cookie{Name: "admin_session", Value: token}) + missingResult := httptest.NewRecorder() + s.Handler().ServeHTTP(missingResult, missing) + if missingResult.Code != http.StatusForbidden { + t.Fatalf("missing csrf status=%d, want 403", missingResult.Code) + } + + valid := httptest.NewRequest(http.MethodDelete, "/admin/api/keys/missing", nil) + valid.AddCookie(&http.Cookie{Name: "admin_session", Value: token}) + valid.AddCookie(&http.Cookie{Name: "csrf_token", Value: csrf}) + valid.Header.Set("X-CSRF-Token", csrf) + validResult := httptest.NewRecorder() + s.Handler().ServeHTTP(validResult, valid) + if validResult.Code != http.StatusOK { + t.Fatalf("valid csrf status=%d body=%s", validResult.Code, validResult.Body.String()) + } +} + +func TestAdminUserDeletionIsTransactionalAcrossOwnedRows(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + s.SetupRoutes() + insertSyncUser(t, s, "user-a") + insertScopedSyncDevice(t, s, "device-a", "user-a", "vault-a", "token-a") + if status, response := serveJSON(t, s, http.MethodPost, "/api/v1/sync/push", "token-a", syncPushBody("device-a", "op-a", "")); status != http.StatusOK { + t.Fatalf("push status=%d body=%#v", status, response) + } + adminToken, csrf, err := s.createSession(sessionScopeAdmin, "admin") + if err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodDelete, "/admin/api/users/user-a", nil) + request.AddCookie(&http.Cookie{Name: "admin_session", Value: adminToken}) + request.AddCookie(&http.Cookie{Name: "csrf_token", Value: csrf}) + request.Header.Set("X-CSRF-Token", csrf) + response := httptest.NewRecorder() + s.Handler().ServeHTTP(response, request) + if response.Code != http.StatusOK { + t.Fatalf("delete status=%d body=%s", response.Code, response.Body.String()) + } + for table, where := range map[string]string{ + "server_users": "id='user-a'", + "server_devices": "user_id='user-a'", + "server_ops": "user_id='user-a'", + "server_blob_refs": "user_id='user-a'", + } { + var count int + if err := s.db.QueryRow("SELECT COUNT(*) FROM " + table + " WHERE " + where).Scan(&count); err != nil { + t.Fatal(err) + } + if count != 0 { + t.Fatalf("%s still has %d rows after user deletion", table, count) + } + } +} + +func TestResetTokenIsHashedAndSingleUse(t *testing.T) { + s, err := newTestServer(t) + if err != nil { + t.Fatal(err) + } + defer s.Close() + insertPairableUser(t, s, "user-a", "alice", "correct horse battery staple") + token, err := issueEmailToken(s.db, "user-a", "reset", time.Hour) + if err != nil { + t.Fatal(err) + } + var stored string + if err := s.db.QueryRow("SELECT token_hash FROM server_email_tokens WHERE user_id=?", "user-a").Scan(&stored); err != nil { + t.Fatal(err) + } + if stored == token || stored != emailTokenHash(token) { + t.Fatalf("stored reset credential is not a token hash: %q", stored) + } + if _, err := s.resetPasswordWithToken(token, "a new secure password"); err != nil { + t.Fatalf("first reset: %v", err) + } + if _, err := s.resetPasswordWithToken(token, "another secure password"); err != errResetTokenInvalid { + t.Fatalf("second reset err=%v, want invalid token", err) + } +} diff --git a/internal/server/health.go b/internal/server/health.go new file mode 100644 index 0000000..5399c6c --- /dev/null +++ b/internal/server/health.go @@ -0,0 +1,116 @@ +package server + +import ( + "context" + "fmt" + "io/fs" + "os" + "path/filepath" + "time" +) + +type HealthStatus struct { + Status string `json:"status"` + Version string `json:"version"` + BuildCommit string `json:"build_commit"` + UptimeSeconds int64 `json:"uptime_seconds"` + DatabaseReachable bool `json:"database_reachable"` + BlobStorageWritable bool `json:"blob_storage_writable"` + SchemaVersion int `json:"schema_version"` + ServerTime string `json:"server_time"` +} + +func (s *Server) healthStatus(ctx context.Context) HealthStatus { + health := HealthStatus{ + Status: "ok", + Version: Version, + BuildCommit: BuildCommit, + UptimeSeconds: int64(time.Since(s.startedAt).Seconds()), + ServerTime: time.Now().UTC().Format(time.RFC3339), + } + if s.db == nil || s.db.PingContext(ctx) != nil { + health.DatabaseReachable = false + health.Status = "degraded" + } else { + health.DatabaseReachable = true + if err := s.db.QueryRowContext(ctx, "PRAGMA user_version").Scan(&health.SchemaVersion); err != nil { + health.DatabaseReachable = false + health.Status = "degraded" + } + } + health.BlobStorageWritable = s.blobStorageWritable() + if !health.BlobStorageWritable { + health.Status = "degraded" + } + return health +} + +func (s *Server) blobStorageWritable() bool { + if s.blobsDir == "" { + return false + } + probe, err := os.CreateTemp(s.blobsDir, ".health-*") + if err != nil { + return false + } + name := probe.Name() + if err := probe.Close(); err != nil { + _ = os.Remove(name) + return false + } + return os.Remove(name) == nil +} + +// ServerStats is intentionally independent from the web UI so a future +// admin panel can expose operational data without coupling to templates. +type ServerStats struct { + Users int `json:"users"` + ActiveDevices int `json:"active_devices"` + RevokedDevices int `json:"revoked_devices"` + Vaults int `json:"vaults"` + Operations int `json:"operations"` + DatabaseBytes int64 `json:"database_bytes"` + BlobBytes int64 `json:"blob_bytes"` + LastSyncAt string `json:"last_sync_activity"` +} + +func (s *Server) Stats(ctx context.Context) (ServerStats, error) { + var stats ServerStats + queries := []struct { + query string + target *int + }{ + {"SELECT COUNT(*) FROM server_users", &stats.Users}, + {"SELECT COUNT(*) FROM server_devices WHERE COALESCE(revoked_at, '') = ''", &stats.ActiveDevices}, + {"SELECT COUNT(*) FROM server_devices WHERE COALESCE(revoked_at, '') != ''", &stats.RevokedDevices}, + {"SELECT COUNT(DISTINCT user_id || ':' || vault_id) FROM server_devices WHERE COALESCE(user_id,'') != '' AND COALESCE(vault_id,'') != ''", &stats.Vaults}, + {"SELECT COUNT(*) FROM server_ops", &stats.Operations}, + } + for _, query := range queries { + if err := s.db.QueryRowContext(ctx, query.query).Scan(query.target); err != nil { + return ServerStats{}, err + } + } + if err := s.db.QueryRowContext(ctx, "SELECT COALESCE(MAX(last_seen), '') FROM server_devices").Scan(&stats.LastSyncAt); err != nil { + return ServerStats{}, err + } + if info, err := os.Stat(s.dbPath); err == nil { + stats.DatabaseBytes = info.Size() + } + if err := filepath.WalkDir(s.blobsDir, func(path string, entry fs.DirEntry, walkErr error) error { + if walkErr != nil { + return walkErr + } + if entry.Type().IsRegular() { + info, err := entry.Info() + if err != nil { + return err + } + stats.BlobBytes += info.Size() + } + return nil + }); err != nil && !os.IsNotExist(err) { + return ServerStats{}, fmt.Errorf("blob storage stats: %w", err) + } + return stats, nil +} diff --git a/internal/server/helpers.go b/internal/server/helpers.go index fab3447..8e3575a 100644 --- a/internal/server/helpers.go +++ b/internal/server/helpers.go @@ -4,8 +4,12 @@ import ( "crypto/sha256" "encoding/hex" "encoding/json" + "errors" + "fmt" + "io" "log" "net/http" + "strings" ) func jsonOK(w http.ResponseWriter, v interface{}) { @@ -13,10 +17,43 @@ func jsonOK(w http.ResponseWriter, v interface{}) { json.NewEncoder(w).Encode(v) } -func jsonErr(w http.ResponseWriter, code int, msg string) { +func jsonOKStatus(w http.ResponseWriter, status int, v interface{}) { w.Header().Set("Content-Type", "application/json") - w.WriteHeader(code) - json.NewEncoder(w).Encode(map[string]string{"error": msg}) + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(v) +} + +func jsonErr(w http.ResponseWriter, code int, msg string) { + jsonErrCode(w, code, defaultErrorCode(code), msg) +} + +// jsonErrCode preserves the legacy human-readable error field while giving +// Desktop a stable machine-readable code for localized messages. +func jsonErrCode(w http.ResponseWriter, status int, machineCode, msg string) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + json.NewEncoder(w).Encode(map[string]string{"error": msg, "code": machineCode}) +} + +func defaultErrorCode(status int) string { + switch status { + case http.StatusBadRequest: + return "invalid_request" + case http.StatusUnauthorized: + return "unauthorized" + case http.StatusForbidden: + return "forbidden" + case http.StatusNotFound: + return "not_found" + case http.StatusMethodNotAllowed: + return "method_not_allowed" + case http.StatusRequestEntityTooLarge: + return "request_too_large" + case http.StatusTooManyRequests: + return "rate_limited" + default: + return "internal_error" + } } func jsonInternalError(w http.ResponseWriter, err error) { @@ -24,6 +61,52 @@ func jsonInternalError(w http.ResponseWriter, err error) { jsonErr(w, http.StatusInternalServerError, "internal error") } +func methodNotAllowed(w http.ResponseWriter, allowed ...string) { + w.Header().Set("Allow", strings.Join(allowed, ", ")) + jsonErrCode(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed") +} + +// decodeJSONBody enforces a hard byte limit, rejects a second JSON value, and +// keeps all request handlers on the same error contract. +func decodeJSONBody(w http.ResponseWriter, r *http.Request, destination interface{}, limit int64) bool { + if limit <= 0 { + limit = 1 + } + r.Body = http.MaxBytesReader(w, r.Body, limit) + decoder := json.NewDecoder(r.Body) + if err := decoder.Decode(destination); err != nil { + var tooLarge *http.MaxBytesError + if errors.As(err, &tooLarge) { + jsonErrCode(w, http.StatusRequestEntityTooLarge, "request_too_large", "request body is too large") + return false + } + jsonErrCode(w, http.StatusBadRequest, "invalid_json", "invalid JSON request") + return false + } + var trailing interface{} + if err := decoder.Decode(&trailing); err != io.EOF { + if err == nil { + jsonErrCode(w, http.StatusBadRequest, "trailing_json", "request must contain one JSON value") + } else { + var tooLarge *http.MaxBytesError + if errors.As(err, &tooLarge) { + jsonErrCode(w, http.StatusRequestEntityTooLarge, "request_too_large", "request body is too large") + } else { + jsonErrCode(w, http.StatusBadRequest, "invalid_json", "invalid JSON request") + } + } + return false + } + return true +} + +func validateStringLength(name, value string, max int) error { + if len(value) > max { + return fmt.Errorf("%s is too long", name) + } + return nil +} + func sha256Hex(s string) string { h := sha256.Sum256([]byte(s)) return hex.EncodeToString(h[:]) diff --git a/internal/server/middleware.go b/internal/server/middleware.go index 96bdee1..4889107 100644 --- a/internal/server/middleware.go +++ b/internal/server/middleware.go @@ -56,7 +56,7 @@ func (s *Server) authenticateDevice(w http.ResponseWriter, r *http.Request) (aut Scan(&device.DeviceID, &userID, &vaultID, &revokedAt) if err != nil { err = s.db.QueryRow(`SELECT id, user_id, vault_id, revoked_at - FROM server_devices WHERE api_key=?`, key). + FROM server_devices WHERE api_key=? AND legacy_api_key=1`, key). Scan(&device.DeviceID, &userID, &vaultID, &revokedAt) } if err != nil { @@ -75,23 +75,25 @@ func (s *Server) authenticateDevice(w http.ResponseWriter, r *http.Request) (aut } if device.UserID != "" { var blocked int - s.db.QueryRow("SELECT blocked FROM server_users WHERE id=?", device.UserID).Scan(&blocked) + if err := s.db.QueryRow("SELECT blocked FROM server_users WHERE id=?", device.UserID).Scan(&blocked); err != nil { + jsonInternalError(w, err) + return authenticatedDevice{}, false + } if blocked != 0 { jsonErr(w, 403, "user blocked") return authenticatedDevice{}, false } } - s.db.Exec("UPDATE server_devices SET last_seen=? WHERE id=?", time.Now().UTC().Format(time.RFC3339), device.DeviceID) + if _, err := s.db.Exec("UPDATE server_devices SET last_seen=? WHERE id=?", time.Now().UTC().Format(time.RFC3339), device.DeviceID); err != nil { + jsonInternalError(w, err) + return authenticatedDevice{}, false + } return device, true } func (s *Server) requireAdmin(w http.ResponseWriter, r *http.Request) bool { - cookie, err := r.Cookie("session") - if err != nil || !s.tokens.Check(cookie.Value) { - http.Redirect(w, r, "/admin/login", http.StatusFound) - return false - } - return true + _, ok := s.requireSession(w, r, sessionScopeAdmin) + return ok } type PasswordError string diff --git a/internal/server/password_reset.go b/internal/server/password_reset.go index d9e76dd..e73ff90 100644 --- a/internal/server/password_reset.go +++ b/internal/server/password_reset.go @@ -22,7 +22,7 @@ func (s *Server) resetPasswordWithToken(token, newPassword string) (string, erro var userID, expiresAt string err = tx.QueryRow(`SELECT user_id, expires_at FROM server_email_tokens - WHERE token=? AND purpose='reset'`, token).Scan(&userID, &expiresAt) + WHERE token_hash=? AND purpose='reset'`, emailTokenHash(token)).Scan(&userID, &expiresAt) if errors.Is(err, sql.ErrNoRows) { return "", errResetTokenInvalid } @@ -35,7 +35,7 @@ func (s *Server) resetPasswordWithToken(token, newPassword string) (string, erro } deleted, err := tx.Exec(`DELETE FROM server_email_tokens - WHERE token=? AND purpose='reset' AND expires_at=?`, token, expiresAt) + WHERE token_hash=? AND purpose='reset' AND expires_at=?`, emailTokenHash(token), expiresAt) if err != nil { return "", err } @@ -62,6 +62,9 @@ func (s *Server) resetPasswordWithToken(token, newPassword string) (string, erro if updatedCount != 1 { return "", errResetTokenInvalid } + if _, err := tx.Exec("DELETE FROM server_email_tokens WHERE user_id=? AND purpose='reset'", userID); err != nil { + return "", err + } if err := tx.Commit(); err != nil { return "", err } diff --git a/internal/server/proxy.go b/internal/server/proxy.go new file mode 100644 index 0000000..1b92874 --- /dev/null +++ b/internal/server/proxy.go @@ -0,0 +1,34 @@ +package server + +import ( + "net" + "net/http" + "net/netip" + "strings" +) + +// clientIPFromPeer accepts forwarding headers only when the TCP peer is a +// configured trusted proxy. The first X-Forwarded-For address is the original +// client under the documented nginx/Caddy single-proxy configuration. +func (s *Server) clientIPFromPeer(peer, forwardedFor string) string { + peerAddr, err := netip.ParseAddr(strings.TrimSpace(peer)) + if err != nil || s == nil || s.cfg == nil || !s.cfg.isTrustedProxy(peerAddr) { + return peer + } + for _, value := range strings.Split(forwardedFor, ",") { + candidate := strings.TrimSpace(value) + if addr, err := netip.ParseAddr(candidate); err == nil { + return addr.String() + } + } + return peerAddr.String() +} + +func (s *Server) remotePeerIsTrusted(r *http.Request) bool { + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + host = r.RemoteAddr + } + addr, err := netip.ParseAddr(strings.Trim(host, "[]")) + return err == nil && s != nil && s.cfg != nil && s.cfg.isTrustedProxy(addr) +} diff --git a/internal/server/rate.go b/internal/server/rate.go new file mode 100644 index 0000000..c41b7f8 --- /dev/null +++ b/internal/server/rate.go @@ -0,0 +1,72 @@ +package server + +import ( + "math" + "net/http" + "strings" + "time" +) + +var ratePolicies = map[string]RatePolicy{ + "pair": {Limit: 5, Window: 15 * time.Minute}, + "device-register": {Limit: 5, Window: 15 * time.Minute}, + "auth-test": {Limit: 10, Window: 10 * time.Minute}, + "login": {Limit: 10, Window: 10 * time.Minute}, + "register": {Limit: 5, Window: time.Hour}, + "forgot": {Limit: 5, Window: time.Hour}, + "reset": {Limit: 8, Window: time.Hour}, + "admin-reset": {Limit: 8, Window: time.Hour}, +} + +// allowRate applies an IP limit and, where a login/account is supplied, an +// additional bounded account bucket. It never logs submitted credentials. +func (s *Server) allowRate(w http.ResponseWriter, r *http.Request, action, account string) bool { + policy, ok := ratePolicies[action] + if !ok { + return true + } + ip := s.clientIP(r) + s.limiter.Cleanup(2 * policy.Window) + keys := []string{action + ":ip:" + ip} + account = strings.ToLower(strings.TrimSpace(account)) + if account != "" { + // Keep attacker-controlled account strings out of the in-memory key and + // audit path while preserving an independent per-account bucket. + keys = append(keys, action+":account:"+sha256Hex(account)) + } + for _, key := range keys { + if allowed, retryAfter := s.limiter.Allow(key, policy); !allowed { + seconds := int(math.Ceil(retryAfter.Seconds())) + if seconds < 1 { + seconds = 1 + } + w.Header().Set("Retry-After", strconvItoa(seconds)) + s.auditLog("rate_limit_exceeded", "", "", ip, "rate limit: "+action) + jsonErr(w, http.StatusTooManyRequests, "too many attempts") + return false + } + } + return true +} + +func strconvItoa(value int) string { + if value == 0 { + return "0" + } + negative := value < 0 + if negative { + value = -value + } + var digits [20]byte + index := len(digits) + for value > 0 { + index-- + digits[index] = byte('0' + value%10) + value /= 10 + } + if negative { + index-- + digits[index] = '-' + } + return string(digits[index:]) +} diff --git a/internal/server/ratelimit.go b/internal/server/ratelimit.go new file mode 100644 index 0000000..c89d8ad --- /dev/null +++ b/internal/server/ratelimit.go @@ -0,0 +1,81 @@ +package server + +import ( + "sync" + "time" +) + +// RatePolicy is a fixed-window policy. Buckets expire naturally and cleanup +// prevents unauthenticated traffic from consuming unbounded memory. +type RatePolicy struct { + Limit int + Window time.Duration +} + +type rateBucket struct { + Count int + WindowStart time.Time + LastSeen time.Time +} + +type rateLimiter struct { + mu sync.Mutex + buckets map[string]rateBucket + now func() time.Time +} + +const maxRateLimitBuckets = 10000 + +func newRateLimiter(now func() time.Time) *rateLimiter { + if now == nil { + now = time.Now + } + return &rateLimiter{buckets: make(map[string]rateBucket), now: now} +} + +// Allow consumes one attempt and reports the remaining wait when limited. +func (l *rateLimiter) Allow(key string, policy RatePolicy) (bool, time.Duration) { + if policy.Limit <= 0 || policy.Window <= 0 { + return true, 0 + } + now := l.now().UTC() + l.mu.Lock() + defer l.mu.Unlock() + bucket := l.buckets[key] + if bucket.WindowStart.IsZero() && len(l.buckets) >= maxRateLimitBuckets { + // An attacker can vary IP/account keys faster than age cleanup runs. + // Evict the least recently seen bucket to keep memory strictly bounded. + var oldestKey string + var oldest time.Time + for candidate, existing := range l.buckets { + if oldestKey == "" || existing.LastSeen.Before(oldest) { + oldestKey, oldest = candidate, existing.LastSeen + } + } + if oldestKey != "" { + delete(l.buckets, oldestKey) + } + } + if bucket.WindowStart.IsZero() || !now.Before(bucket.WindowStart.Add(policy.Window)) { + bucket = rateBucket{WindowStart: now} + } + bucket.LastSeen = now + if bucket.Count >= policy.Limit { + l.buckets[key] = bucket + return false, bucket.WindowStart.Add(policy.Window).Sub(now) + } + bucket.Count++ + l.buckets[key] = bucket + return true, 0 +} + +func (l *rateLimiter) Cleanup(maxAge time.Duration) { + now := l.now().UTC() + l.mu.Lock() + defer l.mu.Unlock() + for key, bucket := range l.buckets { + if now.Sub(bucket.LastSeen) > maxAge { + delete(l.buckets, key) + } + } +} diff --git a/internal/server/retention.go b/internal/server/retention.go new file mode 100644 index 0000000..ea6d984 --- /dev/null +++ b/internal/server/retention.go @@ -0,0 +1,45 @@ +package server + +import ( + "os" + "path/filepath" + "strings" + "time" +) + +// CleanupRetention only removes independently expiring data. It intentionally +// never prunes server_ops or content-addressed blobs: the operation log remains +// required for a newly paired device until a future checkpoint protocol exists. +func (s *Server) CleanupRetention(now time.Time) error { + if err := s.cleanupExpiredSessions(); err != nil { + return err + } + if _, err := s.db.Exec("DELETE FROM server_email_tokens WHERE expires_at <= ?", now.UTC().Format(time.RFC3339)); err != nil { + return err + } + if _, err := s.db.Exec("DELETE FROM server_idempotency_keys WHERE created_at < ?", now.Add(-time.Duration(s.cfg.Retention.IdempotencyHours)*time.Hour).UTC().Format(time.RFC3339)); err != nil { + return err + } + if _, err := s.db.Exec("DELETE FROM server_audit_log WHERE created_at < ?", now.AddDate(0, 0, -s.cfg.Retention.AuditDays).UTC().Format(time.RFC3339)); err != nil { + return err + } + entries, err := os.ReadDir(s.blobsDir) + if err != nil && !os.IsNotExist(err) { + return err + } + for _, entry := range entries { + if !strings.HasPrefix(entry.Name(), ".upload-") { + continue + } + info, err := entry.Info() + if err != nil { + return err + } + if info.ModTime().Before(now.Add(-time.Duration(s.cfg.Retention.TempUploadHours) * time.Hour)) { + if err := os.Remove(filepath.Join(s.blobsDir, entry.Name())); err != nil && !os.IsNotExist(err) { + return err + } + } + } + return nil +} diff --git a/internal/server/routes.go b/internal/server/routes.go index 608da62..77a077d 100644 --- a/internal/server/routes.go +++ b/internal/server/routes.go @@ -2,6 +2,8 @@ package server func (s *Server) routes() { s.mux.HandleFunc("/api/v1/health", s.handleHealth) + s.mux.HandleFunc("/livez", s.handleLiveness) + s.mux.HandleFunc("/readyz", s.handleHealth) s.mux.HandleFunc("/api/v1/device/register", s.handleDeviceRegister) s.mux.HandleFunc("/api/v1/sync/push", s.handleSyncPush) s.mux.HandleFunc("/api/v1/sync/pull", s.handleSyncPull) @@ -23,6 +25,7 @@ func (s *Server) routes() { s.mux.HandleFunc("/reset", s.handleUserWebReset) s.mux.HandleFunc("/logout", s.handleUserWebLogout) s.mux.HandleFunc("/api/v1/user/devices", s.handleUserDevices) + s.mux.HandleFunc("/api/v1/user/devices/", s.handleUserWebDeviceAction) s.mux.HandleFunc("/admin/login", s.handleAdminLogin) s.mux.HandleFunc("/admin/dashboard", s.handleAdminDashboard) s.mux.HandleFunc("/admin/users", s.handleAdminUsers) diff --git a/internal/server/schema.go b/internal/server/schema.go index 8c3a6c1..a0c200e 100644 --- a/internal/server/schema.go +++ b/internal/server/schema.go @@ -24,6 +24,7 @@ CREATE TABLE IF NOT EXISTS server_devices ( token_hash TEXT, token_prefix TEXT, token_suffix TEXT, + legacy_api_key INTEGER NOT NULL DEFAULT 0, user_id TEXT, vault_id TEXT, client_version TEXT, @@ -76,13 +77,23 @@ CREATE TABLE IF NOT EXISTS server_idempotency_keys ( ); CREATE TABLE IF NOT EXISTS server_email_tokens ( - token TEXT PRIMARY KEY, + token_hash TEXT PRIMARY KEY, user_id TEXT NOT NULL, purpose TEXT NOT NULL, expires_at TEXT NOT NULL, created_at TEXT NOT NULL ); +CREATE TABLE IF NOT EXISTS server_sessions ( + token_hash TEXT PRIMARY KEY, + csrf_hash TEXT NOT NULL, + scope TEXT NOT NULL, + subject_id TEXT NOT NULL, + expires_at TEXT NOT NULL, + created_at TEXT NOT NULL, + last_seen TEXT NOT NULL +); + CREATE TABLE IF NOT EXISTS server_revisions ( id INTEGER PRIMARY KEY AUTOINCREMENT, op_id TEXT NOT NULL, @@ -95,6 +106,16 @@ CREATE TABLE IF NOT EXISTS server_blobs ( created_at TEXT NOT NULL ); +CREATE TABLE IF NOT EXISTS server_blob_refs ( + user_id TEXT NOT NULL, + vault_id TEXT NOT NULL, + sha256 TEXT NOT NULL, + size INTEGER NOT NULL, + created_at TEXT NOT NULL, + last_accessed TEXT NOT NULL, + PRIMARY KEY (user_id, vault_id, sha256) +); + CREATE TABLE IF NOT EXISTS server_audit_log ( id INTEGER PRIMARY KEY AUTOINCREMENT, event_type TEXT NOT NULL, @@ -114,6 +135,10 @@ CREATE INDEX IF NOT EXISTS idx_server_users_username ON server_users(username); CREATE INDEX IF NOT EXISTS idx_server_users_email ON server_users(email); CREATE INDEX IF NOT EXISTS idx_server_audit_log_event ON server_audit_log(event_type); CREATE INDEX IF NOT EXISTS idx_server_audit_log_created ON server_audit_log(created_at); +CREATE INDEX IF NOT EXISTS idx_server_blob_refs_scope ON server_blob_refs(user_id, vault_id, sha256); +CREATE INDEX IF NOT EXISTS idx_server_blob_refs_user ON server_blob_refs(user_id, sha256); +CREATE INDEX IF NOT EXISTS idx_server_sessions_expiry ON server_sessions(expires_at); +CREATE INDEX IF NOT EXISTS idx_server_sessions_subject ON server_sessions(scope, subject_id); CREATE TABLE IF NOT EXISTS server_smtp_config ( key TEXT PRIMARY KEY, @@ -125,6 +150,8 @@ type sqliteColumn struct { primaryKeyOrder int } +const schemaVersion = 2 + func migrateServerSchema(db *sql.DB) error { tx, err := db.Begin() if err != nil { @@ -150,6 +177,12 @@ func migrateServerSchema(db *sql.DB) error { if err := backfillDeviceOwners(tx); err != nil { return err } + if err := migrateLegacyDeviceCredentials(tx); err != nil { + return err + } + if err := migrateEmailTokenHashes(tx); err != nil { + return err + } if err := backfillOperationScope(tx); err != nil { return err } @@ -167,10 +200,67 @@ func migrateServerSchema(db *sql.DB) error { ON server_devices(user_id, vault_id)`); err != nil { return err } + if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", schemaVersion)); err != nil { + return err + } return tx.Commit() } +func migrateLegacyDeviceCredentials(tx *sql.Tx) error { + if err := ensureSQLiteColumn(tx, "server_devices", "legacy_api_key", "INTEGER NOT NULL DEFAULT 0"); err != nil { + return err + } + // Existing rows predate the hash-only enrollment path. Their API keys keep + // working only when explicitly marked legacy; devices enrolled by this + // version store a disabled placeholder and can never authenticate with it. + _, err := tx.Exec(`UPDATE server_devices SET legacy_api_key=1 + WHERE legacy_api_key=0 AND api_key NOT LIKE 'disabled:%'`) + return err +} + +func migrateEmailTokenHashes(tx *sql.Tx) error { + columns, err := sqliteTableColumns(tx, "server_email_tokens") + if err != nil { + return err + } + if _, ok := columns["token_hash"]; ok { + return nil + } + if _, err := tx.Exec("ALTER TABLE server_email_tokens RENAME TO server_email_tokens_legacy"); err != nil { + return err + } + if _, err := tx.Exec(`CREATE TABLE server_email_tokens ( + token_hash TEXT PRIMARY KEY, + user_id TEXT NOT NULL, + purpose TEXT NOT NULL, + expires_at TEXT NOT NULL, + created_at TEXT NOT NULL + )`); err != nil { + return err + } + rows, err := tx.Query(`SELECT token, user_id, purpose, expires_at, created_at FROM server_email_tokens_legacy`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var token, userID, purpose, expiresAt, createdAt string + if err := rows.Scan(&token, &userID, &purpose, &expiresAt, &createdAt); err != nil { + return err + } + if _, err := tx.Exec(`INSERT INTO server_email_tokens (token_hash, user_id, purpose, expires_at, created_at) + VALUES (?, ?, ?, ?, ?)`, sha256Hex(token), userID, purpose, expiresAt, createdAt); err != nil { + return err + } + } + if err := rows.Err(); err != nil { + return err + } + _, err = tx.Exec("DROP TABLE server_email_tokens_legacy") + return err +} + func ensureSQLiteColumn(tx *sql.Tx, table, column, definition string) error { columns, err := sqliteTableColumns(tx, table) if err != nil { diff --git a/internal/server/server.go b/internal/server/server.go index 2f84cf9..d7f3863 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1,133 +1,62 @@ package server import ( - "crypto/rand" "database/sql" - "encoding/hex" "fmt" + "net" "net/http" "os" "path/filepath" "strings" - "sync" "time" _ "github.com/mattn/go-sqlite3" ) -type pairRateLimit struct { - mu sync.Mutex - attempts map[string]int -} - -func (p *pairRateLimit) allow(ip string) bool { - p.mu.Lock() - defer p.mu.Unlock() - if p.attempts == nil { - p.attempts = make(map[string]int) - } - p.attempts[ip]++ - return p.attempts[ip] <= 5 -} - -func (p *pairRateLimit) reset(ip string) { - p.mu.Lock() - defer p.mu.Unlock() - delete(p.attempts, ip) -} - -type tokenStore struct { - mu sync.Mutex - tokens map[string]time.Time -} - -func newTokenStore() *tokenStore { - return &tokenStore{tokens: make(map[string]time.Time)} -} - -func (ts *tokenStore) Create() string { - ts.mu.Lock() - defer ts.mu.Unlock() - b := make([]byte, 16) - rand.Read(b) - tok := hex.EncodeToString(b) - ts.tokens[tok] = time.Now().Add(24 * time.Hour) - return tok -} - -func (ts *tokenStore) Check(tok string) bool { - ts.mu.Lock() - defer ts.mu.Unlock() - exp, ok := ts.tokens[tok] - if !ok { - return false - } - if time.Now().After(exp) { - delete(ts.tokens, tok) - return false - } - return true -} - -type userTokenStore struct { - mu sync.Mutex - tokens map[string]userTokenEntry -} - -type userTokenEntry struct { - UserID string - ExpiresAt time.Time -} - -func newUserTokenStore() *userTokenStore { - return &userTokenStore{tokens: make(map[string]userTokenEntry)} -} - -func (uts *userTokenStore) Create(userID string) string { - uts.mu.Lock() - defer uts.mu.Unlock() - b := make([]byte, 16) - rand.Read(b) - tok := hex.EncodeToString(b) - uts.tokens[tok] = userTokenEntry{UserID: userID, ExpiresAt: time.Now().Add(24 * time.Hour)} - return tok -} - -func (uts *userTokenStore) Check(tok string) (string, bool) { - uts.mu.Lock() - defer uts.mu.Unlock() - entry, ok := uts.tokens[tok] - if !ok { - return "", false - } - if time.Now().After(entry.ExpiresAt) { - delete(uts.tokens, tok) - return "", false - } - return entry.UserID, true -} - type Server struct { - db *sql.DB - cfg *Config - tokens *tokenStore - userTokens *userTokenStore - blobsDir string - mux *http.ServeMux - pairLimit *pairRateLimit + db *sql.DB + dbPath string + cfg *Config + blobsDir string + mux *http.ServeMux + limiter *rateLimiter + startedAt time.Time } +// Version and BuildCommit are assigned through -ldflags during release builds. +var ( + Version = "dev" + BuildCommit = "unknown" +) + func (s *Server) auditLog(eventType, userID, deviceID, ip, msg string) { s.db.Exec("INSERT INTO server_audit_log (event_type, user_id, device_id, ip, message, created_at) VALUES (?, ?, ?, ?, ?, ?)", eventType, userID, deviceID, ip, msg, time.Now().UTC().Format(time.RFC3339)) } func NewServer(dbPath, dataDir string, cfg *Config) (*Server, error) { + if cfg == nil { + cfg = DefaultConfig() + } + if err := cfg.normalize(); err != nil { + return nil, fmt.Errorf("config: %w", err) + } db, err := sql.Open("sqlite3", fmt.Sprintf("file:%s?mode=rwc", dbPath)) if err != nil { return nil, fmt.Errorf("open db: %w", err) } db.SetMaxOpenConns(1) + for _, pragma := range []string{ + "PRAGMA foreign_keys = ON", + "PRAGMA busy_timeout = 5000", + "PRAGMA journal_mode = WAL", + "PRAGMA synchronous = NORMAL", + } { + if _, err := db.Exec(pragma); err != nil { + db.Close() + return nil, fmt.Errorf("sqlite %s: %w", pragma, err) + } + } for _, stmt := range strings.Split(serverSchema, ";") { stmt = strings.TrimSpace(stmt) @@ -151,12 +80,12 @@ func NewServer(dbPath, dataDir string, cfg *Config) (*Server, error) { } s := &Server{ - db: db, - cfg: cfg, - tokens: newTokenStore(), - userTokens: newUserTokenStore(), - blobsDir: blobsDir, - pairLimit: &pairRateLimit{}, + db: db, + dbPath: dbPath, + cfg: cfg, + blobsDir: blobsDir, + limiter: newRateLimiter(nil), + startedAt: time.Now().UTC(), } s.mux = http.NewServeMux() return s, nil @@ -174,6 +103,36 @@ func (s *Server) Close() error { return s.db.Close() } -func (s *Server) ListenAndServe(addr string) error { - return http.ListenAndServe(addr, s.mux) +// Handler is the only HTTP entrypoint. Additional request security middleware +// is composed here so tests and production use the same path. +func (s *Server) Handler() http.Handler { + return s.mux +} + +// HTTPServer creates a conservatively configured server suitable for running +// behind nginx or Caddy. It intentionally does not enable TLS itself. +func (s *Server) HTTPServer(addr string) *http.Server { + return &http.Server{ + Addr: addr, + Handler: s.Handler(), + ReadHeaderTimeout: 10 * time.Second, + ReadTimeout: 30 * time.Second, + WriteTimeout: 60 * time.Second, + IdleTimeout: 120 * time.Second, + MaxHeaderBytes: 16 << 10, + } +} + +// ListenAndServe exists for callers that do not need to manage lifecycle. The +// command entrypoint uses HTTPServer plus graceful shutdown instead. +func (s *Server) ListenAndServe(addr string) error { + return s.HTTPServer(addr).ListenAndServe() +} + +func (s *Server) clientIP(r *http.Request) string { + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + host = r.RemoteAddr + } + return s.clientIPFromPeer(host, r.Header.Get("X-Forwarded-For")) } diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 3bb56a2..291bbef 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -111,7 +111,7 @@ func TestSyncPushPullStoresSequencedOps(t *testing.T) { insertSyncUser(t, s, "user-a") now := time.Now().UTC().Format(time.RFC3339) if _, err := s.db.Exec( - "INSERT INTO server_devices (id, name, api_key, user_id, vault_id, last_seen, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + "INSERT INTO server_devices (id, name, api_key, legacy_api_key, user_id, vault_id, last_seen, created_at) VALUES (?, ?, ?, 1, ?, ?, ?, ?)", "device-a", "Device A", "api-key", "user-a", "vault-a", now, now, ); err != nil { t.Fatalf("insert device: %v", err) @@ -185,7 +185,7 @@ func TestRevokedLegacyAPIKeyCannotPushOrPull(t *testing.T) { now := time.Now().UTC().Format(time.RFC3339) if _, err := s.db.Exec( - "INSERT INTO server_devices (id, name, api_key, last_seen, revoked_at, created_at) VALUES (?, ?, ?, ?, ?, ?)", + "INSERT INTO server_devices (id, name, api_key, legacy_api_key, last_seen, revoked_at, created_at) VALUES (?, ?, ?, 1, ?, ?, ?)", "device-revoked", "Revoked Device", "revoked-key", now, now, now, ); err != nil { t.Fatalf("insert device: %v", err) @@ -389,8 +389,8 @@ func TestWebResetRejectsExpiredToken(t *testing.T) { insertPairableUser(t, s, "user-a", "alice", oldPassword) now := time.Now().UTC().Format(time.RFC3339) if _, err := s.db.Exec(`INSERT INTO server_email_tokens - (token, user_id, purpose, expires_at, created_at) - VALUES (?, ?, 'reset', ?, ?)`, "expired-reset-token", "user-a", time.Now().Add(-time.Hour).UTC().Format(time.RFC3339), now); err != nil { + (token_hash, user_id, purpose, expires_at, created_at) + VALUES (?, ?, 'reset', ?, ?)`, emailTokenHash("expired-reset-token"), "user-a", time.Now().Add(-time.Hour).UTC().Format(time.RFC3339), now); err != nil { t.Fatalf("insert reset token: %v", err) } @@ -438,18 +438,18 @@ func TestServerRenderedPagesEscapeStoredValues(t *testing.T) { { name: "user dashboard", path: "/dashboard", - cookie: &http.Cookie{Name: "user_session", Value: s.userTokens.Create("user-a")}, + cookie: testSessionCookie(t, s, sessionScopeUser, "user-a"), containsDevID: true, }, { name: "admin users", path: "/admin/users", - cookie: &http.Cookie{Name: "admin_session", Value: s.tokens.Create()}, + cookie: testSessionCookie(t, s, sessionScopeAdmin, "admin"), }, { name: "admin devices", path: "/admin/devices", - cookie: &http.Cookie{Name: "admin_session", Value: s.tokens.Create()}, + cookie: testSessionCookie(t, s, sessionScopeAdmin, "admin"), containsDevID: true, }, } @@ -476,6 +476,19 @@ func TestServerRenderedPagesEscapeStoredValues(t *testing.T) { } } +func testSessionCookie(t *testing.T, s *Server, scope, subjectID string) *http.Cookie { + t.Helper() + token, _, err := s.createSession(scope, subjectID) + if err != nil { + t.Fatalf("create test session: %v", err) + } + name := "user_session" + if scope == sessionScopeAdmin { + name = "admin_session" + } + return &http.Cookie{Name: name, Value: token} +} + func TestNewServerMigratesLegacyOperationScope(t *testing.T) { dir := t.TempDir() dbPath := filepath.Join(dir, "legacy.db") diff --git a/internal/server/sessions.go b/internal/server/sessions.go new file mode 100644 index 0000000..b46253c --- /dev/null +++ b/internal/server/sessions.go @@ -0,0 +1,173 @@ +package server + +import ( + "crypto/rand" + "crypto/subtle" + "database/sql" + "encoding/hex" + "fmt" + "net/http" + "strings" + "time" +) + +const ( + sessionScopeAdmin = "admin" + sessionScopeUser = "user" + sessionLifetime = 24 * time.Hour +) + +type webSession struct { + Scope string + SubjectID string + CSRFHash string + ExpiresAt time.Time +} + +func randomSecret(bytes int) (string, error) { + b := make([]byte, bytes) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} + +// createSession stores only hashes. The plaintext session and CSRF values are +// returned once to be placed in cookies or an API login response. +func (s *Server) createSession(scope, subjectID string) (token, csrf string, err error) { + if scope != sessionScopeAdmin && scope != sessionScopeUser { + return "", "", fmt.Errorf("unknown session scope") + } + token, err = randomSecret(32) + if err != nil { + return "", "", err + } + csrf, err = randomSecret(32) + if err != nil { + return "", "", err + } + now := time.Now().UTC() + _, err = s.db.Exec(`INSERT INTO server_sessions + (token_hash, csrf_hash, scope, subject_id, expires_at, created_at, last_seen) + VALUES (?, ?, ?, ?, ?, ?, ?)`, + sha256Hex(token), sha256Hex(csrf), scope, subjectID, + now.Add(sessionLifetime).Format(time.RFC3339), now.Format(time.RFC3339), now.Format(time.RFC3339)) + if err != nil { + return "", "", err + } + return token, csrf, nil +} + +func (s *Server) loadSession(token, scope string) (webSession, bool) { + if token == "" { + return webSession{}, false + } + var session webSession + var storedTokenHash string + var expiresAt string + err := s.db.QueryRow(`SELECT token_hash, csrf_hash, scope, subject_id, expires_at + FROM server_sessions WHERE token_hash=? AND scope=?`, sha256Hex(token), scope). + Scan(&storedTokenHash, &session.CSRFHash, &session.Scope, &session.SubjectID, &expiresAt) + if err != nil { + return webSession{}, false + } + expires, err := time.Parse(time.RFC3339, expiresAt) + if err != nil || !time.Now().UTC().Before(expires) { + _, _ = s.db.Exec("DELETE FROM server_sessions WHERE token_hash=?", sha256Hex(token)) + return webSession{}, false + } + session.ExpiresAt = expires + // The database lookup is indexed by a fixed-size hash. Keep a constant-time + // comparison at the final token-hash boundary as well. + if subtle.ConstantTimeCompare([]byte(storedTokenHash), []byte(sha256Hex(token))) != 1 { + return webSession{}, false + } + _, _ = s.db.Exec("UPDATE server_sessions SET last_seen=? WHERE token_hash=?", time.Now().UTC().Format(time.RFC3339), sha256Hex(token)) + return session, true +} + +func (s *Server) deleteSession(token string) error { + if token == "" { + return nil + } + _, err := s.db.Exec("DELETE FROM server_sessions WHERE token_hash=?", sha256Hex(token)) + return err +} + +func (s *Server) deleteSessionsForSubject(scope, subjectID string) error { + _, err := s.db.Exec("DELETE FROM server_sessions WHERE scope=? AND subject_id=?", scope, subjectID) + return err +} + +func (s *Server) setSessionCookies(w http.ResponseWriter, r *http.Request, scope, token, csrf string) { + name := "user_session" + path := "/" + if scope == sessionScopeAdmin { + name = "admin_session" + path = "/admin" + } + secure := s.requestIsHTTPS(r) + http.SetCookie(w, &http.Cookie{ + Name: name, Value: token, Path: path, HttpOnly: true, Secure: secure, + SameSite: http.SameSiteLaxMode, MaxAge: int(sessionLifetime.Seconds()), + }) + http.SetCookie(w, &http.Cookie{ + Name: "csrf_token", Value: csrf, Path: path, HttpOnly: false, Secure: secure, + SameSite: http.SameSiteStrictMode, MaxAge: int(sessionLifetime.Seconds()), + }) +} + +func (s *Server) clearSessionCookies(w http.ResponseWriter, r *http.Request, scope string) { + name := "user_session" + path := "/" + if scope == sessionScopeAdmin { + name = "admin_session" + path = "/admin" + } + secure := s.requestIsHTTPS(r) + for _, cookieName := range []string{name, "csrf_token"} { + http.SetCookie(w, &http.Cookie{Name: cookieName, Value: "", Path: path, HttpOnly: cookieName != "csrf_token", Secure: secure, MaxAge: -1}) + } +} + +func (s *Server) requestIsHTTPS(r *http.Request) bool { + if r.TLS != nil { + return true + } + if s == nil || s.cfg == nil { + return false + } + if !s.remotePeerIsTrusted(r) { + return false + } + return strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https") +} + +func (s *Server) requireSession(w http.ResponseWriter, r *http.Request, scope string) (webSession, bool) { + name := "user_session" + login := "/login" + if scope == sessionScopeAdmin { + name = "admin_session" + login = "/admin/login" + } + cookie, err := r.Cookie(name) + if err != nil { + http.Redirect(w, r, login, http.StatusFound) + return webSession{}, false + } + session, ok := s.loadSession(cookie.Value, scope) + if !ok { + http.Redirect(w, r, login, http.StatusFound) + return webSession{}, false + } + return session, true +} + +func (s *Server) cleanupExpiredSessions() error { + _, err := s.db.Exec("DELETE FROM server_sessions WHERE expires_at <= ?", time.Now().UTC().Format(time.RFC3339)) + return err +} + +func isNoRows(err error) bool { + return err == sql.ErrNoRows +} diff --git a/internal/server/sync_contract.go b/internal/server/sync_contract.go new file mode 100644 index 0000000..e359f4a --- /dev/null +++ b/internal/server/sync_contract.go @@ -0,0 +1,95 @@ +package server + +import ( + "encoding/json" + "fmt" + "strings" +) + +const ( + maxOpIDLength = 128 + maxEntityTypeLength = 64 + maxEntityIDLength = 4096 + maxOpTypeLength = 64 + maxIdempotencyKeyLength = 128 + maxDeviceIDLength = 128 + maxDeviceNameLength = 128 + maxClientVersionLength = 256 + maxVaultIDLength = 256 + maxLoginLength = 320 +) + +type syncPushOperation struct { + OpID string `json:"op_id"` + EntityType string `json:"entity_type"` + EntityID string `json:"entity_id"` + OpType string `json:"op_type"` + PayloadJSON string `json:"payload_json"` + ClientSequence int `json:"client_sequence"` + LastSeenServerSeq int `json:"last_seen_server_seq"` + CreatedAt string `json:"created_at"` +} + +type syncPushRequest struct { + DeviceID string `json:"device_id"` + IdempotencyKey string `json:"idempotency_key"` + Ops []syncPushOperation `json:"ops"` +} + +type syncPullRequest struct { + SinceSequence int `json:"since_sequence"` + PageLimit int `json:"page_limit"` +} + +func (s *Server) validateSyncPush(req syncPushRequest) (code, message string) { + if len(req.Ops) > s.cfg.Limits.MaxPushOperations { + return "too_many_operations", "too many operations in one push" + } + for name, value := range map[string]struct { + value string + max int + }{ + "device_id": {req.DeviceID, maxDeviceIDLength}, + "idempotency_key": {req.IdempotencyKey, maxIdempotencyKeyLength}, + } { + if err := validateStringLength(name, value.value, value.max); err != nil { + return "field_too_long", err.Error() + } + } + for _, op := range req.Ops { + for _, value := range []struct { + name string + value string + max int + }{ + {"op_id", op.OpID, maxOpIDLength}, + {"entity_type", op.EntityType, maxEntityTypeLength}, + {"entity_id", op.EntityID, maxEntityIDLength}, + {"op_type", op.OpType, maxOpTypeLength}, + } { + if strings.TrimSpace(value.value) == "" { + return "invalid_operation", fmt.Sprintf("%s is required", value.name) + } + if err := validateStringLength(value.name, value.value, value.max); err != nil { + return "field_too_long", err.Error() + } + } + if len(op.PayloadJSON) > s.cfg.Limits.MaxPayloadJSON { + return "payload_too_large", "operation payload is too large" + } + if op.PayloadJSON != "" && !json.Valid([]byte(op.PayloadJSON)) { + return "invalid_payload", "operation payload_json must be valid JSON" + } + if op.ClientSequence < 0 || op.LastSeenServerSeq < 0 { + return "invalid_operation", "operation sequences must be non-negative" + } + } + return "", "" +} + +func (s *Server) pullPageLimit(requested int) int { + if requested <= 0 || requested > s.cfg.Limits.MaxPullPage { + return s.cfg.Limits.MaxPullPage + } + return requested +} diff --git a/internal/server/templates.go b/internal/server/templates.go index e59bc37..2f2f3d5 100644 --- a/internal/server/templates.go +++ b/internal/server/templates.go @@ -454,7 +454,7 @@ function testSMTP(){ ) } -func userDashboardHTML(locale, username, deviceRows string) string { +func userDashboardHTML(locale, username, deviceRows, csrf string) string { return fmt.Sprintf(` @@ -479,7 +479,7 @@ a{color:#6366f1}| %[4]s | %[5]s | %[6]s | %[7]s | %[8]s |
|---|