mirror of
https://github.com/sipeed/NanoKVM.git
synced 2026-09-11 00:22:56 -05:00
feat: add secure multi-user support (#876)
This commit is contained in:
@@ -48,7 +48,7 @@ authentication: enable # Whether to enable identity verification fo
|
||||
jwt:
|
||||
secretKey: "" # The secret key used to sign and verify JWT Tokens. If left empty, a random key will be generated automatically on startup
|
||||
refreshTokenDuration: 2678400 # The token refresh duration threshold in seconds before forcing a re-login. Default is `2678400` (~31 days)
|
||||
revokeTokensOnLogout: true # Whether to invalidate all existing tokens upon logout by rotating the SecretKey. Default is `true`
|
||||
revokeTokensOnLogout: true # Whether logout invalidates all sessions belonging to that user. Other users are never logged out. Setting this to false only clears the browser cookie and is not recommended. Default is `true`
|
||||
security:
|
||||
loginLockoutDuration: 0, # The duration (in seconds) to ban an IP from attempting to log in again after reaching the failure limit. If set to `0` or left empty, brute-force protection is disabled. Default is `0`
|
||||
loginMaxFailures: 5, # The maximum number of continuous failed login attempts allowed per IP before triggering protection. Default is `5`
|
||||
@@ -62,6 +62,28 @@ turn:
|
||||
turnCred: example_cred # The credential/password required for authorization to the TURN server
|
||||
```
|
||||
|
||||
## Web Users
|
||||
|
||||
NanoKVM uses two device-wide roles:
|
||||
|
||||
- `admin`: KVM access plus user, system, network, update, storage, terminal, script, MCP, and PicoClaw administration.
|
||||
- `user`: KVM video, keyboard, mouse, paste, power/reset, and Wake-on-LAN access.
|
||||
|
||||
Administrators manage accounts from **Settings > Account**. Account data remains in
|
||||
`/etc/kvm/pwd`; the server migrates the legacy single-account JSON format in place and writes
|
||||
the multi-user format atomically with mode `0600`. Keeping the same path preserves the physical
|
||||
BOOT-button password reset behavior.
|
||||
Users must confirm their current password when changing it themselves; administrators can reset
|
||||
non-owner users from the account manager. Only the device owner can change the device owner's
|
||||
password, because that password is also synchronized to the Linux root account.
|
||||
|
||||
All authenticated sessions are backed by the current account state. Disabling, deleting,
|
||||
changing the role or password of a user invalidates that user's HTTP and real-time connections
|
||||
without affecting other users. Multiple users may watch and control the KVM concurrently; input
|
||||
uses the existing cooperative HID coordinator, so simultaneous input can interleave.
|
||||
Video mode, quality, resolution, and MJPEG frame-detection controls remain shared KVM
|
||||
operations; when several users adjust them concurrently, the latest change applies device-wide.
|
||||
|
||||
## Compile & Deploy
|
||||
|
||||
Note: The manual steps below require a Linux x86-64 host with Go 1.25 or newer; they are not compatible with ARM, Windows or macOS. With Docker you can skip them entirely and use the containerized flow instead — the root [Makefile](../Makefile) (`make shell`) or the dev container (see "Development" in the root [README](../README.md)) — which works on any host OS; run `server/build.sh` inside the container for a release-equivalent build.
|
||||
|
||||
@@ -46,7 +46,7 @@ authentication: enable # 是否开启 HTTP 接口与网页的身份
|
||||
jwt:
|
||||
secretKey: "" # 用于签发和验证 JWT Token 的密钥。如果不填,服务启动时将自动随机生成
|
||||
refreshTokenDuration: 2678400 # 登录超时的刷新周期(单位:秒)。默认为 `2678400`(约31天)
|
||||
revokeTokensOnLogout: true # 退出登录时是否废除所有现存的 Token。启用此项可以在注销时轮换 SecretKey,强迫所有终端重新登录。默认为 `true`
|
||||
revokeTokensOnLogout: true # 退出登录时是否废除该用户的全部会话;不会影响其他用户。设为 false 时仅清除浏览器 Cookie,不推荐使用。默认为 `true`
|
||||
security:
|
||||
loginLockoutDuration: 0, # 达到失败上限后,禁止该 IP 再次尝试登录的持续时间(单位:秒)。如果设为 `0` 或不填,则代表不开启防暴力破解功能。默认为 `0`
|
||||
loginMaxFailures: 5, # 允许触发保护前,单个 IP 连续登录失败的最大次数。默认为 `5`
|
||||
@@ -60,6 +60,22 @@ turn:
|
||||
turnCred: example_cred # TURN 服务器授权连接时使用的凭据/密码
|
||||
```
|
||||
|
||||
## Web 多用户
|
||||
|
||||
NanoKVM 使用两种设备级角色:
|
||||
|
||||
- `admin`:除 KVM 操作外,还可管理用户、系统、网络、更新、存储、终端、脚本、MCP 和 PicoClaw。
|
||||
- `user`:可使用 KVM 视频、键盘、鼠标、粘贴、目标机电源/复位和网络唤醒。
|
||||
|
||||
管理员可在**设置 > 账户**中管理用户。账户数据继续保存在 `/etc/kvm/pwd`;服务会原地迁移旧版单账户
|
||||
JSON,并以 `0600` 权限原子写入多用户格式。沿用同一路径可保持长按 BOOT 键重置密码的现有行为。
|
||||
用户自助修改密码时必须验证当前密码;管理员可在账户管理中重置非设备所有者的密码。由于设备所有者
|
||||
密码还会同步至 Linux root 账户,因此只有设备所有者本人可以修改自己的密码。
|
||||
|
||||
所有登录会话都会校验当前账户状态。禁用、删除、修改角色或密码会立即撤销该用户的 HTTP 与实时连接,
|
||||
且不会影响其他用户。多个用户可同时观看和协作控制 KVM;输入沿用现有 HID 协调器,因此同时输入可能交错。
|
||||
视频模式、画质、分辨率和 MJPEG 帧检测仍属于共享 KVM 操作;多人同时调整时,以最后一次设备级修改为准。
|
||||
|
||||
## 编译部署
|
||||
|
||||
**注意:请使用 Linux 操作系统(x86-64)和 Go 1.25 或更高版本。该工具链无法在 ARM、Windows 或 macOS 下使用。**
|
||||
|
||||
616
server/authn/store.go
Normal file
616
server/authn/store.go
Normal file
@@ -0,0 +1,616 @@
|
||||
package authn
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sync"
|
||||
|
||||
"NanoKVM-Server/utils"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
const (
|
||||
AccountFile = "/etc/kvm/pwd"
|
||||
currentFileVersion = 1
|
||||
defaultUsername = "admin"
|
||||
defaultPassword = "admin"
|
||||
)
|
||||
|
||||
type Role string
|
||||
|
||||
const (
|
||||
RoleAdmin Role = "admin"
|
||||
RoleUser Role = "user"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrUserNotFound = errors.New("user not found")
|
||||
ErrUserExists = errors.New("username already exists")
|
||||
ErrLastAdmin = errors.New("at least one enabled admin is required")
|
||||
ErrSelfModification = errors.New("administrators cannot disable or demote themselves")
|
||||
ErrSelfDelete = errors.New("administrators cannot delete themselves")
|
||||
ErrSystemAccount = errors.New("the device owner account must remain an enabled administrator")
|
||||
|
||||
usernamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_.-]{0,31}$`)
|
||||
)
|
||||
|
||||
type User struct {
|
||||
Username string `json:"username"`
|
||||
PasswordHash string `json:"password"`
|
||||
Role Role `json:"role"`
|
||||
Enabled bool `json:"enabled"`
|
||||
TokenVersion uint64 `json:"tokenVersion"`
|
||||
MustChangePassword bool `json:"mustChangePassword,omitempty"`
|
||||
SystemAccount bool `json:"systemAccount,omitempty"`
|
||||
}
|
||||
|
||||
type UserInfo struct {
|
||||
Username string `json:"username"`
|
||||
Role Role `json:"role"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SystemAccount bool `json:"systemAccount,omitempty"`
|
||||
}
|
||||
|
||||
type UserPatch struct {
|
||||
Role *Role
|
||||
Enabled *bool
|
||||
}
|
||||
|
||||
type database struct {
|
||||
Version int `json:"version"`
|
||||
Users []User `json:"users"`
|
||||
LegacyUsername string `json:"username,omitempty"`
|
||||
LegacyPassword string `json:"password,omitempty"`
|
||||
}
|
||||
|
||||
type legacyAccount struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type Store struct {
|
||||
path string
|
||||
mutex sync.RWMutex
|
||||
}
|
||||
|
||||
var DefaultStore = NewStore(AccountFile)
|
||||
|
||||
func NewStore(path string) *Store {
|
||||
return &Store{path: path}
|
||||
}
|
||||
|
||||
func IsValidRole(role Role) bool {
|
||||
return role == RoleAdmin || role == RoleUser
|
||||
}
|
||||
|
||||
func ValidateUsername(username string) error {
|
||||
if !usernamePattern.MatchString(username) {
|
||||
return errors.New("username must be 1-32 characters and contain only letters, numbers, '.', '_' or '-'")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ValidatePassword(password string) error {
|
||||
length := len([]byte(password))
|
||||
if length < 8 || length > 72 {
|
||||
return errors.New("password must be between 8 and 72 bytes")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) List() ([]UserInfo, error) {
|
||||
s.mutex.RLock()
|
||||
defer s.mutex.RUnlock()
|
||||
|
||||
db, err := s.loadLocked(false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
users := make([]UserInfo, 0, len(db.Users))
|
||||
for _, user := range db.Users {
|
||||
users = append(users, UserInfo{
|
||||
Username: user.Username,
|
||||
Role: user.Role,
|
||||
Enabled: user.Enabled,
|
||||
SystemAccount: user.SystemAccount,
|
||||
})
|
||||
}
|
||||
return users, nil
|
||||
}
|
||||
|
||||
func (s *Store) Get(username string) (*User, error) {
|
||||
s.mutex.RLock()
|
||||
defer s.mutex.RUnlock()
|
||||
|
||||
db, err := s.loadLocked(false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return findUser(db.Users, username)
|
||||
}
|
||||
|
||||
func (s *Store) Authenticate(username, password string) (*User, bool, error) {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
db, err := s.loadLocked(true)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
user, err := findUser(db.Users, username)
|
||||
if err != nil || !user.Enabled {
|
||||
return nil, false, nil
|
||||
}
|
||||
|
||||
if bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)) == nil {
|
||||
return user, true, nil
|
||||
}
|
||||
|
||||
// Older releases stored a reversibly encrypted password. Upgrade it after
|
||||
// the first successful login so the compatibility path is not permanent.
|
||||
legacyPassword, decodeErr := utils.DecodeDecrypt(user.PasswordHash)
|
||||
if decodeErr != nil || legacyPassword != password {
|
||||
return nil, false, nil
|
||||
}
|
||||
|
||||
hash, hashErr := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if hashErr != nil {
|
||||
return nil, false, hashErr
|
||||
}
|
||||
for index := range db.Users {
|
||||
if db.Users[index].Username == username {
|
||||
db.Users[index].PasswordHash = string(hash)
|
||||
if db.Users[index].TokenVersion == 0 {
|
||||
db.Users[index].TokenVersion = 1
|
||||
}
|
||||
user = cloneUser(&db.Users[index])
|
||||
break
|
||||
}
|
||||
}
|
||||
if err = s.saveLocked(db); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return user, true, nil
|
||||
}
|
||||
|
||||
func (s *Store) ValidateToken(username string, tokenVersion uint64) (*User, error) {
|
||||
user, err := s.Get(username)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !user.Enabled || user.TokenVersion != tokenVersion {
|
||||
return nil, errors.New("session revoked")
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (s *Store) Create(username, password string, role Role) error {
|
||||
if err := ValidateUsername(username); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ValidatePassword(password); err != nil {
|
||||
return err
|
||||
}
|
||||
if !IsValidRole(role) {
|
||||
return errors.New("invalid role")
|
||||
}
|
||||
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tokenVersion, err := newTokenVersion()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
db, err := s.loadLocked(true)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err = findUser(db.Users, username); err == nil {
|
||||
return ErrUserExists
|
||||
}
|
||||
db.Users = append(db.Users, User{
|
||||
Username: username,
|
||||
PasswordHash: string(hash),
|
||||
Role: role,
|
||||
Enabled: true,
|
||||
TokenVersion: tokenVersion,
|
||||
})
|
||||
return s.saveLocked(db)
|
||||
}
|
||||
|
||||
func (s *Store) Update(actor, username string, patch UserPatch) (*User, error) {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
db, err := s.loadLocked(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
index := userIndex(db.Users, username)
|
||||
if index < 0 {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
user := db.Users[index]
|
||||
|
||||
if patch.Role != nil && !IsValidRole(*patch.Role) {
|
||||
return nil, errors.New("invalid role")
|
||||
}
|
||||
if actor == username && user.Role == RoleAdmin &&
|
||||
((patch.Role != nil && *patch.Role != RoleAdmin) || (patch.Enabled != nil && !*patch.Enabled)) {
|
||||
return nil, ErrSelfModification
|
||||
}
|
||||
if user.SystemAccount &&
|
||||
((patch.Role != nil && *patch.Role != RoleAdmin) || (patch.Enabled != nil && !*patch.Enabled)) {
|
||||
return nil, ErrSystemAccount
|
||||
}
|
||||
|
||||
changed := false
|
||||
if patch.Role != nil && user.Role != *patch.Role {
|
||||
user.Role = *patch.Role
|
||||
changed = true
|
||||
}
|
||||
if patch.Enabled != nil && user.Enabled != *patch.Enabled {
|
||||
user.Enabled = *patch.Enabled
|
||||
changed = true
|
||||
}
|
||||
if !changed {
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
db.Users[index] = user
|
||||
if enabledAdminCount(db.Users) == 0 {
|
||||
return nil, ErrLastAdmin
|
||||
}
|
||||
db.Users[index].TokenVersion++
|
||||
if err = s.saveLocked(db); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cloneUser(&db.Users[index]), nil
|
||||
}
|
||||
|
||||
func (s *Store) Delete(actor, username string) error {
|
||||
if actor == username {
|
||||
return ErrSelfDelete
|
||||
}
|
||||
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
db, err := s.loadLocked(true)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
index := userIndex(db.Users, username)
|
||||
if index < 0 {
|
||||
return ErrUserNotFound
|
||||
}
|
||||
if db.Users[index].SystemAccount {
|
||||
return ErrSystemAccount
|
||||
}
|
||||
users := append([]User(nil), db.Users[:index]...)
|
||||
users = append(users, db.Users[index+1:]...)
|
||||
if enabledAdminCount(users) == 0 {
|
||||
return ErrLastAdmin
|
||||
}
|
||||
db.Users = users
|
||||
return s.saveLocked(db)
|
||||
}
|
||||
|
||||
func (s *Store) SetPassword(username, password string) (*User, error) {
|
||||
return s.SetPasswordAndRun(username, password, nil)
|
||||
}
|
||||
|
||||
// SetPasswordAndRun commits the web password before running the optional
|
||||
// system-side update. If that update fails, the account record is rolled back
|
||||
// while the store lock is still held.
|
||||
func (s *Store) SetPasswordAndRun(username, password string, afterCommit func() error) (*User, error) {
|
||||
if err := ValidatePassword(password); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
db, err := s.loadLocked(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
index := userIndex(db.Users, username)
|
||||
if index < 0 {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
previous := *db
|
||||
previous.Users = append([]User(nil), db.Users...)
|
||||
db.Users[index].PasswordHash = string(hash)
|
||||
db.Users[index].MustChangePassword = false
|
||||
db.Users[index].TokenVersion++
|
||||
if err = s.saveLocked(db); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if afterCommit != nil {
|
||||
if err = afterCommit(); err != nil {
|
||||
if rollbackErr := s.saveLocked(&previous); rollbackErr != nil {
|
||||
return nil, fmt.Errorf("system password update failed: %v; account rollback failed: %w", err, rollbackErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return cloneUser(&db.Users[index]), nil
|
||||
}
|
||||
|
||||
func (s *Store) Revoke(username string) (*User, error) {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
db, err := s.loadLocked(true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
index := userIndex(db.Users, username)
|
||||
if index < 0 {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
db.Users[index].TokenVersion++
|
||||
if err = s.saveLocked(db); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cloneUser(&db.Users[index]), nil
|
||||
}
|
||||
|
||||
func (s *Store) loadLocked(migrate bool) (*database, error) {
|
||||
data, err := os.ReadFile(s.path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
db, defaultErr := defaultDatabase()
|
||||
if defaultErr != nil {
|
||||
return nil, defaultErr
|
||||
}
|
||||
if migrate {
|
||||
if saveErr := s.saveLocked(db); saveErr != nil {
|
||||
return nil, saveErr
|
||||
}
|
||||
}
|
||||
return db, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var db database
|
||||
if err = json.Unmarshal(data, &db); err == nil && db.Version != 0 {
|
||||
if err = validateDatabase(&db); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &db, nil
|
||||
}
|
||||
|
||||
var legacy legacyAccount
|
||||
if err = json.Unmarshal(data, &legacy); err != nil || legacy.Username == "" || legacy.Password == "" {
|
||||
return nil, errors.New("invalid account file")
|
||||
}
|
||||
if err = ValidateUsername(legacy.Username); err != nil {
|
||||
return nil, fmt.Errorf("invalid legacy username: %w", err)
|
||||
}
|
||||
tokenVersion, err := newTokenVersion()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db = database{
|
||||
Version: currentFileVersion,
|
||||
Users: []User{{
|
||||
Username: legacy.Username,
|
||||
PasswordHash: legacy.Password,
|
||||
Role: RoleAdmin,
|
||||
Enabled: true,
|
||||
TokenVersion: tokenVersion,
|
||||
MustChangePassword: passwordMatches(legacy.Password, defaultPassword),
|
||||
SystemAccount: true,
|
||||
}},
|
||||
}
|
||||
if migrate {
|
||||
if err = s.saveLocked(&db); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return &db, nil
|
||||
}
|
||||
|
||||
func (s *Store) saveLocked(db *database) error {
|
||||
if err := validateDatabase(db); err != nil {
|
||||
return err
|
||||
}
|
||||
syncLegacyAccount(db)
|
||||
data, err := json.MarshalIndent(db, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data = append(data, '\n')
|
||||
|
||||
directory := filepath.Dir(s.path)
|
||||
if err = os.MkdirAll(directory, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
temporary, err := os.CreateTemp(directory, ".pwd-*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
temporaryPath := temporary.Name()
|
||||
defer func() { _ = os.Remove(temporaryPath) }()
|
||||
|
||||
if err = temporary.Chmod(0o600); err == nil {
|
||||
_, err = temporary.Write(data)
|
||||
}
|
||||
if err == nil {
|
||||
err = temporary.Sync()
|
||||
}
|
||||
closeErr := temporary.Close()
|
||||
if err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err = os.Rename(temporaryPath, s.path); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
dir, err := os.Open(directory)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = dir.Close() }()
|
||||
return dir.Sync()
|
||||
}
|
||||
|
||||
func defaultDatabase() (*database, error) {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(defaultPassword), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tokenVersion, err := newTokenVersion()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &database{
|
||||
Version: currentFileVersion,
|
||||
Users: []User{{
|
||||
Username: defaultUsername,
|
||||
PasswordHash: string(hash),
|
||||
Role: RoleAdmin,
|
||||
Enabled: true,
|
||||
TokenVersion: tokenVersion,
|
||||
MustChangePassword: true,
|
||||
SystemAccount: true,
|
||||
}},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newTokenVersion() (uint64, error) {
|
||||
var data [8]byte
|
||||
for {
|
||||
if _, err := rand.Read(data[:]); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
version := binary.LittleEndian.Uint64(data[:])
|
||||
if version != 0 && version != ^uint64(0) {
|
||||
return version, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func validateDatabase(db *database) error {
|
||||
if db.Version != currentFileVersion {
|
||||
return fmt.Errorf("unsupported account file version: %d", db.Version)
|
||||
}
|
||||
if len(db.Users) == 0 {
|
||||
return errors.New("account file contains no users")
|
||||
}
|
||||
seen := make(map[string]struct{}, len(db.Users))
|
||||
systemAccounts := 0
|
||||
for index := range db.Users {
|
||||
user := &db.Users[index]
|
||||
if err := ValidateUsername(user.Username); err != nil {
|
||||
return err
|
||||
}
|
||||
if user.PasswordHash == "" || !IsValidRole(user.Role) {
|
||||
return errors.New("account file contains an invalid user")
|
||||
}
|
||||
if _, exists := seen[user.Username]; exists {
|
||||
return errors.New("account file contains duplicate usernames")
|
||||
}
|
||||
seen[user.Username] = struct{}{}
|
||||
if user.TokenVersion == 0 {
|
||||
user.TokenVersion = 1
|
||||
}
|
||||
if user.SystemAccount {
|
||||
systemAccounts++
|
||||
if user.Role != RoleAdmin || !user.Enabled || systemAccounts > 1 {
|
||||
return ErrSystemAccount
|
||||
}
|
||||
}
|
||||
}
|
||||
if enabledAdminCount(db.Users) == 0 {
|
||||
return ErrLastAdmin
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func findUser(users []User, username string) (*User, error) {
|
||||
index := userIndex(users, username)
|
||||
if index < 0 {
|
||||
return nil, ErrUserNotFound
|
||||
}
|
||||
return cloneUser(&users[index]), nil
|
||||
}
|
||||
|
||||
func userIndex(users []User, username string) int {
|
||||
for index := range users {
|
||||
if users[index].Username == username {
|
||||
return index
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func enabledAdminCount(users []User) int {
|
||||
count := 0
|
||||
for _, user := range users {
|
||||
if user.Role == RoleAdmin && user.Enabled {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func cloneUser(user *User) *User {
|
||||
if user == nil {
|
||||
return nil
|
||||
}
|
||||
copy := *user
|
||||
return ©
|
||||
}
|
||||
|
||||
func passwordMatches(hash, password string) bool {
|
||||
if bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil {
|
||||
return true
|
||||
}
|
||||
legacyPassword, err := utils.DecodeDecrypt(hash)
|
||||
return err == nil && legacyPassword == password
|
||||
}
|
||||
|
||||
// Keep a legacy top-level account mirror so downgrading to a single-user
|
||||
// server still authenticates the original device owner instead of failing
|
||||
// open or forcing an immediate physical reset.
|
||||
func syncLegacyAccount(db *database) {
|
||||
db.LegacyUsername = ""
|
||||
db.LegacyPassword = ""
|
||||
var fallback *User
|
||||
for _, user := range db.Users {
|
||||
if user.SystemAccount {
|
||||
db.LegacyUsername = user.Username
|
||||
db.LegacyPassword = user.PasswordHash
|
||||
return
|
||||
}
|
||||
if fallback == nil && user.Role == RoleAdmin && user.Enabled {
|
||||
copy := user
|
||||
fallback = ©
|
||||
}
|
||||
}
|
||||
if fallback != nil {
|
||||
db.LegacyUsername = fallback.Username
|
||||
db.LegacyPassword = fallback.PasswordHash
|
||||
}
|
||||
}
|
||||
212
server/authn/store_test.go
Normal file
212
server/authn/store_test.go
Normal file
@@ -0,0 +1,212 @@
|
||||
package authn
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func TestLegacyAccountMigratesInPlace(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "pwd")
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte("legacy-password"), bcrypt.MinCost)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
legacy, _ := json.Marshal(legacyAccount{Username: "owner", Password: string(hash)})
|
||||
if err = os.WriteFile(path, legacy, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
store := NewStore(path)
|
||||
user, ok, err := store.Authenticate("owner", "legacy-password")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("authenticate migrated account: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if user.Role != RoleAdmin || !user.SystemAccount || user.TokenVersion == 0 {
|
||||
t.Fatalf("unexpected migrated user: %+v", user)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var db database
|
||||
if err = json.Unmarshal(data, &db); err != nil {
|
||||
t.Fatalf("account file was not migrated: %v", err)
|
||||
}
|
||||
if db.Version != currentFileVersion || len(db.Users) != 1 {
|
||||
t.Fatalf("unexpected database: %+v", db)
|
||||
}
|
||||
var downgrade legacyAccount
|
||||
if err = json.Unmarshal(data, &downgrade); err != nil || downgrade.Username != "owner" || downgrade.Password == "" {
|
||||
t.Fatalf("legacy downgrade mirror is invalid: %+v, %v", downgrade, err)
|
||||
}
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if info.Mode().Perm() != 0o600 {
|
||||
t.Fatalf("account mode = %o, want 600", info.Mode().Perm())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCorruptAccountFileFailsClosed(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "pwd")
|
||||
if err := os.WriteFile(path, []byte("not json"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, ok, err := NewStore(path).Authenticate("admin", "admin")
|
||||
if err == nil || ok {
|
||||
t.Fatalf("corrupt file must fail closed: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserLifecycleAndTokenRevocation(t *testing.T) {
|
||||
store := NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err := store.Create("alice", "correct-horse", RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("alice", "wrong-password"); err != nil || ok {
|
||||
t.Fatalf("wrong password login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
alice, err := store.Get("alice")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = store.ValidateToken("alice", alice.TokenVersion); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = store.SetPassword("alice", "new-password"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = store.ValidateToken("alice", alice.TokenVersion); err == nil {
|
||||
t.Fatal("old token stayed valid after password change")
|
||||
}
|
||||
if _, ok, err := store.Authenticate("alice", "new-password"); err != nil || !ok {
|
||||
t.Fatalf("new password login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
|
||||
disabled := false
|
||||
if _, err = store.Update("admin", "alice", UserPatch{Enabled: &disabled}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("alice", "new-password"); err != nil || ok {
|
||||
t.Fatalf("disabled login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
enabled := true
|
||||
if _, err = store.Update("admin", "alice", UserPatch{Enabled: &enabled}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = store.Delete("admin", "alice"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = store.Get("alice"); !errors.Is(err, ErrUserNotFound) {
|
||||
t.Fatalf("deleted user lookup error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLastAdminAndSelfProtection(t *testing.T) {
|
||||
store := NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
userRole := RoleUser
|
||||
if _, err := store.Update("admin", "admin", UserPatch{Role: &userRole}); !errors.Is(err, ErrSelfModification) {
|
||||
t.Fatalf("self demotion error = %v", err)
|
||||
}
|
||||
if err := store.Delete("admin", "admin"); !errors.Is(err, ErrSelfDelete) {
|
||||
t.Fatalf("self delete error = %v", err)
|
||||
}
|
||||
if err := store.Create("second", "second-password", RoleAdmin); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
disabled := false
|
||||
if _, err := store.Update("second", "admin", UserPatch{Enabled: &disabled}); !errors.Is(err, ErrSystemAccount) {
|
||||
t.Fatalf("disable device owner error = %v", err)
|
||||
}
|
||||
if _, err := store.Update("second", "admin", UserPatch{Role: &userRole}); !errors.Is(err, ErrSystemAccount) {
|
||||
t.Fatalf("demote device owner error = %v", err)
|
||||
}
|
||||
if err := store.Delete("second", "admin"); !errors.Is(err, ErrSystemAccount) {
|
||||
t.Fatalf("delete device owner error = %v", err)
|
||||
}
|
||||
if err := store.Delete("admin", "second"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemovedAccountFileAndRecreatedUserInvalidateOldVersions(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "pwd")
|
||||
store := NewStore(path)
|
||||
admin, ok, err := store.Authenticate("admin", "admin")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err = os.Remove(path); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = store.ValidateToken("admin", admin.TokenVersion); err == nil {
|
||||
t.Fatal("old admin token survived account-file reset")
|
||||
}
|
||||
|
||||
if _, ok, err = store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("login after reset: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err = store.Create("alice", "alice-password", RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
alice, err := store.Get("alice")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = store.Delete("admin", "alice"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err = store.Create("alice", "alice-password", RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = store.ValidateToken("alice", alice.TokenVersion); err == nil {
|
||||
t.Fatal("old token survived deletion and username reuse")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentCreatesDoNotLoseUsers(t *testing.T) {
|
||||
store := NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
|
||||
const count = 5
|
||||
var workers sync.WaitGroup
|
||||
errorsCh := make(chan error, count)
|
||||
for index := 0; index < count; index++ {
|
||||
workers.Add(1)
|
||||
go func(index int) {
|
||||
defer workers.Done()
|
||||
errorsCh <- store.Create(fmt.Sprintf("user%d", index), "valid-password", RoleUser)
|
||||
}(index)
|
||||
}
|
||||
workers.Wait()
|
||||
close(errorsCh)
|
||||
for err := range errorsCh {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
users, err := store.List()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(users) != count+1 {
|
||||
t.Fatalf("got %d users, want %d", len(users), count+1)
|
||||
}
|
||||
}
|
||||
@@ -7,13 +7,6 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// RegenerateSecretKey regenerate secret key when logout
|
||||
func RegenerateSecretKey() {
|
||||
if instance.JWT.RevokeTokensOnLogout {
|
||||
instance.JWT.SecretKey = generateRandomSecretKey()
|
||||
}
|
||||
}
|
||||
|
||||
// Generate random string for secret key.
|
||||
func generateRandomSecretKey() string {
|
||||
b := make([]byte, 64)
|
||||
|
||||
@@ -1,68 +1,160 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/config"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"NanoKVM-Server/config"
|
||||
)
|
||||
|
||||
const (
|
||||
principalContextKey = "principal"
|
||||
tokenContextKey = "token"
|
||||
CookieName = "nano-kvm-token"
|
||||
sessionRecheckDelay = 5 * time.Second
|
||||
)
|
||||
|
||||
type Principal struct {
|
||||
Username string
|
||||
Role authn.Role
|
||||
}
|
||||
|
||||
type Token struct {
|
||||
Username string `json:"username"`
|
||||
Username string `json:"username"`
|
||||
TokenVersion uint64 `json:"tokenVersion"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
func CheckToken() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if allowByToken(c) {
|
||||
principal, token, ok := authenticate(c)
|
||||
if !ok {
|
||||
abortUnauthorized(c)
|
||||
return
|
||||
}
|
||||
|
||||
c.Set(principalContextKey, principal)
|
||||
c.Set(tokenContextKey, token)
|
||||
if token == nil || token.ExpiresAt == nil {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
abortUnauthorized(c)
|
||||
requestContext, cancel := context.WithCancel(c.Request.Context())
|
||||
c.Request = c.Request.WithContext(requestContext)
|
||||
unregister := activeSessions.register(principal.Username, cancel)
|
||||
timer := time.AfterFunc(time.Until(token.ExpiresAt.Time), cancel)
|
||||
go watchSessionState(requestContext, cancel, principal.Username, token.TokenVersion, sessionRecheckDelay)
|
||||
defer func() {
|
||||
timer.Stop()
|
||||
unregister()
|
||||
cancel()
|
||||
}()
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func watchSessionState(ctx context.Context, cancel context.CancelFunc, username string, tokenVersion uint64, interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
if _, err := authn.DefaultStore.ValidateToken(username, tokenVersion); err != nil {
|
||||
cancel()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func RequireRole(roles ...authn.Role) gin.HandlerFunc {
|
||||
allowed := make(map[authn.Role]struct{}, len(roles))
|
||||
for _, role := range roles {
|
||||
allowed[role] = struct{}{}
|
||||
}
|
||||
return func(c *gin.Context) {
|
||||
principal, ok := CurrentPrincipal(c)
|
||||
if !ok {
|
||||
abortUnauthorized(c)
|
||||
return
|
||||
}
|
||||
if _, ok = allowed[principal.Role]; !ok {
|
||||
c.JSON(http.StatusForbidden, "forbidden")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func CurrentPrincipal(c *gin.Context) (Principal, bool) {
|
||||
value, exists := c.Get(principalContextKey)
|
||||
if !exists {
|
||||
return Principal{}, false
|
||||
}
|
||||
principal, ok := value.(Principal)
|
||||
return principal, ok
|
||||
}
|
||||
|
||||
func CheckLoopbackInternalToken() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if allowByLoopbackInternalToken(c.Request) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
abortUnauthorized(c)
|
||||
}
|
||||
}
|
||||
|
||||
func CheckTokenOrLoopbackInternalToken() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if allowByToken(c) || allowByLoopbackInternalToken(c.Request) {
|
||||
if allowByLoopbackInternalToken(c.Request) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
abortUnauthorized(c)
|
||||
principal, token, ok := authenticate(c)
|
||||
if !ok {
|
||||
abortUnauthorized(c)
|
||||
return
|
||||
}
|
||||
c.Set(principalContextKey, principal)
|
||||
c.Set(tokenContextKey, token)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func allowByToken(c *gin.Context) bool {
|
||||
func authenticate(c *gin.Context) (Principal, *Token, bool) {
|
||||
conf := config.GetInstance()
|
||||
|
||||
if conf.Authentication == "disable" {
|
||||
return true
|
||||
return Principal{Username: "admin", Role: authn.RoleAdmin}, nil, true
|
||||
}
|
||||
|
||||
cookie, err := c.Cookie("nano-kvm-token")
|
||||
cookie, err := c.Cookie(CookieName)
|
||||
if err != nil {
|
||||
return false
|
||||
return Principal{}, nil, false
|
||||
}
|
||||
|
||||
_, err = ParseJWT(cookie)
|
||||
return err == nil
|
||||
token, err := ParseJWT(cookie)
|
||||
if err != nil {
|
||||
return Principal{}, nil, false
|
||||
}
|
||||
user, err := authn.DefaultStore.ValidateToken(token.Username, token.TokenVersion)
|
||||
if err != nil {
|
||||
log.Debugf("validate session for %q: %s", token.Username, err)
|
||||
return Principal{}, nil, false
|
||||
}
|
||||
return Principal{Username: user.Username, Role: user.Role}, token, true
|
||||
}
|
||||
|
||||
func abortUnauthorized(c *gin.Context) {
|
||||
@@ -70,37 +162,44 @@ func abortUnauthorized(c *gin.Context) {
|
||||
c.Abort()
|
||||
}
|
||||
|
||||
func GenerateJWT(username string) (string, error) {
|
||||
func GenerateJWT(username string, tokenVersion uint64) (string, error) {
|
||||
conf := config.GetInstance()
|
||||
|
||||
now := time.Now()
|
||||
expireDuration := time.Duration(conf.JWT.RefreshTokenDuration) * time.Second
|
||||
|
||||
claims := Token{
|
||||
Username: username,
|
||||
Username: username,
|
||||
TokenVersion: tokenVersion,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(expireDuration)),
|
||||
Subject: username,
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
ExpiresAt: jwt.NewNumericDate(now.Add(expireDuration)),
|
||||
},
|
||||
}
|
||||
|
||||
t := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
|
||||
return t.SignedString([]byte(conf.JWT.SecretKey))
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(conf.JWT.SecretKey))
|
||||
}
|
||||
|
||||
func ParseJWT(jwtToken string) (*Token, error) {
|
||||
conf := config.GetInstance()
|
||||
|
||||
t, err := jwt.ParseWithClaims(jwtToken, &Token{}, func(token *jwt.Token) (interface{}, error) {
|
||||
return []byte(conf.JWT.SecretKey), nil
|
||||
})
|
||||
parsed, err := jwt.ParseWithClaims(
|
||||
jwtToken,
|
||||
&Token{},
|
||||
func(token *jwt.Token) (interface{}, error) {
|
||||
if token.Method != jwt.SigningMethodHS256 {
|
||||
return nil, errors.New("unexpected signing method")
|
||||
}
|
||||
return []byte(conf.JWT.SecretKey), nil
|
||||
},
|
||||
jwt.WithValidMethods([]string{jwt.SigningMethodHS256.Alg()}),
|
||||
jwt.WithExpirationRequired(),
|
||||
)
|
||||
if err != nil {
|
||||
log.Debugf("parse jwt error: %s", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if claims, ok := t.Claims.(*Token); ok && t.Valid {
|
||||
return claims, nil
|
||||
} else {
|
||||
return nil, err
|
||||
claims, ok := parsed.Claims.(*Token)
|
||||
if !ok || !parsed.Valid || claims.Username == "" || claims.Subject != claims.Username || claims.TokenVersion == 0 {
|
||||
return nil, errors.New("invalid token claims")
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
230
server/middleware/jwt_test.go
Normal file
230
server/middleware/jwt_test.go
Normal file
@@ -0,0 +1,230 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/config"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func TestCheckTokenUsesLiveRoleAndVersion(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
admin, ok, err := store.Authenticate("admin", "admin")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err = store.Create("alice", "valid-password", authn.RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
alice, err := store.Get("alice")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
restore := useTestAuthStore(t, store)
|
||||
defer restore()
|
||||
adminToken, err := GenerateJWT(admin.Username, admin.TokenVersion)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
userToken, err := GenerateJWT(alice.Username, alice.TokenVersion)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.GET("/admin", CheckToken(), RequireRole(authn.RoleAdmin), func(c *gin.Context) {
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
if status := requestWithToken(router, "/admin", ""); status != http.StatusUnauthorized {
|
||||
t.Fatalf("missing token status = %d", status)
|
||||
}
|
||||
if status := requestWithToken(router, "/admin", adminToken); status != http.StatusNoContent {
|
||||
t.Fatalf("admin status = %d", status)
|
||||
}
|
||||
if status := requestWithToken(router, "/admin", userToken); status != http.StatusForbidden {
|
||||
t.Fatalf("user status = %d", status)
|
||||
}
|
||||
if _, err = store.Revoke("admin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if status := requestWithToken(router, "/admin", adminToken); status != http.StatusUnauthorized {
|
||||
t.Fatalf("revoked status = %d", status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseJWTRejectsOtherHMACMethods(t *testing.T) {
|
||||
conf := config.GetInstance()
|
||||
originalSecret := conf.JWT.SecretKey
|
||||
conf.JWT.SecretKey = "test-secret"
|
||||
defer func() { conf.JWT.SecretKey = originalSecret }()
|
||||
|
||||
now := time.Now()
|
||||
claims := Token{
|
||||
Username: "admin",
|
||||
TokenVersion: 1,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: "admin",
|
||||
ExpiresAt: jwt.NewNumericDate(now.Add(time.Hour)),
|
||||
},
|
||||
}
|
||||
token, err := jwt.NewWithClaims(jwt.SigningMethodHS384, claims).SignedString([]byte(conf.JWT.SecretKey))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = ParseJWT(token); err == nil {
|
||||
t.Fatal("HS384 token was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseJWTRejectsExpiredToken(t *testing.T) {
|
||||
conf := config.GetInstance()
|
||||
originalSecret := conf.JWT.SecretKey
|
||||
conf.JWT.SecretKey = "test-secret"
|
||||
defer func() { conf.JWT.SecretKey = originalSecret }()
|
||||
|
||||
claims := Token{
|
||||
Username: "admin",
|
||||
TokenVersion: 1,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: "admin",
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(-time.Minute)),
|
||||
},
|
||||
}
|
||||
token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(conf.JWT.SecretKey))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = ParseJWT(token); err == nil {
|
||||
t.Fatal("expired token was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRevokeUserSessionsCancelsActiveRequests(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
unregister := activeSessions.register("alice", cancel)
|
||||
defer unregister()
|
||||
|
||||
RevokeUserSessions("alice")
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("active session was not cancelled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccountFileResetCancelsActiveSession(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "pwd")
|
||||
store := authn.NewStore(path)
|
||||
user, ok, err := store.Authenticate("admin", "admin")
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
restore := useTestAuthStore(t, store)
|
||||
defer restore()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
go watchSessionState(ctx, cancel, user.Username, user.TokenVersion, 10*time.Millisecond)
|
||||
if err = os.Remove(path); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("active session survived account-file reset")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWatchWebSocketClosesRevokedConnection(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }}
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
connection, err := upgrader.Upgrade(writer, request, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer connection.Close()
|
||||
stop := WatchWebSocket(ctx, connection)
|
||||
defer stop()
|
||||
_, _, _ = connection.ReadMessage()
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
wsURL := "ws" + server.URL[len("http"):]
|
||||
connection, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer connection.Close()
|
||||
cancel()
|
||||
_ = connection.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
_, _, err = connection.ReadMessage()
|
||||
if !websocket.IsCloseError(err, SessionRevokedCloseCode) {
|
||||
t.Fatalf("close error = %v, want code %d", err, SessionRevokedCloseCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckWebSocketOrigin(t *testing.T) {
|
||||
tests := []struct {
|
||||
origin string
|
||||
host string
|
||||
want bool
|
||||
}{
|
||||
{"", "kvm.local", true},
|
||||
{"http://kvm.local", "kvm.local", true},
|
||||
{"http://KVM.local", "kvm.local", true},
|
||||
{"http://kvm.local:8080", "kvm.local:8080", true},
|
||||
{"http://kvm.local:8081", "kvm.local:8080", false},
|
||||
{"http://evil.example", "kvm.local", false},
|
||||
}
|
||||
for _, test := range tests {
|
||||
request := httptest.NewRequest(http.MethodGet, "http://"+test.host+"/api/ws", nil)
|
||||
request.Host = test.host
|
||||
if test.origin != "" {
|
||||
request.Header.Set("Origin", test.origin)
|
||||
}
|
||||
if got := CheckWebSocketOrigin(request); got != test.want {
|
||||
t.Fatalf("origin %q host %q = %v, want %v", test.origin, test.host, got, test.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func requestWithToken(handler http.Handler, path, token string) int {
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
request.AddCookie(&http.Cookie{Name: CookieName, Value: token})
|
||||
handler.ServeHTTP(recorder, request)
|
||||
return recorder.Code
|
||||
}
|
||||
|
||||
func useTestAuthStore(t *testing.T, store *authn.Store) func() {
|
||||
t.Helper()
|
||||
originalStore := authn.DefaultStore
|
||||
conf := config.GetInstance()
|
||||
originalAuthentication := conf.Authentication
|
||||
originalSecret := conf.JWT.SecretKey
|
||||
originalDuration := conf.JWT.RefreshTokenDuration
|
||||
authn.DefaultStore = store
|
||||
conf.Authentication = "enable"
|
||||
conf.JWT.SecretKey = "test-secret"
|
||||
conf.JWT.RefreshTokenDuration = 3600
|
||||
return func() {
|
||||
authn.DefaultStore = originalStore
|
||||
conf.Authentication = originalAuthentication
|
||||
conf.JWT.SecretKey = originalSecret
|
||||
conf.JWT.RefreshTokenDuration = originalDuration
|
||||
}
|
||||
}
|
||||
123
server/middleware/session.go
Normal file
123
server/middleware/session.go
Normal file
@@ -0,0 +1,123 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const SessionRevokedCloseCode = 4401
|
||||
|
||||
type sessionRegistry struct {
|
||||
mutex sync.Mutex
|
||||
nextID atomic.Uint64
|
||||
byUserID map[string]map[uint64]context.CancelFunc
|
||||
}
|
||||
|
||||
var activeSessions = &sessionRegistry{byUserID: make(map[string]map[uint64]context.CancelFunc)}
|
||||
|
||||
func (r *sessionRegistry) register(username string, cancel context.CancelFunc) func() {
|
||||
id := r.nextID.Add(1)
|
||||
r.mutex.Lock()
|
||||
if r.byUserID[username] == nil {
|
||||
r.byUserID[username] = make(map[uint64]context.CancelFunc)
|
||||
}
|
||||
r.byUserID[username][id] = cancel
|
||||
r.mutex.Unlock()
|
||||
|
||||
return func() {
|
||||
r.mutex.Lock()
|
||||
delete(r.byUserID[username], id)
|
||||
if len(r.byUserID[username]) == 0 {
|
||||
delete(r.byUserID, username)
|
||||
}
|
||||
r.mutex.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func RevokeUserSessions(username string) {
|
||||
activeSessions.mutex.Lock()
|
||||
sessions := activeSessions.byUserID[username]
|
||||
delete(activeSessions.byUserID, username)
|
||||
activeSessions.mutex.Unlock()
|
||||
|
||||
for _, cancel := range sessions {
|
||||
cancel()
|
||||
}
|
||||
}
|
||||
|
||||
func WatchWebSocket(ctx context.Context, connection *websocket.Conn) func() {
|
||||
stopped := make(chan struct{})
|
||||
var stopOnce sync.Once
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = connection.WriteControl(
|
||||
websocket.CloseMessage,
|
||||
websocket.FormatCloseMessage(SessionRevokedCloseCode, "session expired or revoked"),
|
||||
time.Now().Add(2*time.Second),
|
||||
)
|
||||
_ = connection.Close()
|
||||
case <-stopped:
|
||||
}
|
||||
}()
|
||||
return func() { stopOnce.Do(func() { close(stopped) }) }
|
||||
}
|
||||
|
||||
func CheckWebSocketOrigin(request *http.Request) bool {
|
||||
origin := strings.TrimSpace(request.Header.Get("Origin"))
|
||||
if origin == "" {
|
||||
return true
|
||||
}
|
||||
parsed, err := url.Parse(origin)
|
||||
if err != nil || parsed.Host == "" {
|
||||
return false
|
||||
}
|
||||
requestScheme := "http"
|
||||
if request.TLS != nil {
|
||||
requestScheme = "https"
|
||||
}
|
||||
if forwarded := strings.TrimSpace(strings.Split(request.Header.Get("X-Forwarded-Proto"), ",")[0]); forwarded != "" {
|
||||
requestScheme = strings.ToLower(forwarded)
|
||||
}
|
||||
return equalOrigin(parsed, request.Host, requestScheme)
|
||||
}
|
||||
|
||||
func equalOrigin(origin *url.URL, requestHost, requestScheme string) bool {
|
||||
originHost, originPort := splitHostPort(origin.Host)
|
||||
host, port := splitHostPort(requestHost)
|
||||
if !strings.EqualFold(originHost, host) {
|
||||
return false
|
||||
}
|
||||
if originPort == "" {
|
||||
originPort = defaultPort(origin.Scheme)
|
||||
}
|
||||
if port == "" {
|
||||
port = defaultPort(requestScheme)
|
||||
}
|
||||
return originPort == port
|
||||
}
|
||||
|
||||
func splitHostPort(value string) (string, string) {
|
||||
host, port, err := net.SplitHostPort(value)
|
||||
if err == nil {
|
||||
return host, port
|
||||
}
|
||||
return strings.Trim(value, "[]"), ""
|
||||
}
|
||||
|
||||
func defaultPort(scheme string) string {
|
||||
switch strings.ToLower(scheme) {
|
||||
case "https", "wss":
|
||||
return "443"
|
||||
default:
|
||||
return "80"
|
||||
}
|
||||
}
|
||||
@@ -5,19 +5,38 @@ type LoginReq struct {
|
||||
Password string `validate:"required"`
|
||||
}
|
||||
|
||||
type LoginRsp struct {
|
||||
Token string `json:"token"`
|
||||
}
|
||||
|
||||
type GetAccountRsp struct {
|
||||
Username string `json:"username"`
|
||||
Role string `json:"role"`
|
||||
}
|
||||
|
||||
type ChangePasswordReq struct {
|
||||
Username string `json:"username" validate:"required"`
|
||||
Password string `json:"password" validate:"required"`
|
||||
CurrentPassword string `json:"currentPassword"`
|
||||
Password string `json:"password" validate:"required"`
|
||||
}
|
||||
|
||||
type IsPasswordUpdatedRsp struct {
|
||||
IsUpdated bool `json:"isUpdated"`
|
||||
}
|
||||
|
||||
type UserInfo struct {
|
||||
Username string `json:"username"`
|
||||
Role string `json:"role"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SystemAccount bool `json:"systemAccount,omitempty"`
|
||||
}
|
||||
|
||||
type ListUsersRsp struct {
|
||||
Users []UserInfo `json:"users"`
|
||||
}
|
||||
|
||||
type CreateUserReq struct {
|
||||
Username string `json:"username" validate:"required"`
|
||||
Password string `json:"password" validate:"required"`
|
||||
Role string `json:"role" validate:"required"`
|
||||
}
|
||||
|
||||
type UpdateUserReq struct {
|
||||
Role *string `json:"role"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
@@ -20,7 +20,7 @@ func ValidateRequest(req interface{}) error {
|
||||
}
|
||||
|
||||
if env == "" || env == "debug" {
|
||||
log.Debugf("request: %+v\n", req)
|
||||
log.Debugf("request validated: %T", req)
|
||||
}
|
||||
|
||||
return nil
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/application"
|
||||
|
||||
@@ -9,7 +10,7 @@ import (
|
||||
|
||||
func applicationRouter(r *gin.Engine) {
|
||||
service := application.NewService()
|
||||
api := r.Group("/api").Use(middleware.CheckToken())
|
||||
api := r.Group("/api").Use(middleware.CheckToken(), middleware.RequireRole(authn.RoleAdmin))
|
||||
|
||||
api.GET("/application/version", service.GetVersion) // get application version
|
||||
api.POST("/application/update", service.Update) // update application
|
||||
|
||||
@@ -3,6 +3,7 @@ package router
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/auth"
|
||||
)
|
||||
@@ -18,4 +19,14 @@ func authRouter(r *gin.Engine) {
|
||||
api.GET("/auth/account", service.GetAccount) // get account
|
||||
api.POST("/auth/password", service.ChangePassword) // change password
|
||||
api.POST("/auth/logout", service.Logout) // logout
|
||||
|
||||
admin := r.Group("/api").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
admin.GET("/auth/users", service.ListUsers)
|
||||
admin.POST("/auth/users", service.CreateUser)
|
||||
admin.PUT("/auth/users/:username", service.UpdateUser)
|
||||
admin.DELETE("/auth/users/:username", service.DeleteUser)
|
||||
admin.POST("/auth/users/:username/password", service.ChangeUserPassword)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/controlmode"
|
||||
"NanoKVM-Server/service/hid"
|
||||
@@ -19,7 +20,10 @@ type setAIControlModeRequest struct {
|
||||
}
|
||||
|
||||
func controlRouter(r *gin.Engine, control *controlmode.Manager, picoclawService *picoclaw.Service) {
|
||||
group := r.Group("/api/ai/control").Use(middleware.CheckToken())
|
||||
group := r.Group("/api/ai/control").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
group.GET("/status", func(c *gin.Context) {
|
||||
status, err := control.Status()
|
||||
if err != nil {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/service/download"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
@@ -9,7 +10,7 @@ import (
|
||||
|
||||
func downloadRouter(r *gin.Engine) {
|
||||
service := download.NewService()
|
||||
api := r.Group("/api").Use(middleware.CheckToken())
|
||||
api := r.Group("/api").Use(middleware.CheckToken(), middleware.RequireRole(authn.RoleAdmin))
|
||||
|
||||
api.POST("/download/image", service.DownloadImage) // download image
|
||||
api.POST("/download/image/cancel", service.CancelDownloadImage) // cancel image download
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/extensions/tailscale"
|
||||
|
||||
@@ -8,7 +9,10 @@ import (
|
||||
)
|
||||
|
||||
func extensionsRouter(r *gin.Engine) {
|
||||
api := r.Group("/api/extensions").Use(middleware.CheckToken())
|
||||
api := r.Group("/api/extensions").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
|
||||
ts := tailscale.NewService()
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ package router
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/hid"
|
||||
)
|
||||
@@ -20,17 +21,21 @@ func hidRouter(r *gin.Engine) {
|
||||
|
||||
api.POST("/hid/paste", service.Paste) // paste
|
||||
|
||||
api.GET("/hid/shortcuts", service.GetShortcuts) // get shortcuts
|
||||
api.POST("/hid/shortcut", service.AddShortcut) // add shortcut
|
||||
api.DELETE("/hid/shortcut", service.DeleteShortcut) // delete shortcut
|
||||
api.GET("/hid/shortcuts", service.GetShortcuts) // get shortcuts
|
||||
api.GET("/hid/shortcut/leader-key", service.GetLeaderKey) // get shortcut leader key
|
||||
|
||||
api.GET("/hid/shortcut/leader-key", service.GetLeaderKey) // set shortcut leader key
|
||||
api.POST("/hid/shortcut/leader-key", service.SetLeaderKey) // set shortcut leader key
|
||||
|
||||
api.GET("/hid/mode", service.GetHidMode) // get hid mode
|
||||
api.POST("/hid/mode", service.SetHidMode) // set hid mode
|
||||
api.POST("/hid/reset", service.ResetHid) // reset hid
|
||||
api.GET("/hid/mode", service.GetHidMode) // get hid mode
|
||||
api.GET("/hid/leds", service.GetKeyboardLedStatus)
|
||||
|
||||
admin := r.Group("/api").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
admin.POST("/hid/mode", service.SetHidMode) // set hid mode
|
||||
admin.POST("/hid/reset", service.ResetHid) // reset hid
|
||||
admin.POST("/hid/shortcut", service.AddShortcut)
|
||||
admin.DELETE("/hid/shortcut", service.DeleteShortcut)
|
||||
admin.POST("/hid/shortcut/leader-key", service.SetLeaderKey)
|
||||
|
||||
localAPI.POST("/usb/recover", service.RecoverUSB)
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package router
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/controlmode"
|
||||
"NanoKVM-Server/service/hid"
|
||||
@@ -23,7 +24,10 @@ func mcpRouter(r *gin.Engine, control *controlmode.Manager, picoclawService *pic
|
||||
picoclawService.PublishControlModeChangedFrom(status, "mcp_config")
|
||||
},
|
||||
)
|
||||
management := r.Group("/api/mcp").Use(middleware.CheckToken())
|
||||
management := r.Group("/api/mcp").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
management.GET("/config", service.GetConfig)
|
||||
management.POST("/config", service.SetConfig)
|
||||
management.POST("/key/regenerate", service.RegenerateAPIKey)
|
||||
|
||||
@@ -3,6 +3,7 @@ package router
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/network"
|
||||
)
|
||||
@@ -15,15 +16,19 @@ func networkRouter(r *gin.Engine) {
|
||||
|
||||
api := r.Group("/api").Use(middleware.CheckToken())
|
||||
|
||||
api.POST("/network/wol", service.WakeOnLAN) // wake on lan
|
||||
api.GET("/network/wol/mac", service.GetMac) // get mac list
|
||||
api.DELETE("/network/wol/mac", service.DeleteMac) // delete mac
|
||||
api.POST("/network/wol/mac/name", service.SetMacName) // set mac name
|
||||
api.POST("/network/wol", service.WakeOnLAN) // wake on lan
|
||||
api.GET("/network/wol/mac", service.GetMac) // get mac list
|
||||
|
||||
api.GET("/network/wifi", service.GetWifi) // get Wi-Fi information
|
||||
api.POST("/network/wifi/connect", service.ConnectWifi) // connect Wi-Fi
|
||||
api.POST("/network/wifi/disconnect", service.DisconnectWifi) // disconnect Wi-Fi
|
||||
admin := r.Group("/api").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
admin.DELETE("/network/wol/mac", service.DeleteMac) // delete mac
|
||||
admin.POST("/network/wol/mac/name", service.SetMacName) // set mac name
|
||||
|
||||
api.GET("/network/dns", service.GetDNS) // get DNS configuration
|
||||
api.POST("/network/dns", service.SetDNS) // set DNS configuration
|
||||
admin.GET("/network/wifi", service.GetWifi) // get Wi-Fi information
|
||||
admin.POST("/network/wifi/connect", service.ConnectWifi) // connect Wi-Fi
|
||||
admin.POST("/network/wifi/disconnect", service.DisconnectWifi) // disconnect Wi-Fi
|
||||
admin.GET("/network/dns", service.GetDNS) // get DNS configuration
|
||||
admin.POST("/network/dns", service.SetDNS) // set DNS configuration
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package router
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/picoclaw"
|
||||
)
|
||||
@@ -39,7 +40,10 @@ func PicoclawLoopbackHTTPAllowedPaths() []string {
|
||||
}
|
||||
|
||||
func picoclawRouter(r *gin.Engine, service *picoclaw.Service) {
|
||||
frontendAPI := r.Group(picoclawBasePath).Use(middleware.CheckToken())
|
||||
frontendAPI := r.Group(picoclawBasePath).Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
localAPI := r.Group(picoclawBasePath).Use(middleware.CheckLoopbackInternalToken())
|
||||
|
||||
localAPI.GET(picoclawScreenshotPath, service.Screenshot)
|
||||
|
||||
@@ -3,13 +3,14 @@ package router
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/storage"
|
||||
)
|
||||
|
||||
func storageRouter(r *gin.Engine) {
|
||||
service := storage.NewService()
|
||||
api := r.Group("/api").Use(middleware.CheckToken())
|
||||
api := r.Group("/api").Use(middleware.CheckToken(), middleware.RequireRole(authn.RoleAdmin))
|
||||
|
||||
api.GET("/storage/image", service.GetImages) // get image list
|
||||
api.GET("/storage/image/mounted", service.GetMountedImage) // get mounted image
|
||||
|
||||
@@ -3,6 +3,7 @@ package router
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/vm"
|
||||
)
|
||||
@@ -11,6 +12,10 @@ func vmRouter(r *gin.Engine) {
|
||||
service := vm.NewService()
|
||||
|
||||
api := r.Group("/api").Use(middleware.CheckToken())
|
||||
admin := r.Group("/api").Use(
|
||||
middleware.CheckToken(),
|
||||
middleware.RequireRole(authn.RoleAdmin),
|
||||
)
|
||||
|
||||
api.GET("/vm/info", service.GetInfo) // get device information
|
||||
api.GET("/vm/hardware", service.GetHardware) // get hardware version
|
||||
@@ -19,55 +24,55 @@ func vmRouter(r *gin.Engine) {
|
||||
api.GET("/vm/gpio", service.GetGpio) // get gpio
|
||||
api.POST("/vm/screen", service.SetScreen) // update screen
|
||||
|
||||
api.GET("/vm/terminal", service.Terminal) // web terminal
|
||||
admin.GET("/vm/terminal", service.Terminal) // web terminal
|
||||
|
||||
api.GET("/vm/script", service.GetScripts) // get script
|
||||
api.POST("/vm/script/upload", service.UploadScript) // upload script
|
||||
api.POST("/vm/script/run", service.RunScript) // run script
|
||||
api.DELETE("/vm/script", service.DeleteScript) // delete script
|
||||
admin.GET("/vm/script", service.GetScripts) // get script
|
||||
admin.POST("/vm/script/upload", service.UploadScript) // upload script
|
||||
admin.POST("/vm/script/run", service.RunScript) // run script
|
||||
admin.DELETE("/vm/script", service.DeleteScript) // delete script
|
||||
|
||||
api.GET("/vm/device/virtual", service.GetVirtualDevice) // get virtual device
|
||||
api.POST("/vm/device/virtual", service.UpdateVirtualDevice) // update virtual device
|
||||
admin.GET("/vm/device/virtual", service.GetVirtualDevice) // get virtual device
|
||||
admin.POST("/vm/device/virtual", service.UpdateVirtualDevice) // update virtual device
|
||||
|
||||
api.GET("/vm/memory/limit", service.GetMemoryLimit) // get memory limit
|
||||
api.POST("/vm/memory/limit", service.SetMemoryLimit) // set memory limit
|
||||
admin.GET("/vm/memory/limit", service.GetMemoryLimit) // get memory limit
|
||||
admin.POST("/vm/memory/limit", service.SetMemoryLimit) // set memory limit
|
||||
|
||||
api.GET("/vm/oled", service.GetOLED) // get OLED configuration
|
||||
api.POST("/vm/oled", service.SetOLED) // set OLED configuration
|
||||
admin.GET("/vm/oled", service.GetOLED) // get OLED configuration
|
||||
admin.POST("/vm/oled", service.SetOLED) // set OLED configuration
|
||||
|
||||
// Only supported by PCIe version
|
||||
api.GET("/vm/hdmi", service.GetHdmiState) // get HDMI state
|
||||
api.POST("/vm/hdmi/reset", service.ResetHdmi) // reset hdmi
|
||||
api.POST("/vm/hdmi/enable", service.EnableHdmi) // enable hdmi
|
||||
api.POST("/vm/hdmi/disable", service.DisableHdmi) // disable hdmi
|
||||
api.POST("/vm/hdmi/timeout", service.SetHdmiIdleTimeout)
|
||||
api.GET("/vm/hdmi", service.GetHdmiState) // get HDMI state
|
||||
api.POST("/vm/hdmi/reset", service.ResetHdmi) // reset hdmi
|
||||
admin.POST("/vm/hdmi/enable", service.EnableHdmi) // enable hdmi
|
||||
admin.POST("/vm/hdmi/disable", service.DisableHdmi) // disable hdmi
|
||||
admin.POST("/vm/hdmi/timeout", service.SetHdmiIdleTimeout)
|
||||
|
||||
api.GET("/vm/ssh", service.GetSSHState) // get SSH state
|
||||
api.POST("/vm/ssh/enable", service.EnableSSH) // enable SSH
|
||||
api.POST("/vm/ssh/disable", service.DisableSSH) // disable SSH
|
||||
admin.GET("/vm/ssh", service.GetSSHState) // get SSH state
|
||||
admin.POST("/vm/ssh/enable", service.EnableSSH) // enable SSH
|
||||
admin.POST("/vm/ssh/disable", service.DisableSSH) // disable SSH
|
||||
|
||||
api.GET("/vm/swap", service.GetSwap) // get swap file size
|
||||
api.POST("/vm/swap", service.SetSwap) // set swap file size
|
||||
admin.GET("/vm/swap", service.GetSwap) // get swap file size
|
||||
admin.POST("/vm/swap", service.SetSwap) // set swap file size
|
||||
|
||||
api.GET("/vm/mouse-jiggler", service.GetMouseJiggler) // get mouse jiggler
|
||||
api.POST("/vm/mouse-jiggler/", service.SetMouseJiggler) // set mouse jiggler
|
||||
admin.GET("/vm/mouse-jiggler", service.GetMouseJiggler) // get mouse jiggler
|
||||
admin.POST("/vm/mouse-jiggler/", service.SetMouseJiggler) // set mouse jiggler
|
||||
|
||||
api.GET("/vm/hostname", service.GetHostname) // Get Hostname
|
||||
api.POST("/vm/hostname", service.SetHostname) // Set Hostname
|
||||
api.GET("/vm/hostname", service.GetHostname) // Get Hostname
|
||||
admin.POST("/vm/hostname", service.SetHostname) // Set Hostname
|
||||
|
||||
api.GET("/vm/web-title", service.GetWebTitle) // Get web title
|
||||
api.POST("/vm/web-title", service.SetWebTitle) // Set web title
|
||||
api.GET("/vm/web-title", service.GetWebTitle) // Get web title
|
||||
admin.POST("/vm/web-title", service.SetWebTitle) // Set web title
|
||||
|
||||
api.GET("/vm/mdns", service.GetMdnsState) // get mDNS state
|
||||
api.POST("/vm/mdns/enable", service.EnableMdns) // enable mDNS
|
||||
api.POST("/vm/mdns/disable", service.DisableMdns) // disable mDNS
|
||||
admin.GET("/vm/mdns", service.GetMdnsState) // get mDNS state
|
||||
admin.POST("/vm/mdns/enable", service.EnableMdns) // enable mDNS
|
||||
admin.POST("/vm/mdns/disable", service.DisableMdns) // disable mDNS
|
||||
|
||||
api.POST("/vm/tls", service.SetTls) // enable/disable TLS
|
||||
admin.POST("/vm/tls", service.SetTls) // enable/disable TLS
|
||||
|
||||
api.GET("/vm/autostart", service.GetAutostart) // get autostart list
|
||||
api.GET("/vm/autostart/:name", service.GetAutostartContent) // get autostart content
|
||||
api.DELETE("/vm/autostart/:name", service.DeleteAutostart) // delete autostart script
|
||||
api.POST("/vm/autostart/:name", service.UploadAutostart) // upload autostart script
|
||||
admin.GET("/vm/autostart", service.GetAutostart) // get autostart list
|
||||
admin.GET("/vm/autostart/:name", service.GetAutostartContent) // get autostart content
|
||||
admin.DELETE("/vm/autostart/:name", service.DeleteAutostart) // delete autostart script
|
||||
admin.POST("/vm/autostart/:name", service.UploadAutostart) // upload autostart script
|
||||
|
||||
api.POST("/vm/system/reboot", service.Reboot) // reboot system
|
||||
admin.POST("/vm/system/reboot", service.Reboot) // reboot system
|
||||
}
|
||||
|
||||
@@ -1,113 +0,0 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/utils"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
const AccountFile = "/etc/kvm/pwd"
|
||||
|
||||
type Account struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"` // should be named HashedPassword for clarity
|
||||
}
|
||||
|
||||
func GetAccount() (*Account, error) {
|
||||
if _, err := os.Stat(AccountFile); err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return getDefaultAccount(), nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(AccountFile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var account Account
|
||||
if err = json.Unmarshal(content, &account); err != nil {
|
||||
log.Errorf("unmarshal account failed: %s", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &account, nil
|
||||
}
|
||||
|
||||
func SetAccount(username string, hashedPassword string) error {
|
||||
account, err := json.Marshal(&Account{
|
||||
Username: username,
|
||||
Password: hashedPassword,
|
||||
})
|
||||
if err != nil {
|
||||
log.Errorf("failed to marshal account information to json: %s", err)
|
||||
return err
|
||||
}
|
||||
|
||||
err = os.MkdirAll(filepath.Dir(AccountFile), 0o644)
|
||||
if err != nil {
|
||||
log.Errorf("create directory %s failed: %s", AccountFile, err)
|
||||
return err
|
||||
}
|
||||
|
||||
err = os.WriteFile(AccountFile, account, 0o644)
|
||||
if err != nil {
|
||||
log.Errorf("write password failed: %s", err)
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func CompareAccount(username string, plainPassword string) bool {
|
||||
account, err := GetAccount()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
if username != account.Username {
|
||||
return false
|
||||
}
|
||||
|
||||
hashedPassword, err := utils.DecodeDecrypt(plainPassword)
|
||||
if err != nil || hashedPassword == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
err = bcrypt.CompareHashAndPassword([]byte(account.Password), []byte(hashedPassword))
|
||||
if err != nil {
|
||||
// Compatible with old versions
|
||||
accountHashedPassword, _ := utils.DecodeDecrypt(account.Password)
|
||||
if accountHashedPassword == hashedPassword {
|
||||
return true
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func DelAccount() error {
|
||||
if err := os.Remove(AccountFile); err != nil {
|
||||
log.Errorf("failed to delete password: %s", err)
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func getDefaultAccount() *Account {
|
||||
hashedPassword, _ := bcrypt.GenerateFromPassword([]byte("admin"), bcrypt.DefaultCost)
|
||||
|
||||
return &Account{
|
||||
Username: "admin",
|
||||
Password: string(hashedPassword),
|
||||
}
|
||||
}
|
||||
@@ -1,11 +1,15 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/config"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/proto"
|
||||
"NanoKVM-Server/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -15,12 +19,9 @@ func (s *Service) Login(c *gin.Context) {
|
||||
var req proto.LoginReq
|
||||
var rsp proto.Response
|
||||
|
||||
// authentication disabled
|
||||
conf := config.GetInstance()
|
||||
if conf.Authentication == "disable" {
|
||||
rsp.OkRspWithData(c, &proto.LoginRsp{
|
||||
Token: "disabled",
|
||||
})
|
||||
rsp.OkRsp(c)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -30,63 +31,98 @@ func (s *Service) Login(c *gin.Context) {
|
||||
rsp.ErrRsp(c, code, msg)
|
||||
return
|
||||
}
|
||||
|
||||
if err := proto.ParseFormRequest(c, &req); err != nil {
|
||||
time.Sleep(3 * time.Second)
|
||||
rsp.ErrRsp(c, -1, "invalid parameters")
|
||||
return
|
||||
}
|
||||
|
||||
if ok := CompareAccount(req.Username, req.Password); !ok {
|
||||
password, err := utils.DecodeDecrypt(req.Password)
|
||||
if err != nil || password == "" {
|
||||
s.loginFailed(c, clientIP, &rsp)
|
||||
return
|
||||
}
|
||||
user, ok, err := authn.DefaultStore.Authenticate(req.Username, password)
|
||||
if err != nil {
|
||||
log.Errorf("load account during login: %s", err)
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
if locked, code, msg := RecordLoginFailure(clientIP); locked {
|
||||
rsp.ErrRsp(c, code, msg)
|
||||
return
|
||||
}
|
||||
|
||||
rsp.ErrRsp(c, -2, "invalid username or password")
|
||||
rsp.ErrRsp(c, -3, "authentication unavailable")
|
||||
return
|
||||
}
|
||||
if !ok {
|
||||
s.loginFailed(c, clientIP, &rsp)
|
||||
return
|
||||
}
|
||||
|
||||
ClearLoginAttempt(clientIP)
|
||||
|
||||
token, err := middleware.GenerateJWT(req.Username)
|
||||
token, err := middleware.GenerateJWT(user.Username, user.TokenVersion)
|
||||
if err != nil {
|
||||
time.Sleep(1 * time.Second)
|
||||
time.Sleep(time.Second)
|
||||
rsp.ErrRsp(c, -3, "generate token failed")
|
||||
return
|
||||
}
|
||||
setSessionCookie(c, token)
|
||||
rsp.OkRsp(c)
|
||||
log.Infof("user logged in: %s", user.Username)
|
||||
}
|
||||
|
||||
rsp.OkRspWithData(c, &proto.LoginRsp{
|
||||
Token: token,
|
||||
})
|
||||
|
||||
log.Debugf("login success, username: %s", req.Username)
|
||||
func (s *Service) loginFailed(c *gin.Context, clientIP string, rsp *proto.Response) {
|
||||
time.Sleep(2 * time.Second)
|
||||
if locked, code, msg := RecordLoginFailure(clientIP); locked {
|
||||
rsp.ErrRsp(c, code, msg)
|
||||
return
|
||||
}
|
||||
rsp.ErrRsp(c, -2, "invalid username or password")
|
||||
}
|
||||
|
||||
func (s *Service) Logout(c *gin.Context) {
|
||||
conf := config.GetInstance()
|
||||
|
||||
if conf.JWT.RevokeTokensOnLogout {
|
||||
config.RegenerateSecretKey()
|
||||
}
|
||||
|
||||
var rsp proto.Response
|
||||
principal, ok := middleware.CurrentPrincipal(c)
|
||||
conf := config.GetInstance()
|
||||
if ok && conf.Authentication != "disable" && conf.JWT.RevokeTokensOnLogout {
|
||||
if _, err := authn.DefaultStore.Revoke(principal.Username); err != nil {
|
||||
rsp.ErrRsp(c, -1, "failed to revoke session")
|
||||
return
|
||||
}
|
||||
middleware.RevokeUserSessions(principal.Username)
|
||||
log.Infof("user logged out: %s", principal.Username)
|
||||
}
|
||||
clearSessionCookie(c)
|
||||
rsp.OkRsp(c)
|
||||
}
|
||||
|
||||
func (s *Service) GetAccount(c *gin.Context) {
|
||||
var rsp proto.Response
|
||||
|
||||
account, err := GetAccount()
|
||||
if err != nil {
|
||||
principal, ok := middleware.CurrentPrincipal(c)
|
||||
if !ok {
|
||||
rsp.ErrRsp(c, -1, "get account failed")
|
||||
return
|
||||
}
|
||||
|
||||
rsp.OkRspWithData(c, &proto.GetAccountRsp{
|
||||
Username: account.Username,
|
||||
Username: principal.Username,
|
||||
Role: string(principal.Role),
|
||||
})
|
||||
log.Debugf("get account successful")
|
||||
}
|
||||
|
||||
func setSessionCookie(c *gin.Context, token string) {
|
||||
conf := config.GetInstance()
|
||||
secure := conf.Proto == "https" || c.Request.TLS != nil ||
|
||||
strings.EqualFold(c.GetHeader("X-Forwarded-Proto"), "https")
|
||||
c.SetSameSite(http.SameSiteStrictMode)
|
||||
c.SetCookie(
|
||||
middleware.CookieName,
|
||||
token,
|
||||
int(conf.JWT.RefreshTokenDuration),
|
||||
"/",
|
||||
"",
|
||||
secure,
|
||||
true,
|
||||
)
|
||||
}
|
||||
|
||||
func clearSessionCookie(c *gin.Context) {
|
||||
conf := config.GetInstance()
|
||||
secure := conf.Proto == "https" || c.Request.TLS != nil ||
|
||||
strings.EqualFold(c.GetHeader("X-Forwarded-Proto"), "https")
|
||||
c.SetSameSite(http.SameSiteStrictMode)
|
||||
c.SetCookie(middleware.CookieName, "", -1, "/", "", secure, true)
|
||||
}
|
||||
|
||||
@@ -1,124 +1,126 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/proto"
|
||||
"NanoKVM-Server/utils"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"time"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/config"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/proto"
|
||||
"NanoKVM-Server/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
var systemPasswordUpdater = changeRootPassword
|
||||
|
||||
func (s *Service) ChangePassword(c *gin.Context) {
|
||||
var req proto.ChangePasswordReq
|
||||
var rsp proto.Response
|
||||
|
||||
if err := proto.ParseFormRequest(c, &req); err != nil {
|
||||
rsp.ErrRsp(c, -1, "invalid parameters")
|
||||
return
|
||||
}
|
||||
|
||||
password, err := utils.DecodeDecrypt(req.Password)
|
||||
if err != nil || password == "" {
|
||||
rsp.ErrRsp(c, -2, "invalid password")
|
||||
return
|
||||
}
|
||||
|
||||
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
rsp.ErrRsp(c, -3, "failed to hash password")
|
||||
return
|
||||
}
|
||||
|
||||
if err = SetAccount(req.Username, string(hashedPassword)); err != nil {
|
||||
rsp.ErrRsp(c, -4, "failed to save password")
|
||||
return
|
||||
}
|
||||
|
||||
// change root password
|
||||
err = changeRootPassword(password)
|
||||
if err != nil {
|
||||
_ = DelAccount()
|
||||
rsp.ErrRsp(c, -5, "failed to change password")
|
||||
principal, ok := middleware.CurrentPrincipal(c)
|
||||
if !ok {
|
||||
rsp.ErrRsp(c, -2, "invalid session")
|
||||
return
|
||||
}
|
||||
currentPassword, err := utils.DecodeDecrypt(req.CurrentPassword)
|
||||
if err != nil || currentPassword == "" {
|
||||
rsp.ErrRsp(c, -3, "current password is required")
|
||||
return
|
||||
}
|
||||
if _, authenticated, authErr := authn.DefaultStore.Authenticate(principal.Username, currentPassword); authErr != nil {
|
||||
rsp.ErrRsp(c, -4, "authentication unavailable")
|
||||
return
|
||||
} else if !authenticated {
|
||||
rsp.ErrRsp(c, -3, "current password is incorrect")
|
||||
return
|
||||
}
|
||||
if err := changeUserPassword(principal.Username, req.Password); err != nil {
|
||||
rsp.ErrRsp(c, -5, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
middleware.RevokeUserSessions(principal.Username)
|
||||
clearSessionCookie(c)
|
||||
rsp.OkRsp(c)
|
||||
log.Debugf("change password success, username: %s", req.Username)
|
||||
log.Infof("password changed for user: %s", principal.Username)
|
||||
}
|
||||
|
||||
func (s *Service) IsPasswordUpdated(c *gin.Context) {
|
||||
var rsp proto.Response
|
||||
if config.GetInstance().Authentication == "disable" {
|
||||
rsp.OkRspWithData(c, &proto.IsPasswordUpdatedRsp{IsUpdated: true})
|
||||
return
|
||||
}
|
||||
principal, ok := middleware.CurrentPrincipal(c)
|
||||
if !ok {
|
||||
rsp.ErrRsp(c, -1, "invalid session")
|
||||
return
|
||||
}
|
||||
user, err := authn.DefaultStore.Get(principal.Username)
|
||||
if err != nil {
|
||||
rsp.ErrRsp(c, -2, "failed to get password state")
|
||||
return
|
||||
}
|
||||
rsp.OkRspWithData(c, &proto.IsPasswordUpdatedRsp{IsUpdated: !user.MustChangePassword})
|
||||
}
|
||||
|
||||
if _, err := os.Stat(AccountFile); err != nil {
|
||||
rsp.OkRspWithData(c, &proto.IsPasswordUpdatedRsp{
|
||||
IsUpdated: false,
|
||||
func changeUserPassword(username, encryptedPassword string) error {
|
||||
password, err := utils.DecodeDecrypt(encryptedPassword)
|
||||
if err != nil || password == "" {
|
||||
return errInvalidPassword
|
||||
}
|
||||
if err = authn.ValidatePassword(password); err != nil {
|
||||
return err
|
||||
}
|
||||
user, err := authn.DefaultStore.Get(username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if user.SystemAccount && user.Role == authn.RoleAdmin && user.Enabled {
|
||||
_, err = authn.DefaultStore.SetPasswordAndRun(username, password, func() error {
|
||||
return systemPasswordUpdater(password)
|
||||
})
|
||||
return
|
||||
return err
|
||||
}
|
||||
|
||||
account, err := GetAccount()
|
||||
if err != nil || account == nil {
|
||||
rsp.ErrRsp(c, -1, "failed to get password")
|
||||
return
|
||||
}
|
||||
|
||||
err = bcrypt.CompareHashAndPassword([]byte(account.Password), []byte("admin"))
|
||||
|
||||
rsp.OkRspWithData(c, &proto.IsPasswordUpdatedRsp{
|
||||
// If the hash is not valid, still assume it's not updated
|
||||
// The error we want to see is password and hash not matching
|
||||
IsUpdated: errors.Is(err, bcrypt.ErrMismatchedHashAndPassword),
|
||||
})
|
||||
_, err = authn.DefaultStore.SetPassword(username, password)
|
||||
return err
|
||||
}
|
||||
|
||||
func changeRootPassword(password string) error {
|
||||
err := passwd(password)
|
||||
if err != nil {
|
||||
if err := passwd(password); err != nil {
|
||||
log.Errorf("failed to change root password: %s", err)
|
||||
return err
|
||||
}
|
||||
|
||||
log.Debugf("change root password successful.")
|
||||
log.Debug("change root password successful")
|
||||
return nil
|
||||
}
|
||||
|
||||
func passwd(password string) error {
|
||||
cmd := exec.Command("passwd", "root")
|
||||
|
||||
stdin, err := cmd.StdinPipe()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
_ = stdin.Close()
|
||||
}()
|
||||
|
||||
defer func() { _ = stdin.Close() }()
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
if err = cmd.Start(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err = io.WriteString(stdin, password+"\n"); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
if _, err = io.WriteString(stdin, password+"\n"); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err = cmd.Wait(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
return cmd.Wait()
|
||||
}
|
||||
|
||||
307
server/service/auth/service_test.go
Normal file
307
server/service/auth/service_test.go
Normal file
@@ -0,0 +1,307 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/config"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/proto"
|
||||
"NanoKVM-Server/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/mervick/aes-everywhere/go/aes256"
|
||||
)
|
||||
|
||||
func TestLoginCookieAndUserAuthorizationLifecycle(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
restore := useTestStore(store)
|
||||
defer restore()
|
||||
|
||||
service := NewService()
|
||||
router := gin.New()
|
||||
router.POST("/login", service.Login)
|
||||
authenticated := router.Group("/").Use(middleware.CheckToken())
|
||||
authenticated.GET("/account", service.GetAccount)
|
||||
authenticated.POST("/logout", service.Logout)
|
||||
admin := router.Group("/").Use(middleware.CheckToken(), middleware.RequireRole(authn.RoleAdmin))
|
||||
admin.GET("/users", service.ListUsers)
|
||||
admin.POST("/users", service.CreateUser)
|
||||
admin.PUT("/users/:username", service.UpdateUser)
|
||||
|
||||
adminCookie := loginCookie(t, router, "admin", "admin")
|
||||
if !adminCookie.HttpOnly || adminCookie.SameSite != http.SameSiteStrictMode || adminCookie.Path != "/" {
|
||||
t.Fatalf("unsafe session cookie: %+v", adminCookie)
|
||||
}
|
||||
|
||||
createBody := map[string]any{
|
||||
"username": "alice",
|
||||
"password": encryptForRequest("valid-password"),
|
||||
"role": "user",
|
||||
}
|
||||
if code := requestJSON(router, http.MethodPost, "/users", createBody, adminCookie); code != http.StatusOK {
|
||||
t.Fatalf("create status = %d", code)
|
||||
}
|
||||
userCookie := loginCookie(t, router, "alice", "valid-password")
|
||||
if code := requestJSON(router, http.MethodGet, "/users", nil, userCookie); code != http.StatusForbidden {
|
||||
t.Fatalf("user management status = %d, want 403", code)
|
||||
}
|
||||
|
||||
disabled := false
|
||||
if code := requestJSON(router, http.MethodPut, "/users/alice", map[string]any{"enabled": disabled}, adminCookie); code != http.StatusOK {
|
||||
t.Fatalf("disable status = %d", code)
|
||||
}
|
||||
if code := requestJSON(router, http.MethodGet, "/account", nil, userCookie); code != http.StatusUnauthorized {
|
||||
t.Fatalf("disabled session status = %d, want 401", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogoutDoesNotInvalidateAnotherUser(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
restore := useTestStore(store)
|
||||
defer restore()
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err := store.Create("alice", "alice-password", authn.RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.Create("bob", "bob-password", authn.RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
service := NewService()
|
||||
router := gin.New()
|
||||
router.POST("/login", service.Login)
|
||||
authenticated := router.Group("/").Use(middleware.CheckToken())
|
||||
authenticated.GET("/account", service.GetAccount)
|
||||
authenticated.POST("/logout", service.Logout)
|
||||
|
||||
aliceCookie := loginCookie(t, router, "alice", "alice-password")
|
||||
bobCookie := loginCookie(t, router, "bob", "bob-password")
|
||||
if code := requestJSON(router, http.MethodPost, "/logout", nil, aliceCookie); code != http.StatusOK {
|
||||
t.Fatalf("logout status = %d", code)
|
||||
}
|
||||
if code := requestJSON(router, http.MethodGet, "/account", nil, aliceCookie); code != http.StatusUnauthorized {
|
||||
t.Fatalf("logged out session status = %d", code)
|
||||
}
|
||||
if code := requestJSON(router, http.MethodGet, "/account", nil, bobCookie); code != http.StatusOK {
|
||||
t.Fatalf("other user's session status = %d", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalUserPasswordNeverChangesSystemPassword(t *testing.T) {
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
restore := useTestStore(store)
|
||||
defer restore()
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err := store.Create("alice", "alice-password", authn.RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
originalUpdater := systemPasswordUpdater
|
||||
updates := 0
|
||||
systemPasswordUpdater = func(string) error {
|
||||
updates++
|
||||
return nil
|
||||
}
|
||||
defer func() { systemPasswordUpdater = originalUpdater }()
|
||||
|
||||
if err := changeUserPassword("alice", encryptForRequest("new-password")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if updates != 0 {
|
||||
t.Fatalf("system password updated %d times for a normal user", updates)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("alice", "new-password"); err != nil || !ok {
|
||||
t.Fatalf("new web password login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelfPasswordChangeRequiresCurrentPassword(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
restore := useTestStore(store)
|
||||
defer restore()
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err := store.Create("alice", "alice-password", authn.RoleUser); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
service := NewService()
|
||||
router := gin.New()
|
||||
router.POST("/login", service.Login)
|
||||
authenticated := router.Group("/").Use(middleware.CheckToken())
|
||||
authenticated.POST("/password", service.ChangePassword)
|
||||
cookie := loginCookie(t, router, "alice", "alice-password")
|
||||
|
||||
recorder := requestJSONRecorder(router, http.MethodPost, "/password", map[string]any{
|
||||
"currentPassword": encryptForRequest("wrong-password"),
|
||||
"password": encryptForRequest("new-password"),
|
||||
}, cookie)
|
||||
var response proto.Response
|
||||
if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
|
||||
t.Fatalf("decode wrong-current-password response: %v", err)
|
||||
}
|
||||
if response.Code != -3 || response.Msg != "current password is incorrect" {
|
||||
t.Fatalf("wrong-current-password response = %+v", response)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("alice", "new-password"); err != nil || ok {
|
||||
t.Fatalf("password changed without current credential: ok=%v err=%v", ok, err)
|
||||
}
|
||||
|
||||
requestJSON(router, http.MethodPost, "/password", map[string]any{
|
||||
"currentPassword": encryptForRequest("alice-password"),
|
||||
"password": encryptForRequest("new-password"),
|
||||
}, cookie)
|
||||
if _, ok, err := store.Authenticate("alice", "new-password"); err != nil || !ok {
|
||||
t.Fatalf("valid password change failed: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdministratorCannotResetDeviceOwnerPassword(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
restore := useTestStore(store)
|
||||
defer restore()
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if err := store.Create("second-admin", "second-password", authn.RoleAdmin); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
originalUpdater := systemPasswordUpdater
|
||||
updates := 0
|
||||
systemPasswordUpdater = func(string) error {
|
||||
updates++
|
||||
return nil
|
||||
}
|
||||
defer func() { systemPasswordUpdater = originalUpdater }()
|
||||
|
||||
service := NewService()
|
||||
router := gin.New()
|
||||
router.POST("/login", service.Login)
|
||||
admin := router.Group("/").Use(middleware.CheckToken(), middleware.RequireRole(authn.RoleAdmin))
|
||||
admin.POST("/users/:username/password", service.ChangeUserPassword)
|
||||
cookie := loginCookie(t, router, "second-admin", "second-password")
|
||||
requestJSON(router, http.MethodPost, "/users/admin/password", map[string]any{
|
||||
"password": encryptForRequest("new-password"),
|
||||
}, cookie)
|
||||
|
||||
if updates != 0 {
|
||||
t.Fatalf("administrator changed the device owner's system password %d times", updates)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("device owner password changed: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("admin", "new-password"); err != nil || ok {
|
||||
t.Fatalf("unauthorized password became active: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemPasswordFailureRollsBackWebPassword(t *testing.T) {
|
||||
store := authn.NewStore(filepath.Join(t.TempDir(), "pwd"))
|
||||
restore := useTestStore(store)
|
||||
defer restore()
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("default login: ok=%v err=%v", ok, err)
|
||||
}
|
||||
|
||||
originalUpdater := systemPasswordUpdater
|
||||
systemPasswordUpdater = func(string) error { return errors.New("passwd failed") }
|
||||
defer func() { systemPasswordUpdater = originalUpdater }()
|
||||
|
||||
if err := changeUserPassword("admin", encryptForRequest("new-password")); err == nil {
|
||||
t.Fatal("system password failure was ignored")
|
||||
}
|
||||
if _, ok, err := store.Authenticate("admin", "admin"); err != nil || !ok {
|
||||
t.Fatalf("old web password was not restored: ok=%v err=%v", ok, err)
|
||||
}
|
||||
if _, ok, err := store.Authenticate("admin", "new-password"); err != nil || ok {
|
||||
t.Fatalf("failed password stayed active: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func loginCookie(t *testing.T, handler http.Handler, username, password string) *http.Cookie {
|
||||
t.Helper()
|
||||
body := map[string]any{"username": username, "password": encryptForRequest(password)}
|
||||
recorder := httptest.NewRecorder()
|
||||
request := jsonRequest(http.MethodPost, "/login", body)
|
||||
handler.ServeHTTP(recorder, request)
|
||||
if recorder.Code != http.StatusOK {
|
||||
t.Fatalf("login status = %d", recorder.Code)
|
||||
}
|
||||
if bytes.Contains(recorder.Body.Bytes(), []byte("token")) {
|
||||
t.Fatalf("login exposed JWT in response: %s", recorder.Body.String())
|
||||
}
|
||||
for _, cookie := range recorder.Result().Cookies() {
|
||||
if cookie.Name == middleware.CookieName {
|
||||
return cookie
|
||||
}
|
||||
}
|
||||
t.Fatal("login did not set session cookie")
|
||||
return nil
|
||||
}
|
||||
|
||||
func requestJSON(handler http.Handler, method, path string, body any, cookie *http.Cookie) int {
|
||||
return requestJSONRecorder(handler, method, path, body, cookie).Code
|
||||
}
|
||||
|
||||
func requestJSONRecorder(handler http.Handler, method, path string, body any, cookie *http.Cookie) *httptest.ResponseRecorder {
|
||||
recorder := httptest.NewRecorder()
|
||||
request := jsonRequest(method, path, body)
|
||||
if cookie != nil {
|
||||
request.AddCookie(cookie)
|
||||
}
|
||||
handler.ServeHTTP(recorder, request)
|
||||
return recorder
|
||||
}
|
||||
|
||||
func jsonRequest(method, path string, body any) *http.Request {
|
||||
var data []byte
|
||||
if body != nil {
|
||||
data, _ = json.Marshal(body)
|
||||
}
|
||||
request := httptest.NewRequest(method, path, bytes.NewReader(data))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
return request
|
||||
}
|
||||
|
||||
func encryptForRequest(password string) string {
|
||||
return url.QueryEscape(aes256.Encrypt(password, utils.SecretKey))
|
||||
}
|
||||
|
||||
func useTestStore(store *authn.Store) func() {
|
||||
originalStore := authn.DefaultStore
|
||||
conf := config.GetInstance()
|
||||
originalAuthentication := conf.Authentication
|
||||
originalSecret := conf.JWT.SecretKey
|
||||
originalDuration := conf.JWT.RefreshTokenDuration
|
||||
originalRevokeOnLogout := conf.JWT.RevokeTokensOnLogout
|
||||
authn.DefaultStore = store
|
||||
conf.Authentication = "enable"
|
||||
conf.JWT.SecretKey = "test-secret"
|
||||
conf.JWT.RefreshTokenDuration = 3600
|
||||
conf.JWT.RevokeTokensOnLogout = true
|
||||
return func() {
|
||||
authn.DefaultStore = originalStore
|
||||
conf.Authentication = originalAuthentication
|
||||
conf.JWT.SecretKey = originalSecret
|
||||
conf.JWT.RefreshTokenDuration = originalDuration
|
||||
conf.JWT.RevokeTokensOnLogout = originalRevokeOnLogout
|
||||
}
|
||||
}
|
||||
127
server/service/auth/users.go
Normal file
127
server/service/auth/users.go
Normal file
@@ -0,0 +1,127 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"NanoKVM-Server/authn"
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/proto"
|
||||
"NanoKVM-Server/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
var errInvalidPassword = errors.New("invalid password")
|
||||
|
||||
func (s *Service) ListUsers(c *gin.Context) {
|
||||
var rsp proto.Response
|
||||
users, err := authn.DefaultStore.List()
|
||||
if err != nil {
|
||||
rsp.ErrRsp(c, -1, "failed to load users")
|
||||
return
|
||||
}
|
||||
result := make([]proto.UserInfo, 0, len(users))
|
||||
for _, user := range users {
|
||||
result = append(result, proto.UserInfo{
|
||||
Username: user.Username,
|
||||
Role: string(user.Role),
|
||||
Enabled: user.Enabled,
|
||||
SystemAccount: user.SystemAccount,
|
||||
})
|
||||
}
|
||||
rsp.OkRspWithData(c, &proto.ListUsersRsp{Users: result})
|
||||
}
|
||||
|
||||
func (s *Service) CreateUser(c *gin.Context) {
|
||||
var req proto.CreateUserReq
|
||||
var rsp proto.Response
|
||||
if err := proto.ParseFormRequest(c, &req); err != nil {
|
||||
rsp.ErrRsp(c, -1, "invalid parameters")
|
||||
return
|
||||
}
|
||||
password, err := decodePassword(req.Password)
|
||||
if err != nil {
|
||||
rsp.ErrRsp(c, -2, err.Error())
|
||||
return
|
||||
}
|
||||
if err = authn.DefaultStore.Create(req.Username, password, authn.Role(req.Role)); err != nil {
|
||||
rsp.ErrRsp(c, -3, err.Error())
|
||||
return
|
||||
}
|
||||
rsp.OkRsp(c)
|
||||
log.Infof("user created: %s (%s)", req.Username, req.Role)
|
||||
}
|
||||
|
||||
func (s *Service) UpdateUser(c *gin.Context) {
|
||||
var req proto.UpdateUserReq
|
||||
var rsp proto.Response
|
||||
if err := proto.ParseFormRequest(c, &req); err != nil || (req.Role == nil && req.Enabled == nil) {
|
||||
rsp.ErrRsp(c, -1, "invalid parameters")
|
||||
return
|
||||
}
|
||||
principal, _ := middleware.CurrentPrincipal(c)
|
||||
patch := authn.UserPatch{Enabled: req.Enabled}
|
||||
if req.Role != nil {
|
||||
role := authn.Role(*req.Role)
|
||||
patch.Role = &role
|
||||
}
|
||||
username := c.Param("username")
|
||||
if _, err := authn.DefaultStore.Update(principal.Username, username, patch); err != nil {
|
||||
rsp.ErrRsp(c, -2, err.Error())
|
||||
return
|
||||
}
|
||||
middleware.RevokeUserSessions(username)
|
||||
rsp.OkRsp(c)
|
||||
log.Infof("user updated: %s", username)
|
||||
}
|
||||
|
||||
func (s *Service) DeleteUser(c *gin.Context) {
|
||||
var rsp proto.Response
|
||||
principal, _ := middleware.CurrentPrincipal(c)
|
||||
username := c.Param("username")
|
||||
if err := authn.DefaultStore.Delete(principal.Username, username); err != nil {
|
||||
rsp.ErrRsp(c, -1, err.Error())
|
||||
return
|
||||
}
|
||||
middleware.RevokeUserSessions(username)
|
||||
rsp.OkRsp(c)
|
||||
log.Infof("user deleted: %s", username)
|
||||
}
|
||||
|
||||
func (s *Service) ChangeUserPassword(c *gin.Context) {
|
||||
var req proto.ChangePasswordReq
|
||||
var rsp proto.Response
|
||||
if err := proto.ParseFormRequest(c, &req); err != nil {
|
||||
rsp.ErrRsp(c, -1, "invalid parameters")
|
||||
return
|
||||
}
|
||||
username := c.Param("username")
|
||||
user, err := authn.DefaultStore.Get(username)
|
||||
if err != nil {
|
||||
rsp.ErrRsp(c, -2, err.Error())
|
||||
return
|
||||
}
|
||||
if user.SystemAccount {
|
||||
rsp.ErrRsp(c, -3, "the device owner must change its own password")
|
||||
return
|
||||
}
|
||||
if err := changeUserPassword(username, req.Password); err != nil {
|
||||
rsp.ErrRsp(c, -4, err.Error())
|
||||
return
|
||||
}
|
||||
middleware.RevokeUserSessions(username)
|
||||
rsp.OkRsp(c)
|
||||
log.Infof("password reset for user: %s", username)
|
||||
}
|
||||
|
||||
func decodePassword(encrypted string) (string, error) {
|
||||
password, err := utils.DecodeDecrypt(encrypted)
|
||||
if err != nil || password == "" {
|
||||
return "", errInvalidPassword
|
||||
}
|
||||
if err = authn.ValidatePassword(password); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return password, nil
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/controlmode"
|
||||
"NanoKVM-Server/service/stream/mjpeg"
|
||||
|
||||
@@ -19,9 +20,7 @@ import (
|
||||
var gatewayUpgrader = websocket.Upgrader{
|
||||
ReadBufferSize: 4096,
|
||||
WriteBufferSize: 4096,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
CheckOrigin: middleware.CheckWebSocketOrigin,
|
||||
}
|
||||
|
||||
type relayResult struct {
|
||||
@@ -88,6 +87,8 @@ func (s *Service) ConnectGateway(c *gin.Context) {
|
||||
GetSessionManager().Remove(sessionID)
|
||||
return
|
||||
}
|
||||
stopSessionWatcher := middleware.WatchWebSocket(c.Request.Context(), downstream)
|
||||
defer stopSessionWatcher()
|
||||
|
||||
GetSessionManager().AttachUpstream(sessionID, upstream)
|
||||
GetSessionManager().AttachDownstream(sessionID, downstream)
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
package direct
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/middleware"
|
||||
"NanoKVM-Server/service/stream"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
@@ -17,9 +17,7 @@ var (
|
||||
streamer = newStreamer()
|
||||
upgrader = websocket.Upgrader{
|
||||
WriteBufferSize: 256 * 1024,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
CheckOrigin: middleware.CheckWebSocketOrigin,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -29,6 +27,8 @@ func Connect(c *gin.Context) {
|
||||
log.Errorf("failed to upgrade to websocket: %s", err)
|
||||
return
|
||||
}
|
||||
stopSessionWatcher := middleware.WatchWebSocket(c.Request.Context(), ws)
|
||||
defer stopSessionWatcher()
|
||||
client := newClient(ws)
|
||||
if flowWindow, err := strconv.Atoi(c.Query("flow")); err == nil && flowWindow > 0 {
|
||||
client.queue.enableFlowControl(flowWindow)
|
||||
|
||||
@@ -2,8 +2,8 @@ package webrtc
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/config"
|
||||
"NanoKVM-Server/middleware"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -17,9 +17,7 @@ import (
|
||||
var (
|
||||
upgrader = websocket.Upgrader{
|
||||
WriteBufferSize: 256 * 1024,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
CheckOrigin: middleware.CheckWebSocketOrigin,
|
||||
}
|
||||
globalManager *WebRTCManager
|
||||
managerOnce sync.Once
|
||||
@@ -39,6 +37,8 @@ func Connect(c *gin.Context) {
|
||||
log.Errorf("failed to create h264 websocket: %s", err)
|
||||
return
|
||||
}
|
||||
stopSessionWatcher := middleware.WatchWebSocket(c.Request.Context(), wsConn)
|
||||
defer stopSessionWatcher()
|
||||
defer func() {
|
||||
_ = wsConn.Close()
|
||||
log.Debugf("h264 websocket disconnected: %s", c.ClientIP())
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
package vm
|
||||
|
||||
import (
|
||||
"NanoKVM-Server/middleware"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"time"
|
||||
@@ -26,9 +26,7 @@ type WinSize struct {
|
||||
var upgrader = websocket.Upgrader{
|
||||
ReadBufferSize: maxMessageSize,
|
||||
WriteBufferSize: maxMessageSize,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
CheckOrigin: middleware.CheckWebSocketOrigin,
|
||||
}
|
||||
|
||||
func (s *Service) Terminal(c *gin.Context) {
|
||||
@@ -37,6 +35,8 @@ func (s *Service) Terminal(c *gin.Context) {
|
||||
log.Errorf("failed to init websocket: %s", err)
|
||||
return
|
||||
}
|
||||
stopSessionWatcher := middleware.WatchWebSocket(c.Request.Context(), ws)
|
||||
defer stopSessionWatcher()
|
||||
defer func() {
|
||||
_ = ws.Close()
|
||||
}()
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"NanoKVM-Server/middleware"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
@@ -13,9 +13,7 @@ type Service struct{}
|
||||
var upgrader = websocket.Upgrader{
|
||||
ReadBufferSize: 1024,
|
||||
WriteBufferSize: 1024,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true
|
||||
},
|
||||
CheckOrigin: middleware.CheckWebSocketOrigin,
|
||||
}
|
||||
|
||||
func NewService() *Service {
|
||||
@@ -28,6 +26,8 @@ func (s *Service) Connect(c *gin.Context) {
|
||||
log.Errorf("create websocket failed: %s", err)
|
||||
return
|
||||
}
|
||||
stopSessionWatcher := middleware.WatchWebSocket(c.Request.Context(), ws)
|
||||
defer stopSessionWatcher()
|
||||
|
||||
log.Debug("websocket connected")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user