46 Commits

Author SHA1 Message Date
GL.iNet-Yongping.Xie 133f344713 Merge branch 'bugfix/v2.1.0' 2026-02-27 18:28:24 -08:00
GL.iNet-Yongping.Xie 7fd9c39ffb fix: resolve issues in v2.1.0
Fix reported bugs and stability issues in v2.1.0.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-27 18:27:43 -08:00
pengyu.lu 96936ecdd8 Merge branch 'dev-ui-0129' 2026-02-28 10:10:40 +08:00
pengyu.lu a44c9a4833 fix: Fix some UI issues 2026-02-28 10:10:11 +08:00
GL.iNet-Yongping.Xie e6da02f711 feat: allow customizing system admin username
Add support to change the default system admin login name via config file.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-25 18:23:43 -08:00
GL.iNet-Yongping.Xie 75a0416221 feat: bump glkvm cloud to v2.0.0
Release glkvm cloud version v2.0.0.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-23 23:22:57 -08:00
GL.iNet-Yongping.Xie dc15e534b8 feat: bump glkvm cloud to v2.0.0
Release glkvm cloud version v2.0.0.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-23 23:08:41 -08:00
GL.iNet-Yongping.Xie ea96caef90 fix: remove test server IP
Remove the test server IP from the Makefile.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-23 18:48:16 -08:00
GL.iNet-Yongping.Xie 13dca8b190 Merge remote-tracking branch 'origin/dev-ui-0129' 2026-02-23 18:14:45 -08:00
GL.iNet-Yongping.Xie 3810f9e72c fix: resolve branch merge conflicts
Resolve merge conflicts between branches.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-23 18:12:38 -08:00
pengyu.lu a8414dd025 fix: Optimize the script logic for adding devices 2026-02-12 10:11:43 +08:00
GL.iNet-Yongping.Xie cea6b19f9a fix: prevent duplicate installation on Linux/macOS clients
Fix an issue where the Linux and macOS clients could be installed multiple times.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-11 06:36:10 -08:00
GL.iNet-Yongping.Xie ba07f6386b fix: correct device migration script logic
Fix the script logic for migrating devices between platforms.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-09 19:55:10 -08:00
GL.iNet-Yongping.Xie 206b0fa68b fix: correct OIDC login flow and admin defaults
Fix the OIDC authentication flow.
Adjust default system admin logic for LDAP and OIDC.
Update LICENSE validity period.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-09 19:31:27 -08:00
GL.iNet-Yongping.Xie e3857e9f61 Merge remote-tracking branch 'origin/dev-ui-0129' into feature/user-device-group 2026-02-09 17:53:24 -08:00
pengyu.lu f3c32b186c fix: Fixed some ui issues 2026-02-10 09:50:01 +08:00
GL.iNet-Yongping.Xie 769b298707 feat: add Linux, macOS, and Windows clients support
GLKVM Light Cloud now supports managing clients on Linux, macOS, and Windows.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-09 02:02:44 -08:00
pengyu.lu 3a503d0e77 feat: Add functions related to device permissions and user permissions 2026-02-06 18:07:17 +08:00
GL.iNet-Yongping.Xie 929e9526d9 fix: reset device status on startup and optimize sorting
Reset all devices to offline on startup and optimize device list sorting.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-06 00:23:40 -08:00
GL.iNet-Yongping.Xie 82a8037244 fix: kick active sessions on user deletion
Terminate active sessions when deleting a user.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-05 23:22:33 -08:00
GL.iNet-Yongping.Xie 9c56f78493 fix: optimize user list logic
Improve the user listing logic.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-05 23:03:12 -08:00
GL.iNet-Yongping.Xie a23912fd89 fix: optimize user deletion logic
Refine user deletion logic to avoid inconsistent states.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-05 20:29:07 -08:00
GL.iNet-Yongping.Xie dff4d06107 fix: sync with frontend and fix bugs
- Align API behavior with frontend integration
- Fix issues found during joint debugging

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-05 18:22:29 -08:00
GL.iNet-Yongping.Xie 359f29484a fix: resolve merge conflicts
- Resolve conflicts introduced during branch merge

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-03 19:16:38 -08:00
GL.iNet-Yongping.Xie f1fb5b52fb fix: remove unused control permission
- Remove obsolete control permission that is no longer needed

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-03 18:59:25 -08:00
GL.iNet-Yongping.Xie c1b987e6b3 fix: remove redundant query parameters from redirect URL
Remove unnecessary and ineffective query parameters when redirecting to the remote web interface.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-03 18:33:16 -08:00
GL.iNet-Yongping.Xie 1320b58cbb fix: resolve frontend integration issues and remove unused APIs
- Fix bugs found during frontend integration
- Clean up redundant and unused endpoints

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-03 18:14:51 -08:00
GL.iNet-Yongping.Xie a1ebef3326 feat: add groupId filtering to device list
Add support for filtering the device list by groupId.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-02-01 18:43:46 -08:00
GL.iNet-Yongping.Xie fdba252fa8 fix: switch Docker Compose network mode to bridge
Change Docker Compose from host network mode to bridge mode.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-29 19:33:17 -08:00
GL.iNet-Yongping.Xie 9faeac6ff9 fix: remove redundant query parameters from redirect URL
Remove unnecessary and ineffective query parameters when redirecting to the remote web interface.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-29 19:21:11 -08:00
GL.iNet-Yongping.Xie 3e8d898a67 fix: return created device group ID in create API
Update the device group creation API to return the newly created device group ID.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-29 02:37:25 -08:00
GL.iNet-Yongping.Xie b2a9c6c7dd fix: remove and adjust legacy APIs
Remove obsolete interfaces and adjust remaining legacy APIs for consistency.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-28 20:21:49 -08:00
GL.iNet-Yongping.Xie 9c4592e697 fix: migrate legacy APIs to new project structure
Move deprecated interfaces into the updated project layout to align with the new architecture.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-28 19:41:38 -08:00
GL.iNet-Yongping.Xie f783fa14fe fix: resolve issues found during API self-testing
Fix bugs identified during API self-testing and validation.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-27 01:23:44 -08:00
GL.iNet-Yongping.Xie 6a01778bec feat: introduce device and user group management APIs
Add API endpoints to manage devices, device groups, users, and user groups.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-27 00:27:01 -08:00
GL.iNet-Yongping.Xie 3b5360478d git commit -m "fix: adjust Makefile for local debugging
Update the Makefile to improve convenience and usability for local development and debugging.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>"
2026-01-22 11:09:21 +08:00
GL.iNet-Yongping.Xie 39dc768b9c fix: clean up unused files and update Makefile
Remove obsolete files and adjust the Makefile to simplify the project structure.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-21 18:48:42 -08:00
GL.iNet-Yongping.Xie 606377fbf8 fix: remove unused legacy main entry
Remove the obsolete legacy main entry to avoid confusion and reduce maintenance overhead.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-21 02:33:41 -08:00
GL.iNet-Yongping.Xie b331a7c7f3 fix: remove unused legacy main entry
Remove the obsolete legacy main entry to avoid confusion and reduce maintenance overhead.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-21 02:28:05 -08:00
GL.iNet-Yongping.Xie 2b6a6eff77 fix: remove unused legacy main entry
Remove the obsolete legacy main entry to avoid confusion and reduce maintenance overhead.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-21 02:16:29 -08:00
GL.iNet-Yongping.Xie d93e5953ad fix: refactor rtty project structure
Clean up and reorganize the project layout to improve readability and maintainability.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-21 02:04:36 -08:00
GL.iNet-Yongping.Xie 0d130ebc29 fix: remove command-line config parsing logic
Remove redundant and unnecessary configuration parsing from command-line arguments to simplify startup logic.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-21 00:34:36 -08:00
GL.iNet-Yongping.Xie 2d2658383b fix: clarify handler naming to avoid role confusion
Rename handlers to remove the misleading 'admin' term and better reflect their actual responsibilities, avoiding confusion with admin roles.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-19 18:03:04 -08:00
GL.iNet-Yongping.Xie 7931fa0efd fix: refactor legacy login endpoint
Refactor legacy login endpoint and keep compatibility with existing authentication logic.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-19 02:15:07 -08:00
GL.iNet-Yongping.Xie ab77f53a77 feat: add user, group, and device relationship management APIs
Implement editing APIs for user-to-user-group, user-to-device, and device-group-to-device relationships, enabling flexible management of access and resource associations.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-15 18:11:13 -08:00
GL.iNet-Yongping.Xie a5bce3b289 feat: implement resource permission framework
Complete the initial framework for resource-based permission management, providing the foundation for role-based access control and future permission extensions.

Signed-off-by: GL.iNet-Yongping.Xie <yongping.xie@gl-inet.com>
2026-01-15 01:29:02 -08:00
160 changed files with 12305 additions and 4646 deletions
+2
View File
@@ -12,3 +12,5 @@
# Project-local glide cache, RE: https://github.com/Masterminds/glide/issues/736
.glide/
.vscode
+8
View File
@@ -0,0 +1,8 @@
{
"cSpell.words": [
"ddns",
"glkvm",
"repassword",
"webrtc"
]
}
+1 -1
View File
@@ -9,7 +9,7 @@ Additional Use Grant: You may use the Licensed Work free of charge for
non-production purposes, including development,
testing, personal, or academic use.
Change Date: 2029-01-01
Change Date: 2030-01-01
Change License: GNU General Public License, version 3 (GPLv3)
+17 -55
View File
@@ -2,94 +2,46 @@
# ---------------- Project ----------------
BINARY_NAME ?= rttys
UI_DIR ?= ui
CONF_FILE ?= ./rttys.conf
GO_MAIN ?= ./cmd/glkvm-cloud
# Go build flags
BUILD_FLAGS ?= -ldflags "-s -w"
# Output dir for cross builds
DIST_DIR ?= dist
# Image name
IMAGE_NAME ?= glkvm-cloud
IMAGE_TAG ?= build
UNAME_S := $(shell uname -s)
UNAME_M := $(shell uname -m)
GOOS ?= $(shell go env GOOS)
GOARCH ?= $(shell go env GOARCH)
# Map uname -m -> goarch
ifeq ($(UNAME_M),x86_64)
HOST_GOARCH := amd64
else ifeq ($(UNAME_M),aarch64)
HOST_GOARCH := arm64
else ifeq ($(UNAME_M),arm64)
HOST_GOARCH := arm64
else
HOST_GOARCH := $(GOARCH)
endif
# ---------------- Commands ----------------
GO_BUILD_CMD = go build $(BUILD_FLAGS) -o $(BINARY_NAME)
.PHONY: all ui build run build-all build-run full-run \
.PHONY: all ui debug-local \
build-linux-amd64 build-linux-arm64 build-linux-all \
docker-build docker-fullbuild docker-buildx docker-buildx-full
docker-buildx docker-buildx-full
all: build
all: build-linux-amd64 build-linux-arm64
# Build frontend files only
ui:
cd $(UI_DIR) && npm install && npm run build
# Build for current env (native)
build:
CGO_ENABLED=0 GOOS=$(GOOS) GOARCH=$(GOARCH) $(GO_BUILD_CMD)
# Run Go program only (native binary)
run:
./$(BINARY_NAME) -c $(CONF_FILE)
# Build frontend and Go binary
build-all: ui build
# Build Go binary and run
build-run: build run
# Build frontend, build Go binary, and run
full-run: ui build run
# ---------------- Cross compile (Linux) ----------------
# Produce: dist/rttys-linux-amd64 , dist/rttys-linux-arm64
build-linux-amd64:
@mkdir -p $(DIST_DIR)
CGO_ENABLED=0 GOOS=linux GOARCH=amd64 \
go build $(BUILD_FLAGS) -o $(DIST_DIR)/$(BINARY_NAME)-linux-amd64
go build $(BUILD_FLAGS) -o $(DIST_DIR)/$(BINARY_NAME)-linux-amd64 $(GO_MAIN)
build-linux-arm64:
@mkdir -p $(DIST_DIR)
CGO_ENABLED=0 GOOS=linux GOARCH=arm64 \
go build $(BUILD_FLAGS) -o $(DIST_DIR)/$(BINARY_NAME)-linux-arm64
build-linux-all: build-linux-amd64 build-linux-arm64
# ---------------- Docker (single-arch) ----------------
# Build Docker image using current host arch
docker-build: build
docker build -t $(IMAGE_NAME):$(IMAGE_TAG) .
# Full build Docker image
docker-fullbuild: ui build
docker build -t $(IMAGE_NAME):$(IMAGE_TAG) .
go build $(BUILD_FLAGS) -o $(DIST_DIR)/$(BINARY_NAME)-linux-arm64 $(GO_MAIN)
# ---------------- Docker Buildx ----------------
# Multi-arch build
# Usage:
# make docker-buildx GOARCH=amd64 IMAGE_TAG=build-amd64
# make docker-buildx GOARCH=arm64 IMAGE_TAG=build-arm64
PLATFORMS ?= linux/amd64,linux/arm64
REGISTRY ?=
# If REGISTRY is set, tag becomes: REGISTRY/IMAGE_NAME:IMAGE_TAG
@@ -107,6 +59,16 @@ docker-buildx:
-t $(IMAGE_REF) \
--load .
docker-buildx-full: ui
@$(MAKE) docker-buildx
DEBUG_HOST ?= root@xxxxxxxxxx
DEBUG_PATH ?= /root/glkvmcloudbuild.tar
# Local debug bundle (amd64 image + save tar), upload to debug host, then load
debug-local: build-linux-amd64 docker-buildx
docker save $(IMAGE_NAME):$(IMAGE_TAG) -o glkvmcloudbuild.tar
ssh $(DEBUG_HOST) "rm -f $(DEBUG_PATH)"
scp glkvmcloudbuild.tar $(DEBUG_HOST):$(DEBUG_PATH)
ssh $(DEBUG_HOST) "docker load < $(DEBUG_PATH)"
ssh $(DEBUG_HOST) "cd /root/glkvm_cloud && docker-compose down && docker-compose up -d"
+16 -11
View File
@@ -6,7 +6,7 @@ Self-Deployed Lightweight Cloud is a lightweight KVM remote cloud platform tailo
#### Main Functions and Features
- **Device Management** - Online device list monitoring
- **User Groups and Device Support** - Supports user groups managing specific device groups, enabling different users to manage different devices
- **Script Deployment** - Convenient script-based device addition
- **Remote SSH** - Web SSH remote connections
- **Remote Control** - Web remote desktop control
@@ -17,8 +17,8 @@ Self-Deployed Lightweight Cloud is a lightweight KVM remote cloud platform tailo
- **Lightweight Design** - Optimized for small businesses and individual users
- **Enterprise Authentication** - Supports both **LDAP** and **OIDC** login methods for enterprise users.
- **Deployment** - Supports both **internal network** and **public internet** deployments
- **Platform Compatibility** - Supports both **x86_64** and **arm64** platforms
- **Deployment & Platform Compatibility** - Supports both **internal network** and **public internet** deployments on **x86_64** and **arm64** platforms
- **HTTP/HTTPS Web Proxy Support** - Supports onboarding OpenWrt, ImmortalWrt, Raspberry Pi, Linux VPS, macOS, and Windows hosts into self-hosted GLKVM Cloud for centralized management, and using them as HTTP/HTTPS web proxy nodes for NAT traversal access
## Self-Hosting Guide
@@ -84,21 +84,22 @@ Run **as root**:
### 🌐 Platform Access
Once the installation is complete, access the platform via:
Once the installation is complete, the installer will print the platform URL and admin login credentials in the console. You can access the platform via:
```
https://<your_server_public_ip>
```
⚠️ **Note**: Accessing via IP address will trigger a **browser certificate warning**.
To eliminate the warning, it's recommended to configure a **custom domain** with a valid SSL certificate.
⚠️ **Note**: Accessing via an IP address will trigger a **browser certificate warning**.
To remove the warning, configure your own domain and a valid SSL certificate.
### 🔑 Web UI Login Password
### 🔑 Web UI Login Credentials
The default login password for the Web UI will be displayed in the installation script output:
At the end of the installation script, the console will display the Web UI administrator username and password (for example):
```
🔐 Please check the installation console for your web login password.
```text
👤 Admin username: admin
🔑 Admin password: <auto-generated-password>
```
![](img/password.png)
@@ -128,6 +129,10 @@ The default login password for the Web UI will be displayed in the installation
![Remote Control Screenshot](img/web.png)
### Web Proxy
![Web Proxy Screenshot](img/httpproxy.png)
## Use your own SSL Certificate (Optional)
@@ -234,4 +239,4 @@ Once everything is configured, you can access the platform via your domain:
```
https://www.your-domain.com
```
```
+14 -9
View File
@@ -7,7 +7,7 @@
#### 主要功能与特性
* **设备管理** - 实时查看设备在线状态
* **用户组和设备支持** - 支持用户组管理特定的设备组设备,实现不同用户管理不同设备
* **脚本部署** - 通过脚本快速添加设备
* **远程 SSH** - Web SSH 远程连接
* **远程控制** - Web远程桌面控制
@@ -18,8 +18,8 @@
* **轻量设计** - 专为小型企业和个人优化
* **企业级认证** - 同时支持 **LDAP** 和 **OIDC** 登录方式,适用于企业用户。
- **部署方式** - 同时支持 **内网部署** 和 **公网部署**
- **平台兼容性** - 同时支持 **x86_64** 和 **arm64** 平台
- **部署与平台兼容性** - 同时支持 **内网部署** 和 **公网部署**,并兼容 **x86_64** 与 **arm64** 平台
- **HTTP/HTTPS Web代理功能支持** - 支持 OpenWrt、ImmortalWrt、树莓派、Linux VPS、macOS、Windows 等主机接入自部署 GLKVM Cloud 进行统一管理,并可作为 HTTP/HTTPS Web 代理节点实现内网穿透访问
## 自部署指南
@@ -85,21 +85,22 @@
### 🌐 平台访问
安装完成后,你可以通过以下方式访问平台:
安装完成后,安装脚本会在控制台输出平台访问地址和管理员登录信息。你可以通过以下方式访问平台:
```
https://<你的服务器公网IP>
```
⚠️ **提示**:通过 IP 访问时,浏览器会提示 **证书不受信任**。
如果想消除该提示,建议配置 **自定义域名 + 有效 SSL 证书**。
如需消除该提示,建议配置 **自定义域名 + 有效 SSL 证书**。
### 🔑 Web UI 登录密码
### 🔑 Web UI 登录信息
Web UI 的默认登录密码会在安装脚本运行结束时显示:
安装脚本运行结束后,安装控制台会显示 Web UI 管理员用户名和密码(示例):
```
🔐 请在安装控制台查看 Web 登录密码
```text
👤 管理员用户名:admin
🔑 管理员密码:<自动生成密码>
```
![](img/password.png)
@@ -128,6 +129,10 @@ Web UI 的默认登录密码会在安装脚本运行结束时显示:
![远程桌面截图](img/web.png)
#### Web代理功能
![Web代理功能截图](img/httpproxy.png)
## 使用自有 SSL 证书(可选)
⚠️ **可选配置**:
-644
View File
@@ -1,644 +0,0 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
import (
"crypto/tls"
"embed"
"io/fs"
"net"
"net/http"
"path"
"sort"
"strings"
"time"
"rttys/utils"
"github.com/fanjindong/go-cache"
"github.com/gin-contrib/cors"
"github.com/gin-gonic/gin"
"github.com/rs/zerolog/log"
)
var httpSessions = cache.NewMemCache(cache.WithClearInterval(time.Minute))
const httpSessionExpire = 30 * time.Minute
//go:embed all:ui/dist
var staticFs embed.FS
func (srv *RttyServer) ListenAPI() error {
cfg := &srv.cfg
gin.SetMode(gin.ReleaseMode)
r := gin.New()
r.Use(func(c *gin.Context) {
hi := getHostInfoFromRequest(c.Request)
host := hi.Host
allowedHost := cfg.WebUIHost
// If WebUIHost is configured, enforce host validation
if allowedHost != "" && !isIPHost(host) {
if !domainAllowed(host, allowedHost) {
html := generateErrorHTML("invalid")
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(html))
c.Abort()
return
}
}
c.Next()
})
if cfg.AllowOrigins {
log.Debug().Msg("Allow all origins")
r.Use(cors.Default())
}
authorized := r.Group("/", func(c *gin.Context) {
if !cfg.LocalAuth && isLocalRequest(c) {
return
}
if !httpAuth(cfg, c) {
c.AbortWithStatus(http.StatusUnauthorized)
return
}
})
authorized.GET("/connect/:devid", func(c *gin.Context) {
if !callUserHookUrl(cfg, c) {
c.Status(http.StatusForbidden)
return
}
if c.GetHeader("Upgrade") != "websocket" {
group := c.Query("group")
devid := c.Param("devid")
if dev := srv.GetDevice(group, devid); dev == nil {
c.Redirect(http.StatusFound, "/error/offline")
return
}
url := "/rtty/" + devid
if group != "" {
url += "?group=" + group
}
c.Redirect(http.StatusFound, url)
} else {
handleUserConnection(srv, c)
}
})
authorized.GET("/counts", func(c *gin.Context) {
count := 0
srv.groups.Range(func(key, value any) bool {
count += int(value.(*DeviceGroup).count.Load())
return true
})
c.JSON(http.StatusOK, gin.H{"count": count})
})
authorized.GET("/groups", func(c *gin.Context) {
groups := []string{""}
srv.groups.Range(func(key, value any) bool {
if key != "" {
groups = append(groups, key.(string))
}
return true
})
c.JSON(http.StatusOK, groups)
})
authorized.GET("/devs", func(c *gin.Context) {
devs := make([]*DeviceInfo, 0)
keyword := c.Query("keyword")
// 1. Query all device metadata from DB (offline + online)
metas, err := GetAllDeviceMeta(keyword)
if err != nil || len(metas) == 0 {
c.JSON(http.StatusOK, devs)
return
}
// 2. Build online device map from memory
onlineMap := make(map[string]*Device)
g := srv.GetGroup("", false)
if g != nil {
g.devices.Range(func(key, value any) bool {
dev := value.(*Device)
onlineMap[dev.id] = dev
return true
})
}
now := time.Now().Unix()
// 3. Iterate metas (DB is the source of truth)
for _, meta := range metas {
info := &DeviceInfo{
ID: meta.DeviceID,
Mac: meta.Mac,
Connected: 0,
Uptime: 0,
Desc: meta.Description,
Proto: 0,
IPaddr: meta.IP, // fallback: last known IP
}
// 4. If device is online, override with in-memory data
if dev, ok := onlineMap[meta.DeviceID]; ok {
info.Connected = uint32(now - dev.timestamp)
info.Uptime = dev.uptime
info.Proto = dev.proto
if addr, ok := dev.conn.RemoteAddr().(*net.TCPAddr); ok {
info.IPaddr = addr.IP.String()
} else if host, _, err := net.SplitHostPort(dev.conn.RemoteAddr().String()); err == nil {
info.IPaddr = host
}
}
devs = append(devs, info)
}
// Sort devices:
// 1. Online devices first (Connected > 0)
// 2. Within the same online/offline group, sort by device ID alphabetically
sort.Slice(devs, func(i, j int) bool {
di := devs[i]
dj := devs[j]
// Determine online status
diOnline := di.Connected > 0
djOnline := dj.Connected > 0
if diOnline != djOnline {
return diOnline
}
// If both devices are in the same state (online or offline),
// sort by device ID in ascending alphabetical order
return di.ID < dj.ID
})
c.JSON(http.StatusOK, devs)
})
// UpdateDeviceMetaRequest defines the JSON payload to update device metadata.
// Only DeviceID is mandatory; other fields are optional and will be updated
// only when provided.
type UpdateDeviceMetaRequest struct {
DeviceID string `json:"deviceId" binding:"required"` // DeviceID is the unique device identifier (immutable).
Description string `json:"description,omitempty"` // Description can be updated if provided.
}
// Update device metadata (new interface)
authorized.POST("/devs/update", func(c *gin.Context) {
var req UpdateDeviceMetaRequest
// 1. Parse JSON body
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{
"code": 400,
"msg": "invalid request body",
"err": err.Error(),
})
return
}
// 2. Load existing metadata by device_id
meta, err := GetDeviceMetaByDeviceID(req.DeviceID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"code": 500,
"msg": "failed to query device meta",
"err": err.Error(),
})
return
}
if meta == nil {
c.JSON(http.StatusNotFound, gin.H{
"code": 404,
"msg": "device meta not found",
})
return
}
// 3. Merge data: deviceID/mac/ip, now only description
newDesc := meta.Description
if req.Description != "" {
newDesc = req.Description
}
// 4. Reuse SaveOrUpdateDeviceMeta for UPSERT
if err := SaveOrUpdateDeviceMeta(
meta.DeviceID, // keep original device_id
meta.Mac, // keep original MAC, not editable
newDesc, // new description from request
meta.IP,
); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"code": 500,
"msg": "failed to update device meta",
"err": err.Error(),
})
return
}
c.JSON(http.StatusOK, gin.H{
"code": 0,
"msg": "ok",
})
})
// DeleteDeviceMetaRequest is used to logically delete a device meta record.
// Only DeviceID is required.
type DeleteDeviceMetaRequest struct {
DeviceID string `json:"deviceId" binding:"required"` // DeviceID is the unique device identifier (immutable).
}
// Delete device metadata (physical delete)
authorized.POST("/devs/delete", func(c *gin.Context) {
var req DeleteDeviceMetaRequest
// 1. Parse JSON body
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{
"code": 400,
"msg": "invalid request body",
"err": err.Error(),
})
return
}
// 2. Check existence first (optional but recommended)
meta, err := GetDeviceMetaByDeviceID(req.DeviceID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"code": 500,
"msg": "failed to query device meta",
"err": err.Error(),
})
return
}
if meta == nil {
c.JSON(http.StatusNotFound, gin.H{
"code": 404,
"msg": "device meta not found",
})
return
}
// 3. Physical delete
if err := DeleteDeviceMetaByDeviceID(req.DeviceID); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"code": 500,
"msg": "failed to delete device meta",
"err": err.Error(),
})
return
}
// 4. Success response
c.JSON(http.StatusOK, gin.H{
"code": 0,
"msg": "ok",
})
})
authorized.GET("/dev/:devid", func(c *gin.Context) {
if dev := srv.GetDevice(c.Query("group"), c.Param("devid")); dev != nil {
info := &DeviceInfo{
ID: dev.id,
Desc: dev.desc,
Connected: uint32(time.Now().Unix() - dev.timestamp),
Uptime: dev.uptime,
Proto: dev.proto,
IPaddr: dev.conn.RemoteAddr().(*net.TCPAddr).IP.String(),
}
c.JSON(http.StatusOK, info)
} else {
c.Status(http.StatusNotFound)
}
})
authorized.POST("/cmd/:devid", func(c *gin.Context) {
if !callUserHookUrl(cfg, c) {
c.Status(http.StatusForbidden)
return
}
cmdInfo := &CommandReqInfo{}
err := c.BindJSON(&cmdInfo)
if err != nil || cmdInfo.Cmd == "" || cmdInfo.Username == "" {
cmdErrResp(c, rttyCmdErrInvalid)
return
}
dev := srv.GetDevice(c.Query("group"), c.Param("devid"))
if dev == nil {
cmdErrResp(c, rttyCmdErrOffline)
return
}
dev.handleCmdReq(c, cmdInfo)
})
authorized.Any("/web/:devid/:proto/:addr/*path", func(c *gin.Context) {
httpProxyRedirect(srv, c, "")
})
authorized.Any("/web2/:group/:devid/:proto/:addr/*path", func(c *gin.Context) {
group := c.Param("group")
httpProxyRedirect(srv, c, group)
})
authorized.GET("/signout", func(c *gin.Context) {
sid, err := c.Cookie("sid")
if err != nil || !httpSessions.Exists(sid) {
return
}
httpSessions.Del(sid)
c.Status(http.StatusOK)
})
r.POST("/signin", func(c *gin.Context) {
type credentials struct {
Username string `json:"username"`
Password string `json:"password"`
AuthMethod string `json:"authMethod"`
}
creds := credentials{}
err := c.BindJSON(&creds)
if err != nil {
c.Status(http.StatusBadRequest)
return
}
// 自动确定认证方法或使用指定的方法 (Auto-determine auth method or use specified method)
authMethod := creds.AuthMethod
if authMethod == "" {
// 基于是否提供用户名进行自动检测 (Auto-detect based on whether username is provided)
if creds.Username != "" && cfg.LdapEnabled {
authMethod = "ldap"
} else {
authMethod = "legacy"
}
}
success, errorType := AuthenticateUserWithError(cfg, creds.Username, creds.Password, authMethod)
if success {
sid := utils.GenUniqueID()
httpSessions.Set(sid, true, cache.WithEx(httpSessionExpire))
c.SetCookie("sid", sid, 0, "", "", false, true)
c.Status(http.StatusOK)
return
}
// 根据错误类型返回适当的错误信息 (Return appropriate error message based on error type)
if errorType == "authorization" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "user not authorized"})
} else {
c.JSON(http.StatusUnauthorized, gin.H{"error": "authentication failed"})
}
})
r.GET("/auth-config", func(c *gin.Context) {
authConfig := gin.H{
"ldapEnabled": cfg.LdapEnabled,
"legacyPassword": cfg.Password != "",
"oidcEnabled": cfg.OIDCEnabled,
"kvmCloudVersion": KVMCloudVersion,
}
c.JSON(http.StatusOK, authConfig)
})
r.GET("/alive", func(c *gin.Context) {
if !httpAuth(cfg, c) {
c.AbortWithStatus(http.StatusUnauthorized)
} else {
c.Status(http.StatusOK)
}
})
// ===== 添加OIDC路由 =====
RegisterOIDCRoutes(r, cfg)
fs, err := fs.Sub(staticFs, "ui/dist")
if err != nil {
return err
}
root := http.FS(fs)
fh := http.FileServer(root)
r.NoRoute(func(c *gin.Context) {
upath := path.Clean(c.Request.URL.Path)
if strings.HasSuffix(upath, ".js") || strings.HasSuffix(upath, ".css") {
if strings.Contains(c.Request.Header.Get("Accept-Encoding"), "gzip") {
f, err := root.Open(upath + ".gz")
if err == nil {
f.Close()
c.Request.URL.Path += ".gz"
if strings.HasSuffix(upath, ".js") {
c.Writer.Header().Set("Content-Type", "application/javascript")
} else if strings.HasSuffix(upath, ".css") {
c.Writer.Header().Set("Content-Type", "text/css")
}
c.Writer.Header().Set("Content-Encoding", "gzip")
}
}
} else if upath != "/" {
f, err := root.Open(upath)
if err != nil {
c.Request.URL.Path = "/"
r.HandleContext(c)
return
}
defer f.Close()
}
fh.ServeHTTP(c.Writer, c.Request)
})
r.GET("/get/scriptInfo", func(c *gin.Context) {
// Get domain info
host := c.Request.Host
hostname, _, err := net.SplitHostPort(host)
if err != nil {
hostname = host // Use host directly if no port
}
chosen := hostname
// -------- Reverse proxy mode: force IP ----------
if cfg.ReverseProxyEnabled {
// Reverse proxy mode: always use configured WebRTC IP
if strings.TrimSpace(cfg.WebrtcIP) != "" {
chosen = strings.TrimSpace(cfg.WebrtcIP)
}
} else {
// -------- 3) Original behavior (unchanged) ----------
// 1) If hostname is domain, keep it
// 2) If hostname is IP and cfg.WebrtcIP is set, use cfg.WebrtcIP
if isIP(hostname) && cfg.WebrtcIP != "" {
chosen = cfg.WebrtcIP
}
}
c.JSON(http.StatusOK, gin.H{
"hostname": chosen, // reuse the same chosen value
"port": cfg.AddrDev,
"token": cfg.Token,
"webrtcIP": chosen, // same as hostname
"webrtcPort": cfg.WebrtcPort,
"webrtcUsername": cfg.WebrtcUsername,
"webrtcPassword": cfg.WebrtcPassword,
})
})
ln, err := net.Listen("tcp", cfg.AddrUser)
if err != nil {
return err
}
defer ln.Close()
// If we're behind a reverse proxy (TLS terminated by nginx), never enable TLS here.
enableTLS := !cfg.ReverseProxyEnabled && cfg.SslCert != "" && cfg.SslKey != ""
if enableTLS {
crt, err := tls.LoadX509KeyPair(cfg.SslCert, cfg.SslKey)
if err != nil {
log.Fatal().Msg(err.Error())
}
tlsConfig := &tls.Config{Certificates: []tls.Certificate{crt}}
ln = tls.NewListener(ln, tlsConfig)
}
log.Info().Msgf("Listen users on: %s", ln.Addr().(*net.TCPAddr))
return r.RunListener(ln)
}
func isIP(addr string) bool {
return net.ParseIP(addr) != nil
}
func callUserHookUrl(cfg *Config, c *gin.Context) bool {
if cfg.UserHookUrl == "" {
return true
}
upath := c.Request.URL.RawPath
// Create HTTP request with original headers
req, err := http.NewRequest("GET", cfg.UserHookUrl, nil)
if err != nil {
log.Error().Err(err).Msgf("create hook request for \"%s\" fail", upath)
return false
}
// Copy all headers from original request
for key, values := range c.Request.Header {
lowerKey := strings.ToLower(key)
if lowerKey == "upgrade" || lowerKey == "connection" || lowerKey == "accept-encoding" {
continue
}
for _, value := range values {
req.Header.Add(key, value)
}
}
// Add custom headers for hook identification
req.Header.Set("X-Rttys-Hook", "true")
req.Header.Set("X-Original-Method", c.Request.Method)
req.Header.Set("X-Original-URL", c.Request.URL.String())
cli := &http.Client{
Timeout: 3 * time.Second,
}
resp, err := cli.Do(req)
if err != nil {
log.Error().Err(err).Msgf("call user hook url for \"%s\" fail", upath)
return false
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
log.Error().Msgf("call user hook url for \"%s\", StatusCode: %d", upath, resp.StatusCode)
return false
}
return true
}
func httpLogin(cfg *Config, password string) bool {
return cfg.Password == password
}
func isLocalRequest(c *gin.Context) bool {
addr, _ := net.ResolveTCPAddr("tcp", c.Request.RemoteAddr)
return addr.IP.IsLoopback()
}
func httpAuth(cfg *Config, c *gin.Context) bool {
if !cfg.LocalAuth && isLocalRequest(c) {
return true
}
if cfg.Password == "" {
return true
}
sid, err := c.Cookie("sid")
if err != nil || !httpSessions.Exists(sid) {
return false
}
httpSessions.Expire(sid, httpSessionExpire)
return true
}
-88
View File
@@ -1,88 +0,0 @@
#!/bin/bash
# RTTYS Ubuntu Package Build Script
set -e
# Configuration
PACKAGE_NAME="rttys"
VERSION=$(grep 'const RttysVersion' main.go | cut -d'"' -f2 | sed 's/^v//')
MAINTAINER="Jianhui Zhao <zhaojh329@gmail.com>"
DESCRIPTION="Access your device's terminal from anywhere via the web"
URL="https://github.com/zhaojh329/rttys"
GitCommit=$(git log --pretty=format:"%h" -1)
BuildTime=$(date +%FT%T%z)
ARCH="$1" # Pass architecture as an argument, e.g., amd64 or arm64
[ -z "$ARCH" ] && {
echo "Usage: $0 <arch>";
echo "Example: $0 amd64"
exit 1;
}
# Build directory
BUILD_DIR="build-deb"
INSTALL_DIR="$BUILD_DIR/usr"
echo "Building RTTYS v$VERSION for Ubuntu..."
# Clean and create build directory
rm -rf $BUILD_DIR
mkdir -p $BUILD_DIR/{usr/bin,etc/rttys,lib/systemd/system}
# Build the binary
echo "Building binary..."
CGO_ENABLED=0 GOOS=linux GOARCH=$ARCH go build -ldflags "-s -w -X main.GitCommit=$GitCommit -X main.BuildTime=$BuildTime" -o $INSTALL_DIR/bin/rttys .
# Copy configuration files
echo "Copying configuration files..."
cp rttys.conf $BUILD_DIR/etc/rttys/
cp rttys.service $BUILD_DIR/lib/systemd/system/
# Create postinstall script
cat > $BUILD_DIR/postinstall.sh << 'EOF'
#!/bin/bash
# Enable and start the service
systemctl daemon-reload
systemctl enable rttys
echo "RTTYS installed successfully!"
echo "Configuration file: /etc/rttys/rttys.conf"
echo "Start service: sudo systemctl start rttys"
echo "View logs: sudo journalctl -u rttys -f"
EOF
# Create preremove script
cat > $BUILD_DIR/preremove.sh << 'EOF'
#!/bin/bash
# Stop and disable the service
systemctl stop rttys || true
systemctl disable rttys || true
EOF
# Create postremove script
cat > $BUILD_DIR/postremove.sh << 'EOF'
#!/bin/bash
EOF
# Build the package
echo "Creating .deb package..."
fpm -s dir -t deb \
--name "$PACKAGE_NAME" \
--version "$VERSION" \
--maintainer "$MAINTAINER" \
--description "$DESCRIPTION" \
--url "$URL" \
--architecture "$ARCH" \
--depends "systemd" \
--deb-no-default-config-files \
--config-files "/etc/rttys/rttys.conf" \
--after-install "$BUILD_DIR/postinstall.sh" \
--before-remove "$BUILD_DIR/preremove.sh" \
--after-remove "$BUILD_DIR/postremove.sh" \
-C "$BUILD_DIR" \
.
echo "Package created: ${PACKAGE_NAME}_${VERSION}_${ARCH}.deb"
echo "Install with: sudo dpkg -i ${PACKAGE_NAME}_${VERSION}_${ARCH}.deb"
-37
View File
@@ -1,37 +0,0 @@
#!/bin/sh
VERSION=$(grep 'const RttysVersion' main.go | cut -d'"' -f2 | sed 's/^v//')
GitCommit=$(git log --pretty=format:"%h" -1)
BuildTime=$(date +%FT%T%z)
[ $# -lt 2 ] && {
echo "Usage: $0 linux amd64"
exit 1
}
generate() {
local os="$1"
local arch="$2"
local dir="rttys-$VERSION-$os-$arch"
local bin="rttys"
rm -rf $dir
mkdir $dir
cp rttys.conf $dir
[ "$os" = "windows" ] && {
bin="rttys.exe"
}
GOOS=$os GOARCH=$arch CGO_ENABLED=0 go build -ldflags="-s -w -X main.GitCommit=$GitCommit -X main.BuildTime=$BuildTime" -o $dir/$bin && cp rttys.service $dir
[ -n "$COMPRESS" ] && {
tar -jcvf $dir.tar.bz2 $dir
rm -rf $dir
}
exit 0
}
generate $1 $2
+14
View File
@@ -0,0 +1,14 @@
package main
import (
_ "net/http/pprof"
"rttys/internal/server"
"github.com/rs/zerolog/log"
)
func main() {
if err := server.RunFromEnv(); err != nil {
log.Fatal().Msg(err.Error())
}
}
-16
View File
@@ -1,16 +0,0 @@
package db
// DeviceMeta represents a record in the gl_device table.
type DeviceMeta struct {
DeviceID string `gorm:"primaryKey;column:device_id"` // DeviceID is the globally unique and immutable ID of the device.
Mac string `gorm:"uniqueIndex;column:mac"` // Mac is the unique and immutable MAC address of the device.
IP string `gorm:"column:ip"` // IP is the current IP address of the device.
Description string `gorm:"column:description"` // Description is a human-readable description of the device.
CreateTime int64 `gorm:"column:create_time"` // CreateTime is the creation timestamp (Unix time).
UpdateTime int64 `gorm:"column:update_time"` // UpdateTime is the last update timestamp (Unix time).
}
// TableName sets the name of the table in the database that this struct binds to.
func (DeviceMeta) TableName() string {
return "devices"
}
-41
View File
@@ -1,41 +0,0 @@
package db
import (
"github.com/glebarez/sqlite"
"github.com/rs/zerolog/log"
"gorm.io/gorm"
)
const dbFileName = "/home/database/glkvm-cloud.db"
var deviceDB *gorm.DB
// GetDbClient returns the database client instance.
func GetDbClient() *gorm.DB {
return deviceDB
}
// Init initializes the SQLite database connection and sets up logging.
func Init() {
// Open a SQLite database connection
db, err := gorm.Open(sqlite.Open(dbFileName), &gorm.Config{})
if err != nil {
log.Info().Msg(err.Error())
// Panic if the database connection fails
panic("failed to connect database")
}
// Set the global database client
deviceDB = db
// Auto-migrate the Device schema
err = db.AutoMigrate(&DeviceMeta{})
if err != nil {
// Panic if auto-migration fails
panic(err)
}
// Retrieve and log the initial data records
list := make([]DeviceMeta, 0)
db.Find(&list)
log.Info().Msgf("==== SQLite init done ====, data record:%d \n", len(list))
}
-633
View File
@@ -1,633 +0,0 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
import (
"bufio"
"bytes"
"context"
"crypto/tls"
"encoding/binary"
"fmt"
"io"
"net"
"net/http"
"strings"
"sync"
"time"
"rttys/utils"
"github.com/gorilla/websocket"
jsoniter "github.com/json-iterator/go"
"github.com/rs/zerolog/log"
"github.com/valyala/bytebufferpool"
)
type DeviceInfo struct {
ID string `json:"id"`
Mac string `json:"mac"`
Connected uint32 `json:"connected"`
Uptime uint32 `json:"uptime"`
Desc string `json:"description"`
Proto uint8 `json:"proto"`
IPaddr string `json:"ipaddr"`
}
type Device struct {
group string
id string
proto uint8
desc string
timestamp int64
uptime uint32
token string
heartbeat time.Duration
users sync.Map
pending sync.Map
commands sync.Map
https sync.Map
conn net.Conn
br *bufio.Reader
readBuf []byte
close sync.Once
ctx context.Context
cancel context.CancelFunc
}
const (
msgTypeRegister = byte(iota)
msgTypeLogin
msgTypeLogout
msgTypeTermData
msgTypeWinsize
msgTypeCmd
msgTypeHeartbeat
msgTypeFile
msgTypeHttp
msgTypeAck
)
const (
msgTypeFileSend = byte(iota)
msgTypeFileRecv
msgTypeFileInfo
msgTypeFileData
msgTypeFileAck
msgTypeFileAbort
)
const (
msgRegAttrHeartbeat = iota
msgRegAttrDevid
msgRegAttrDescription
msgRegAttrToken
msgRegAttrGroup
)
const (
msgHeartbeatAttrUptime = iota
)
const (
devRegErrUnsupportedProto = iota + 1
devRegErrInvalidToken
devRegErrHookFailed
devRegErrIdConflicting
)
const (
RttyProtoRequired uint8 = 3
WaitRegistTimeout = 5 * time.Second
DefaultHeartbeat = 5 * time.Second
TermLoginTimeout = 5 * time.Second
CommandTimeout = 30
)
var DevRegErrMsg = map[byte]string{
0: "Success",
devRegErrUnsupportedProto: "Unsupported protocol",
devRegErrInvalidToken: "Invalid token",
devRegErrHookFailed: "Hook failed",
devRegErrIdConflicting: "ID conflict",
}
var DeviceMsgHandlers = map[byte]func(*Device, []byte) error{
msgTypeHeartbeat: handleHeartbeatMsg,
msgTypeLogin: handleLoginMsg,
msgTypeLogout: handleLogoutMsg,
msgTypeTermData: handleTermDataMsg,
msgTypeFile: handleFileMsg,
msgTypeCmd: handleCmdMsg,
msgTypeHttp: handleHttpMsg,
}
func (srv *RttyServer) ListenDevices() {
cfg := &srv.cfg
ln, err := net.Listen("tcp", cfg.AddrDev)
if err != nil {
log.Fatal().Msg(err.Error())
}
defer ln.Close()
if cfg.SslCert != "" && cfg.SslKey != "" {
crt, err := tls.LoadX509KeyPair(cfg.SslCert, cfg.SslKey)
if err != nil {
log.Fatal().Msg(err.Error())
}
tlsConfig := &tls.Config{
GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) {
// 忽略 SNI,始终返回唯一证书
return &crt, nil
},
}
ln = tls.NewListener(ln, tlsConfig)
}
log.Info().Msgf("Listen devices on: %s", ln.Addr().(*net.TCPAddr))
for {
conn, err := ln.Accept()
if err != nil {
log.Error().Msg(err.Error())
continue
}
go handleDeviceConnection(srv, conn)
}
}
func handleDeviceConnection(srv *RttyServer, conn net.Conn) {
defer logPanic()
dev := &Device{
conn: conn,
heartbeat: DefaultHeartbeat,
timestamp: time.Now().Unix(),
br: bufio.NewReader(conn),
}
defer dev.Close(srv)
dev.ctx, dev.cancel = context.WithCancel(context.Background())
log.Debug().Msgf("new device '%s' connected", conn.RemoteAddr())
conn.SetReadDeadline(time.Now().Add(WaitRegistTimeout))
typ, data, err := dev.ReadMsg()
if err != nil {
log.Error().Msgf("read register msg fail: %v", err)
return
}
if typ != msgTypeRegister {
log.Error().Msg("register msg expected first")
return
}
if !dev.ParseRegister(data) {
log.Error().Msg("invalid device info")
return
}
code := dev.Register(srv)
err = dev.WriteMsg(msgTypeRegister, "", append([]byte{code}, DevRegErrMsg[code]...))
if err != nil {
log.Printf("send register to device '%s' fail: %v", dev.id, err)
return
}
if code != 0 {
return
}
deviceRemoteIP := ""
if addr, ok := dev.conn.RemoteAddr().(*net.TCPAddr); ok {
deviceRemoteIP = addr.IP.String()
} else if host, _, err := net.SplitHostPort(dev.conn.RemoteAddr().String()); err == nil {
deviceRemoteIP = host
}
log.Info().Msgf("device '%s' registered, group '%s' proto %d, heartbeat %v, remoteIP '%s'",
dev.id, dev.group, dev.proto, dev.heartbeat, deviceRemoteIP)
// 2. Load existing metadata by device_id
description := ""
meta, err := GetDeviceMetaByDeviceID(dev.id)
if err == nil && meta != nil {
description = meta.Description
}
if err := SaveOrUpdateDeviceMeta(
dev.id,
dev.desc, // device register mac info with desc filed
description,
deviceRemoteIP,
); err != nil {
return
}
for {
conn.SetReadDeadline(time.Now().Add(dev.heartbeat * 3 / 2))
typ, data, err = dev.ReadMsg()
if err != nil {
if err != io.EOF {
log.Error().Msgf("read msg from device '%s' fail: %v", dev.id, err)
}
return
}
log.Debug().Msgf("device msg %s from device %s", msgTypeName(typ), dev.id)
handler, ok := DeviceMsgHandlers[typ]
if !ok {
log.Error().Msgf("unexpected message '%s' from device '%s'", msgTypeName(typ), dev.id)
return
}
err = handler(dev, data)
if err != nil {
log.Error().Msg(err.Error())
return
}
}
}
func msgTypeName(typ byte) string {
switch typ {
case msgTypeRegister:
return "register"
case msgTypeLogin:
return "login"
case msgTypeLogout:
return "logout"
case msgTypeTermData:
return "termdata"
case msgTypeWinsize:
return "winsize"
case msgTypeCmd:
return "cmd"
case msgTypeHeartbeat:
return "heartbeat"
case msgTypeFile:
return "file"
case msgTypeHttp:
return "http"
case msgTypeAck:
return "ack"
default:
return fmt.Sprintf("unknown(%d)", typ)
}
}
func (dev *Device) ReadMsg() (byte, []byte, error) {
head := make([]byte, 3)
br := dev.br
_, err := io.ReadFull(br, head)
if err != nil {
return 0, nil, err
}
typ := head[0]
msgLen := binary.BigEndian.Uint16(head[1:])
if cap(dev.readBuf) < int(msgLen) {
dev.readBuf = make([]byte, msgLen)
} else {
dev.readBuf = dev.readBuf[:msgLen]
}
_, err = io.ReadFull(br, dev.readBuf)
if err != nil {
return 0, nil, err
}
return typ, dev.readBuf, nil
}
func (dev *Device) WriteMsg(typ byte, sid string, data []byte) error {
bb := bytebufferpool.Get()
defer bytebufferpool.Put(bb)
b := []byte{typ, 0, 0}
binary.BigEndian.PutUint16(b[1:], uint16(len(sid)+len(data)))
bb.Write(b)
bb.WriteString(sid)
bb.Write(data)
_, err := bb.WriteTo(dev.conn)
return err
}
func (dev *Device) WriteFileMsg(typ byte, sid string, fileType byte, data []byte) error {
bb := bytebufferpool.Get()
defer bytebufferpool.Put(bb)
bb.WriteByte(fileType)
bb.Write(data)
return dev.WriteMsg(typ, sid, bb.Bytes())
}
func (dev *Device) Close(srv *RttyServer) {
dev.close.Do(func() {
log.Error().Msgf("device '%s' disconnected", dev.id)
srv.DelDevice(dev)
dev.cancel()
dev.conn.Close()
})
}
func (dev *Device) ParseRegister(b []byte) bool {
if len(b) < 1 {
return false
}
dev.proto = b[0]
if dev.proto > 4 {
attrs := utils.ParseTLV(b[1:])
if attrs == nil {
return false
}
for typ, val := range attrs {
switch typ {
case msgRegAttrHeartbeat:
dev.heartbeat = time.Duration(val[0]) * time.Second
case msgRegAttrDevid:
dev.id = string(val)
case msgRegAttrDescription:
dev.desc = string(val)
case msgRegAttrToken:
dev.token = string(val)
case msgRegAttrGroup:
dev.group = string(val)
}
}
return true
}
b = b[1:]
fields := bytes.Split(b, []byte{0})
if len(fields) < 3 {
return false
}
dev.id = string(fields[0])
dev.desc = string(fields[1])
dev.token = string(fields[2])
return true
}
func (dev *Device) Register(srv *RttyServer) byte {
cfg := &srv.cfg
if dev.proto < RttyProtoRequired {
log.Error().Msgf("minimum proto required %d, found %d for device '%s'", RttyProtoRequired, dev.proto, dev.id)
return devRegErrHookFailed
}
log.Info().Msgf("cfg.Token:%s,dev.token:%s", cfg.Token, dev.token)
if cfg.Token != "" && dev.token != cfg.Token {
log.Error().Msgf("invalid token for device '%s'", dev.id)
return devRegErrInvalidToken
}
devHookUrl := cfg.DevHookUrl
if devHookUrl != "" {
cli := &http.Client{
Timeout: 3 * time.Second,
}
data := fmt.Sprintf(`{"group":"%s", "devid":"%s", "token":"%s"}`, dev.group, dev.id, dev.token)
resp, err := cli.Post(devHookUrl, "application/json", strings.NewReader(data))
if err != nil {
log.Error().Msgf("call device hook url fail for device %s: %v", dev.id, err)
return devRegErrHookFailed
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
log.Error().Msgf("call device hook url for device '%s', StatusCode: %d", dev.id, resp.StatusCode)
return devRegErrHookFailed
}
}
if !srv.AddDevice(dev) {
return devRegErrIdConflicting
}
return 0
}
func handleHeartbeatMsg(dev *Device, data []byte) error {
if !parseHeartbeat(dev, data) {
return fmt.Errorf("invalid heartbeat msg from device '%s'", dev.id)
}
return dev.WriteMsg(msgTypeHeartbeat, "", nil)
}
func parseHeartbeat(dev *Device, data []byte) bool {
if dev.proto > 4 {
attrs := utils.ParseTLV(data)
if attrs == nil {
return false
}
for typ, val := range attrs {
switch typ {
case msgHeartbeatAttrUptime:
dev.uptime = binary.BigEndian.Uint32(val)
}
}
} else {
if len(data) < 4 {
return false
}
dev.uptime = binary.BigEndian.Uint32(data[:4])
}
return true
}
func handleLogoutMsg(dev *Device, data []byte) error {
if len(data) < 32 {
return fmt.Errorf("invalid logout msg from device '%s'", dev.id)
}
sid := string(data[:32])
if val, loaded := dev.users.LoadAndDelete(sid); loaded {
user := val.(*User)
user.Close()
}
return nil
}
func handleLoginMsg(dev *Device, data []byte) error {
if len(data) < 33 {
return fmt.Errorf("invalid login msg from device '%s'", dev.id)
}
sid := string(data[:32])
code := data[32]
if val, loaded := dev.pending.LoadAndDelete(sid); loaded {
user := val.(*User)
ok := code == 0
errCode := 0
if ok {
log.Debug().Msgf("login session '%s' for device '%s' success", sid, dev.id)
dev.users.Store(sid, user)
} else {
errCode = LoginErrorBusy
log.Error().Msgf("login session '%s' for device '%s' fail, due to device busy", sid, dev.id)
}
if errCode == 0 {
user.WriteMsg(websocket.TextMessage, []byte(fmt.Appendf(nil, `{"type":"login"}`)))
} else {
user.SendCloseMsg(LoginErrorBusy, "device busy")
}
user.pending <- ok
}
return nil
}
func handleTermDataMsg(dev *Device, data []byte) error {
if len(data) < 32 {
return fmt.Errorf("invalid term data msg from device '%s'", dev.id)
}
sid := string(data[:32])
if val, ok := dev.users.Load(sid); ok {
user := val.(*User)
data[31] = 0
user.WriteMsg(websocket.BinaryMessage, data[31:])
}
return nil
}
func handleFileMsg(dev *Device, data []byte) error {
if len(data) < 33 {
return fmt.Errorf("invalid file msg from device '%s'", dev.id)
}
sid := string(data[:32])
typ := data[32]
if val, ok := dev.users.Load(sid); ok {
user := val.(*User)
switch typ {
case msgTypeFileSend:
user.WriteMsg(websocket.TextMessage,
fmt.Appendf(nil, `{"type":"sendfile", "name": "%s"}`, string(data[33:])))
case msgTypeFileRecv:
user.WriteMsg(websocket.TextMessage, []byte(`{"type":"recvfile"}`))
case msgTypeFileData:
data[32] = 1
user.WriteMsg(websocket.BinaryMessage, data[32:])
case msgTypeFileAck:
user.WriteMsg(websocket.TextMessage, []byte(`{"type":"fileAck"}`))
case msgTypeFileAbort:
user.WriteMsg(websocket.BinaryMessage, []byte{1})
}
}
return nil
}
func handleHttpMsg(dev *Device, data []byte) error {
if len(data) < 18 {
return fmt.Errorf("invalid http msg from device '%s'", dev.id)
}
addr := data[:18]
data = data[18:]
if c, ok := dev.https.Load(string(addr)); ok {
c := c.(net.Conn)
if len(data) == 0 {
c.Close()
} else {
c.Write(data)
}
}
return nil
}
func handleCmdMsg(dev *Device, data []byte) error {
info := &CommandRespInfo{}
err := jsoniter.Unmarshal(data, info)
if err != nil {
return fmt.Errorf("parse command resp info error: %v", err)
}
var attrs map[string]any
err = jsoniter.Unmarshal(info.Attrs, &attrs)
if err != nil {
return fmt.Errorf("parse command resp attrs error: %v", err)
}
attrs["devid"] = dev.id
if val, ok := dev.commands.Load(info.Token); ok {
req := val.(*CommandReq)
req.acked = true
req.c.JSON(http.StatusOK, attrs)
req.cancel()
}
return nil
}
+3
View File
@@ -71,6 +71,9 @@ GLKVM_ACCESS_IP=
# rttys
RTTYS_TOKEN=DeviceTokenYouCanChangeMe
RTTYS_PASS=StrongP@ssw0rd
# Admin username (leave empty to default to "admin")
# Only letters and digits are allowed (e.g. admin, Admin01). No spaces or special characters.
RTTYS_ADMIN_NAME=
RTTYS_DEVICE_PORT=5912
RTTYS_WEBUI_PORT=443
RTTYS_HTTP_PROXY_PORT=10443
+3
View File
@@ -70,6 +70,9 @@ GLKVM_ACCESS_IP=
# rttys
RTTYS_TOKEN=DeviceTokenYouCanChangeMe
RTTYS_PASS=StrongP@ssw0rd
# Admin username (leave empty to default to "admin")
# Only letters and digits are allowed (e.g. admin, Admin01). No spaces or special characters.
RTTYS_ADMIN_NAME=
RTTYS_DEVICE_PORT=5912
RTTYS_WEBUI_PORT=443
RTTYS_HTTP_PROXY_PORT=10443
+100
View File
@@ -0,0 +1,100 @@
PRAGMA foreign_keys = ON;
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT NOT NULL UNIQUE,
email TEXT,
description TEXT NOT NULL DEFAULT '',
password_hash TEXT NOT NULL,
role TEXT NOT NULL CHECK (role IN ('admin','user')),
status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active','disabled')),
is_system INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE INDEX IF NOT EXISTS idx_users_role ON users(role);
CREATE INDEX IF NOT EXISTS idx_users_status ON users(status);
CREATE TABLE IF NOT EXISTS user_groups (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
description TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE TABLE IF NOT EXISTS user_group_members (
user_id INTEGER NOT NULL,
group_id INTEGER NOT NULL,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
PRIMARY KEY (user_id, group_id),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY (group_id) REFERENCES user_groups(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_ugm_group_id ON user_group_members(group_id);
CREATE TABLE IF NOT EXISTS device_groups (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
description TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE TABLE IF NOT EXISTS user_group_device_group_links (
user_group_id INTEGER NOT NULL,
device_group_id INTEGER NOT NULL,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
PRIMARY KEY (user_group_id, device_group_id),
FOREIGN KEY (user_group_id) REFERENCES user_groups(id) ON DELETE CASCADE,
FOREIGN KEY (device_group_id) REFERENCES device_groups(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_ug_dg_device_group_id
ON user_group_device_group_links(device_group_id);
CREATE TABLE IF NOT EXISTS devices (
id INTEGER PRIMARY KEY AUTOINCREMENT,
ddns TEXT NOT NULL UNIQUE,
mac TEXT NOT NULL UNIQUE,
name TEXT NOT NULL DEFAULT '',
description TEXT NOT NULL DEFAULT '',
ip TEXT NOT NULL DEFAULT '',
client TEXT NOT NULL DEFAULT '',
device_group_id INTEGER NULL,
status TEXT NOT NULL DEFAULT 'online' CHECK (status IN ('online','offline','disabled')),
last_seen_at INTEGER NULL,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch()),
FOREIGN KEY (device_group_id) REFERENCES device_groups(id) ON DELETE SET NULL
);
CREATE INDEX IF NOT EXISTS idx_devices_group_id ON devices(device_group_id);
CREATE INDEX IF NOT EXISTS idx_devices_status ON devices(status);
CREATE INDEX IF NOT EXISTS idx_devices_last_seen ON devices(last_seen_at);
CREATE TRIGGER IF NOT EXISTS trg_users_updated_at
AFTER UPDATE ON users
FOR EACH ROW
BEGIN
UPDATE users SET updated_at = unixepoch() WHERE id = OLD.id;
END;
CREATE TRIGGER IF NOT EXISTS trg_user_groups_updated_at
AFTER UPDATE ON user_groups
FOR EACH ROW
BEGIN
UPDATE user_groups SET updated_at = unixepoch() WHERE id = OLD.id;
END;
CREATE TRIGGER IF NOT EXISTS trg_device_groups_updated_at
AFTER UPDATE ON device_groups
FOR EACH ROW
BEGIN
UPDATE device_groups SET updated_at = unixepoch() WHERE id = OLD.id;
END;
CREATE TRIGGER IF NOT EXISTS trg_devices_updated_at
AFTER UPDATE ON devices
FOR EACH ROW
BEGIN
UPDATE devices SET updated_at = unixepoch() WHERE id = OLD.id;
END;
+9 -3
View File
@@ -5,7 +5,6 @@ services:
image: ${GLKVM_IMAGE:-glzhitong/glkvm-cloud:latest}
container_name: glkvm_cloud
restart: always
network_mode: "host"
environment:
# Preferred: set GLKVM_ACCESS_IP explicitly; if empty, entrypoint will auto-detect once.
GLKVM_ACCESS_IP: ${GLKVM_ACCESS_IP:-}
@@ -13,6 +12,7 @@ services:
# ---- rttys ----
RTTYS_TOKEN: ${RTTYS_TOKEN:-DeviceTokenYouCanChangeMe}
RTTYS_PASS: ${RTTYS_PASS:-StrongP@ssw0rd}
RTTYS_ADMIN_NAME: ${RTTYS_ADMIN_NAME:-}
# Ports inside container (mirrored to host via `ports` below)
RTTYS_DEVICE_PORT: ${RTTYS_DEVICE_PORT:-5912} # addr-dev
@@ -66,12 +66,15 @@ services:
- ./database:/home/database:rw
entrypoint: ["/bin/sh", "/docker-entrypoint.sh"]
command: ["rttys"]
ports:
- "${RTTYS_WEBUI_PORT:-443}:${RTTYS_WEBUI_PORT:-443}"
- "${RTTYS_HTTP_PROXY_PORT:-10443}:${RTTYS_HTTP_PROXY_PORT:-10443}"
- "${RTTYS_DEVICE_PORT:-5912}:${RTTYS_DEVICE_PORT:-5912}"
coturn:
image: ${COTURN_IMAGE:-coturn/coturn:edge-alpine}
container_name: glkvm_coturn
restart: always
network_mode: "host"
environment:
# Same semantics as above: prefer explicit value, else auto-detect
GLKVM_ACCESS_IP: ${GLKVM_ACCESS_IP:-}
@@ -82,4 +85,7 @@ services:
command: ["coturn"]
volumes:
- ./templates/turnserver.conf.template:/tpl/turnserver.conf.tmpl:ro
- ./scripts/docker-entrypoint.sh:/docker-entrypoint.sh:ro
- ./scripts/docker-entrypoint.sh:/docker-entrypoint.sh:ro
ports:
- "${TURN_PORT:-3478}:3478/tcp"
- "${TURN_PORT:-3478}:3478/udp"
+1 -1
View File
@@ -60,7 +60,7 @@ case "$1" in
: "${TURN_PORT:=3478}"
render /tpl/rttys.conf.tmpl /home/rttys.conf \
RTTYS_TOKEN RTTYS_PASS \
RTTYS_TOKEN RTTYS_PASS RTTYS_ADMIN_NAME \
GLKVM_ACCESS_IP TURN_PORT TURN_USER TURN_PASS \
RTTYS_DEVICE_PORT RTTYS_WEBUI_PORT RTTYS_HTTP_PROXY_PORT \
LDAP_ENABLED LDAP_SERVER LDAP_PORT LDAP_USE_TLS \
@@ -4,6 +4,9 @@ token: {{RTTYS_TOKEN}}
# Web management password
password: {{RTTYS_PASS}}
# Admin username (leave empty to default to "admin")
admin-name: {{RTTYS_ADMIN_NAME}}
# WebRTC
webrtc-ip: {{GLKVM_ACCESS_IP}}
webrtc-port: {{TURN_PORT}}
-739
View File
@@ -1,739 +0,0 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
import (
"bufio"
"context"
"crypto/tls"
"encoding/binary"
"errors"
"fmt"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"rttys/utils"
"github.com/gin-gonic/gin"
"github.com/rs/zerolog/log"
"github.com/valyala/bytebufferpool"
)
type HttpProxySession struct {
expire atomic.Int64
ctx context.Context
cancel context.CancelFunc
devid string
group string
destaddr string
https bool
}
var httpProxySessions = sync.Map{}
const httpProxySessionsExpire = 15 * time.Minute
func (ses *HttpProxySession) Expire() {
ses.expire.Store(time.Now().Add(httpProxySessionsExpire).Unix())
}
func (ses *HttpProxySession) String() string {
return fmt.Sprintf("{devid: %s, group: %s, destaddr: %s, https: %v}",
ses.devid, ses.group, ses.destaddr, ses.https)
}
func (srv *RttyServer) ListenHttpProxy() {
cfg := &srv.cfg
if cfg.AddrHttpProxy != "" {
addr, err := net.ResolveTCPAddr("tcp", cfg.AddrHttpProxy)
if err != nil {
log.Warn().Msg("invalid http proxy addr: " + err.Error())
} else {
srv.httpProxyPort = addr.Port
}
}
ln, err := net.Listen("tcp", cfg.AddrHttpProxy)
if err != nil {
log.Fatal().Msg(err.Error())
}
defer ln.Close()
// In reverse proxy mode (TLS terminated by nginx), never enable TLS here.
enableTLS := !cfg.ReverseProxyEnabled && cfg.SslCert != "" && cfg.SslKey != ""
if enableTLS {
crt, err := tls.LoadX509KeyPair(cfg.SslCert, cfg.SslKey)
if err != nil {
log.Fatal().Msg(err.Error())
}
tlsConfig := &tls.Config{Certificates: []tls.Certificate{crt}}
ln = tls.NewListener(ln, tlsConfig)
}
srv.httpProxyPort = ln.Addr().(*net.TCPAddr).Port
log.Info().Msgf("Listen http proxy on: %s", ln.Addr().(*net.TCPAddr))
go httpProxySessionsClean()
for {
c, err := ln.Accept()
if err != nil {
log.Error().Msg(err.Error())
continue
}
go doHttpProxy(srv, c)
}
}
func httpProxySessionsClean() {
for {
time.Sleep(time.Second * 30)
httpProxySessions.Range(func(key, value any) bool {
ses := value.(*HttpProxySession)
if time.Now().Unix() > ses.expire.Load() {
log.Debug().Msgf("Http proxy session '%s' expired", key)
ses.cancel()
httpProxySessions.Delete(key)
}
return true
})
}
}
func doHttpProxy(srv *RttyServer, c net.Conn) {
defer logPanic()
defer c.Close()
br := bufio.NewReader(c)
req, err := http.ReadRequest(br)
if err != nil {
return
}
domain, port, proto := getRequestHostInfo(req)
log.Debug().Msgf("http proxy incoming host=%s port=%s proto=%s uri=%s",
domain, port, proto, req.URL.String())
devID, ok := extractDeviceIDFromHost(domain)
if ok {
log.Debug().Msgf("parsed deviceId from host: %s", devID)
} else {
log.Debug().Msgf("host is IP or invalid, skip deviceId parsing")
}
// 获取 URL 查询参数
queryParams := req.URL.Query()
name := queryParams.Get("sid")
if name != "" {
location := "/"
location += fmt.Sprintf("?_=%d", time.Now().Unix())
Write302WithCookie(c, location, "rtty-http-sid", name)
return
}
cookie, err := req.Cookie("rtty-http-sid")
if err != nil {
log.Debug().Msgf(`not found cookie "rtty-http-sid"`)
sendHTTPErrorResponse(c, "invalid")
return
}
sid := cookie.Value
sesVal, ok := httpProxySessions.Load(sid)
if !ok {
log.Debug().Msgf(`not found httpProxySession "%s"`, sid)
sendHTTPErrorResponse(c, "unauthorized")
return
}
ses := sesVal.(*HttpProxySession)
dev := srv.GetDevice(ses.group, ses.devid)
if dev == nil {
log.Debug().Msgf(`device "%s" group "%s" offline`, ses.devid, ses.group)
sendHTTPErrorResponse(c, "offline")
return
}
// 3) match hostDevID vs session devid, and optionally lookup by hostDevID
if devID != "" {
match := devID == ses.devid
log.Debug().Msgf(
"http proxy devid check: hostDevID=%s sessionDevid=%s match=%v hostDevFound=%v sid=%s group=%s",
devID, ses.devid, match, domain, sid, ses.group,
)
// If you want, you can also log when mismatch happens
if !match {
log.Info().Msgf(
"http proxy devid mismatch: hostDevID=%s sessionDevid=%s sid=%s group=%s host=%s uri=%s",
devID, ses.devid, sid, ses.group, domain, req.URL.String(),
)
sendHTTPErrorResponse(c, "invalid")
}
} else {
log.Debug().Msgf(
"http proxy devid check skipped: no hostDevID (host=%s) sid=%s group=%s sessionDevid=%s",
domain, sid, ses.group, ses.devid,
)
}
hostHeaderRewrite := ses.destaddr
destAddr := genDestAddr(hostHeaderRewrite)
srcAddr := tcpAddr2Bytes(c.RemoteAddr().(*net.TCPAddr))
ctx, cancel := context.WithCancel(ses.ctx)
defer cancel()
go func() {
<-ctx.Done()
c.Close()
log.Debug().Msgf("http proxy conn closed: %s", ses)
dev.https.Delete(string(srcAddr))
sendHttpReq(dev, ses.https, srcAddr[:], destAddr, nil)
}()
log.Debug().Msgf("new http proxy conn: %s", ses)
dev.https.Store(string(srcAddr), c)
hpw := &HttpProxyWriter{destAddr, srcAddr, hostHeaderRewrite, dev, ses.https}
req.Host = hostHeaderRewrite
hpw.WriteRequest(req)
if req.Header.Get("Upgrade") == "websocket" {
b := make([]byte, 4096)
for {
n, err := c.Read(b)
if err != nil {
return
}
sendHttpReq(dev, ses.https, srcAddr, destAddr, b[:n])
ses.Expire()
}
} else {
for {
req, err := http.ReadRequest(br)
if err != nil {
return
}
hpw.WriteRequest(req)
ses.Expire()
}
}
}
func httpProxyRedirect(srv *RttyServer, c *gin.Context, group string) {
cfg := &srv.cfg
devid := c.Param("devid")
proto := c.Param("proto")
addr := c.Param("addr")
rawPath := c.Param("path")
log.Info().Msgf("httpProxyRedirect devid: %s, proto: %s, addr: %s, path: %s", devid, proto, addr, rawPath)
if !callUserHookUrl(cfg, c) {
c.Status(http.StatusForbidden)
return
}
log.Debug().Msgf("httpProxyRedirect devid: %s, proto: %s, addr: %s, path: %s", devid, proto, addr, rawPath)
_, _, err := httpProxyVaildAddr(addr)
if err != nil {
log.Debug().Msgf("invalid addr: %s", addr)
c.Status(http.StatusBadRequest)
return
}
path, err := url.Parse(rawPath)
if err != nil {
log.Debug().Msgf("invalid path: %s", rawPath)
c.Status(http.StatusBadRequest)
return
}
dev := srv.GetDevice(group, devid)
if dev == nil {
c.Redirect(http.StatusFound, "/error/offline")
return
}
location := c.Request.Header.Get("HttpProxyRedir")
log.Info().Msgf("HttpProxyRedir location: %s, devid: %s", location, devid)
if location == "" {
location = cfg.HttpProxyRedirURL
if location != "" {
log.Debug().Msgf("use HttpProxyRedirURL from config: %s, devid: %s", location, devid)
}
} else {
log.Debug().Msgf("use HttpProxyRedir from HTTP header: %s, devid: %s", location, devid)
}
if location == "" {
host, _, err := net.SplitHostPort(c.Request.Host)
if err != nil {
host = c.Request.Host
}
location = "http://" + host
if srv.httpProxyPort != 80 {
location += fmt.Sprintf(":%d", srv.httpProxyPort)
}
}
location += path.Path
location += fmt.Sprintf("?_=%d", time.Now().Unix())
if path.RawQuery != "" {
location += "&" + path.RawQuery
}
sid, err := c.Cookie("rtty-http-sid")
log.Info().Msgf("rtty-http-sid: %s", sid)
if err == nil {
if v, loaded := httpProxySessions.LoadAndDelete(sid); loaded {
s := v.(*HttpProxySession)
s.cancel()
log.Debug().Msgf(`del old httpProxySession "%s" for device "%s"`, sid, devid)
}
}
sid = utils.GenUniqueID()
log.Info().Msgf("rtty-http-sid: %s", sid)
ctx, cancel := context.WithCancel(dev.ctx)
ses := &HttpProxySession{
ctx: ctx,
cancel: cancel,
devid: devid,
group: group,
destaddr: addr,
https: proto == "https",
}
ses.Expire()
httpProxySessions.Store(sid, ses)
log.Debug().Msgf(`new httpProxySession "%s" for device "%s"`, sid, devid)
domain := c.Request.Header.Get("HttpProxyRedirDomain")
if domain == "" {
domain = cfg.HttpProxyRedirDomain
if domain != "" {
log.Debug().Msgf("set cookie domain from config: %s, devid: %s", domain, devid)
}
} else {
log.Debug().Msgf("set cookie domain from HTTP header: %s, devid: %s", domain, devid)
}
// Get domain info
host := c.Request.Host
hostname, _, err := net.SplitHostPort(host)
if err != nil {
hostname = host
}
log.Info().Msgf("hostname: %s", hostname)
ip := net.ParseIP(hostname)
isIP := ip != nil
if isIP {
location = fmt.Sprintf("https://%s%s?sid=%s", hostname, cfg.AddrHttpProxy, sid)
log.Info().Msgf("Using IP redirect: %s", location)
} else {
redirHost := buildRedirectHost(hostname, devid)
// Keep original behavior when NOT in reverse proxy mode
if !cfg.ReverseProxyEnabled {
location = fmt.Sprintf("https://%s%s?sid=%s", redirHost, cfg.AddrHttpProxy, sid)
log.Info().Msgf("Using domain redirect: %s", location)
} else {
// ---- verify forwarded headers from reverse proxy ----
rawHost := c.GetHeader("Host")
xfHost := c.GetHeader("X-Forwarded-Host")
xfProto := c.GetHeader("X-Forwarded-Proto")
xfPort := c.GetHeader("X-Forwarded-Port")
xRealIP := c.GetHeader("X-Real-IP")
xFF := c.GetHeader("X-Forwarded-For")
log.Info().Msgf(
"reverse-proxy info: method=%s uri=%s host=%q tls=%v remoteIP=%q",
c.Request.Method,
c.Request.URL.String(),
rawHost,
c.Request.TLS != nil,
c.ClientIP(),
)
log.Info().Msgf(
"reverse-proxy headers: Host=%q X-Forwarded-Host=%q X-Forwarded-Proto=%q X-Forwarded-Port=%q X-Real-IP=%q X-Forwarded-For=%q",
rawHost, xfHost, xfProto, xfPort, xRealIP, xFF,
)
// -------------------------------------------------
// Proxy mode:
// 1) If DEVICE_ENDPOINT_HOST is configured, use it directly
// 2) Otherwise, fallback to forwarded-header logic
// -------------------------------------------------
// 0) scheme: follow reverse proxy
scheme := ""
if v := strings.TrimSpace(c.GetHeader("X-Forwarded-Proto")); v != "" {
scheme = strings.ToLower(strings.Split(v, ",")[0])
} else if c.Request.TLS != nil {
scheme = "https"
} else {
scheme = "http"
}
// [A] Prefer explicit DEVICE_ENDPOINT_HOST if set
if v := strings.TrimSpace(cfg.DeviceEndpointHost); v != "" {
endpoint := v // already normalized when reading env: host[:port] only
baseHost := endpoint
port := ""
if h, p, err := net.SplitHostPort(endpoint); err == nil {
baseHost = h
port = p
}
// Build device host: <deviceId>.<baseHost>
// NOTE: DEVICE_ENDPOINT_HOST is a base domain (host[:port]) for device access,
baseHost = strings.TrimSuffix(strings.TrimSpace(baseHost), ".")
deviceHost := devid
if baseHost != "" {
deviceHost = devid + "." + baseHost
}
hostPort := joinHostPortIfNeeded(deviceHost, scheme, port)
redirectPath := c.Request.URL.Path
location = buildRedirectLocation(scheme, hostPort, redirectPath, sid)
log.Info().Msgf("Using domain redirect (proxy mode, DEVICE_ENDPOINT_HOST): %s", location)
} else {
// 1) external port: prefer the one user actually accessed
port := ""
if fp := strings.TrimSpace(c.GetHeader("X-Forwarded-Port")); fp != "" {
port = strings.TrimSpace(strings.Split(fp, ",")[0])
} else if fh := strings.TrimSpace(c.GetHeader("X-Forwarded-Host")); fh != "" {
fh = strings.TrimSpace(strings.Split(fh, ",")[0])
if _, p, err := net.SplitHostPort(fh); err == nil && p != "" {
port = p
}
}
log.Info().Msgf("port: %s", port)
// 3) Build host: in proxy mode redirect domain to be redirHost
hostPort := joinHostPortIfNeeded(redirHost, scheme, port)
redirectPath := c.Request.URL.Path
location = buildRedirectLocation(scheme, hostPort, redirectPath, sid)
log.Info().Msgf("Using domain redirect (proxy mode): %s", location)
}
}
}
log.Info().Msgf("Final redirect location: %s", location)
c.Redirect(http.StatusFound, location)
}
func sendHttpReq(dev *Device, https bool, srcAddr []byte, destAddr []byte, data []byte) {
bb := bytebufferpool.Get()
defer bytebufferpool.Put(bb)
if dev.proto > 3 {
if https {
bb.WriteByte(1)
} else {
bb.WriteByte(0)
}
}
bb.Write(srcAddr)
bb.Write(destAddr)
bb.Write(data)
dev.WriteMsg(msgTypeHttp, "", bb.Bytes())
}
func genDestAddr(addr string) []byte {
destIP, destPort, err := httpProxyVaildAddr(addr)
if err != nil {
return nil
}
b := make([]byte, 6)
copy(b, destIP)
binary.BigEndian.PutUint16(b[4:], destPort)
return b
}
func tcpAddr2Bytes(addr *net.TCPAddr) []byte {
b := make([]byte, 18)
binary.BigEndian.PutUint16(b[:2], uint16(addr.Port))
copy(b[2:], addr.IP)
return b
}
func httpProxyVaildAddr(addr string) (net.IP, uint16, error) {
ips, ports, err := net.SplitHostPort(addr)
if err != nil {
ips = addr
ports = "80"
}
ip := net.ParseIP(ips)
if ip == nil {
return nil, 0, errors.New("invalid IPv4 Addr")
}
ip = ip.To4()
if ip == nil {
return nil, 0, errors.New("invalid IPv4 Addr")
}
port, _ := strconv.Atoi(ports)
return ip, uint16(port), nil
}
type HttpProxyWriter struct {
destAddr []byte
srcAddr []byte
hostHeaderRewrite string
dev *Device
https bool
}
func (rw *HttpProxyWriter) Write(p []byte) (n int, err error) {
sendHttpReq(rw.dev, rw.https, rw.srcAddr, rw.destAddr, p)
return len(p), nil
}
func (rw *HttpProxyWriter) WriteRequest(req *http.Request) {
req.Host = rw.hostHeaderRewrite
req.Write(rw)
}
func generateErrorHTML(errorType string) string {
return fmt.Sprintf(
`<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>RTTY</title>
<style>
body {
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, "Helvetica Neue", Arial, sans-serif;
background-color: #555;
line-height: 1.6;
}
.error-container {
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
min-height: 60vh;
text-align: center;
}
.error-icon {
margin-bottom: 2rem;
animation: fadeIn 0.8s ease-in-out;
}
.error-icon svg {
width: 90px;
height: 90px;
fill: #f56565;
}
.error-content {
max-width: 700px;
animation: slideUp 0.8s ease-out 0.2s both;
}
.error-title {
font-size: 1.8rem;
font-weight: 600;
color: #7a8fb0;
margin-bottom: 1rem;
line-height: 1.2;
}
.error-message {
font-size: 1rem;
color: #b6c1d3;
margin-bottom: 2rem;
line-height: 1.6;
text-align: left;
}
@keyframes fadeIn {
from {
opacity: 0;
transform: scale(0.8);
}
to {
opacity: 1;
transform: scale(1);
}
}
@keyframes slideUp {
from {
opacity: 0;
transform: translateY(20px);
}
to {
opacity: 1;
transform: translateY(0);
}
}
</style>
</head>
<body>
<div class="error-container">
<div class="error-icon">
<svg viewBox="0 0 24 24">
<path d="M1 21h22L12 2 1 21zm12-3h-2v-2h2v2zm0-4h-2v-4h2v4z"/>
</svg>
</div>
<div class="error-content">
<h2 class="error-title" id="errorTitle"></h2>
<p class="error-message" id="errorMessage"></p>
</div>
</div>
<script>
const translations = {
en: {
'Device Unavailable': 'Device Unavailable',
'Invalid Request': 'Invalid Request',
'Unauthorized Access': 'Unauthorized Access',
'Device offline message': 'The device is currently offline. Please check the device status and try again.',
'Invalid request message': 'The request is invalid or malformed',
'Unauthorized request message': 'You are not authorized to access this resource. Please check your session and try again.'
},
'zh-CN': {
'Device Unavailable': '设备不可用',
'Invalid Request': '无效请求',
'Unauthorized Access': '未授权访问',
'Device offline message': '设备当前离线,请检查设备状态后重试。',
'Invalid request message': '请求无效或格式错误',
'Unauthorized request message': '您无权访问此资源。请检查您的会话并重试。'
}
};
function t(key, lang) {
return translations[lang][key] || translations.en[key] || key;
}
function updateContent() {
const errorType = '%s';
const lang = navigator.language === 'zh-CN' ? 'zh-CN' : 'en';
let title = '', message = '';
switch (errorType) {
case 'offline':
title = t('Device Unavailable', lang);
message = t('Device offline message', lang);
break;
case 'invalid':
title = t('Invalid Request', lang);
message = t('Invalid request message', lang);
break;
case 'unauthorized':
title = t('Unauthorized Access', lang);
message = t('Unauthorized request message', lang);
break;
}
document.getElementById('errorTitle').textContent = title;
document.getElementById('errorMessage').textContent = message;
// Update page title
if (title) {
document.title = title + ' - RTTY';
} else {
document.title = 'Error - RTTY';
}
}
// Initialize page on load
document.addEventListener('DOMContentLoaded', updateContent);
</script>
</body>
</html>`, errorType)
}
func sendHTTPErrorResponse(conn net.Conn, errorType string) {
htmlContent := generateErrorHTML(errorType)
response := "HTTP/1.1 200 OK\r\n"
response += "Content-Type: text/html; charset=utf-8\r\n"
response += fmt.Sprintf("Content-Length: %d\r\n", len(htmlContent))
response += "Connection: close\r\n"
response += "\r\n"
response += htmlContent
conn.Write([]byte(response))
}
func Write302WithCookie(conn net.Conn, location, cookieName, cookieValue string) {
cookie := fmt.Sprintf("%s=%s; Path=/; HttpOnly", cookieName, cookieValue)
response := fmt.Sprintf(
"HTTP/1.1 302 Found\r\n"+
"Location: %s\r\n"+
"Set-Cookie: %s\r\n"+
"Content-Length: 0\r\n"+
"Connection: close\r\n"+
"\r\n",
location, cookie,
)
_, _ = conn.Write([]byte(response))
}
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 55 KiB

After

Width:  |  Height:  |  Size: 52 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 38 KiB

After

Width:  |  Height:  |  Size: 44 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 58 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 61 KiB

After

Width:  |  Height:  |  Size: 14 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.5 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 418 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 400 KiB

+82
View File
@@ -0,0 +1,82 @@
package app
import (
"context"
"database/sql"
"rttys/internal/pkg/password"
)
type App struct {
srv *Server
}
func (a *App) Start(ctx context.Context) error { return a.srv.Run(ctx) }
func (a *App) Shutdown(ctx context.Context) error { return a.srv.Shutdown(ctx) }
func SeedIfEmpty(ctx context.Context, db *sql.DB) error {
var n int
if err := db.QueryRowContext(ctx, `SELECT COUNT(1) FROM users`).Scan(&n); err != nil {
return err
}
if n > 0 {
return nil
}
// users
adminHash, err := password.HashPassword("admin")
if err != nil {
return err
}
userHash, err := password.HashPassword("user")
if err != nil {
return err
}
if _, err := db.ExecContext(ctx, `
INSERT INTO users(username, description, password_hash, role, status, is_system) VALUES
('admin','Admin', ?, 'admin', 'active', 1),
('user1','User One', ?, 'user', 'active', 0)`,
adminHash, userHash,
); err != nil {
return err
}
// device groups
if _, err := db.ExecContext(ctx, `
INSERT INTO device_groups(name, description) VALUES
('dg1','Device Group 1'),
('dg2','Device Group 2')`); err != nil {
return err
}
// user groups
if _, err := db.ExecContext(ctx, `
INSERT INTO user_groups(name, description) VALUES
('ug1','User Group 1'),
('ug2','User Group 2')`); err != nil {
return err
}
// memberships: user1 in ug1
if _, err := db.ExecContext(ctx, `
INSERT INTO user_group_members(user_id, group_id) VALUES (2, 1)`); err != nil {
return err
}
// links: ug1 -> dg1
if _, err := db.ExecContext(ctx, `
INSERT INTO user_group_device_group_links(user_group_id, device_group_id) VALUES (1, 1)`); err != nil {
return err
}
// devices: one in dg1, one in dg2, one ungrouped
if _, err := db.ExecContext(ctx, `
INSERT INTO devices(ddns, mac, name, description, device_group_id, status, last_seen_at) VALUES
('dev-001','00:11:22:33:44:55','Device 1','in dg1', 1, 'online', unixepoch()),
('dev-002','00:11:22:33:44:66','Device 2','in dg2', 2, 'offline', unixepoch()),
('dev-003','00:11:22:33:44:77','Device 3','ungrouped', NULL, 'online', unixepoch())`); err != nil {
return err
}
return nil
}
+42
View File
@@ -0,0 +1,42 @@
package app
import (
"context"
"net/http"
"time"
)
type Server struct {
httpServer *http.Server
}
func NewServer(addr string, handler http.Handler) *Server {
return &Server{
httpServer: &http.Server{
Addr: addr,
Handler: handler,
ReadHeaderTimeout: 10 * time.Second,
},
}
}
func (s *Server) Run(ctx context.Context) error {
errCh := make(chan error, 1)
go func() {
errCh <- s.httpServer.ListenAndServe()
}()
select {
case <-ctx.Done():
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = s.httpServer.Shutdown(shutdownCtx)
return nil
case err := <-errCh:
return err
}
}
func (s *Server) Shutdown(ctx context.Context) error {
return s.httpServer.Shutdown(ctx)
}
+16
View File
@@ -0,0 +1,16 @@
package config
import "time"
type Config struct {
HTTP struct{ Addr string }
Auth struct{ SessionTTL time.Duration }
DB struct{ DSN string }
}
func MustLoad() Config {
var cfg Config
cfg.Auth.SessionTTL = 24 * time.Hour
cfg.DB.DSN = "file:rttys.db?_pragma=foreign_keys(1)"
return cfg
}
+22
View File
@@ -0,0 +1,22 @@
package device
type Status string
const (
StatusOnline Status = "online"
StatusOffline Status = "offline"
StatusDisabled Status = "disabled"
)
type Device struct {
ID int64
Ddns string
Mac string
Name string
Description string
IP string
Client string
DeviceGroupID *int64 // nil means ungrouped (admin-only visibility)
Status Status
LastSeenAt *int64
}
+8
View File
@@ -0,0 +1,8 @@
package device
import "context"
type Repository interface {
ListAll(ctx context.Context) ([]Device, error)
ListByDeviceGroupIDs(ctx context.Context, groupIDs []int64) ([]Device, error)
}
+36
View File
@@ -0,0 +1,36 @@
package device
import (
"context"
"rttys/internal/domain/identity"
"rttys/internal/domain/group"
)
type Service struct {
repo Repository
grepo group.Repository
}
func NewService(repo Repository, grepo group.Repository) *Service {
return &Service{repo: repo, grepo: grepo}
}
// ListVisible implements the scope rule:
// - admin: all devices, including ungrouped
// - user : only devices whose device_group_id is linked via (user -> user_groups -> device_groups)
func (s *Service) ListVisible(ctx context.Context, role identity.Role, userID int64) ([]Device, error) {
if role == identity.RoleAdmin {
return s.repo.ListAll(ctx)
}
dgIDs, err := s.grepo.ListDeviceGroupIDsByUser(ctx, userID)
if err != nil || len(dgIDs) == 0 {
return []Device{}, nil
}
return s.repo.ListByDeviceGroupIDs(ctx, dgIDs)
}
func (s *Service) ListByDeviceGroupIDs(ctx context.Context, groupIDs []int64) ([]Device, error) {
return s.repo.ListByDeviceGroupIDs(ctx, groupIDs)
}
+7
View File
@@ -0,0 +1,7 @@
package devicegroup
type DeviceGroup struct {
ID int64
Name string
Description string
}
+7
View File
@@ -0,0 +1,7 @@
package devicegroup
import "context"
type Repository interface {
ListDeviceGroupsVisibleToUser(ctx context.Context, userID int64, isAdmin bool) ([]DeviceGroup, error)
}
+7
View File
@@ -0,0 +1,7 @@
package group
type UserGroup struct {
ID int64
Name string
Description string
}
+8
View File
@@ -0,0 +1,8 @@
package group
import "context"
type Repository interface {
ListUserGroupsVisibleToUser(ctx context.Context, userID int64, isAdmin bool) ([]UserGroup, error)
ListDeviceGroupIDsByUser(ctx context.Context, userID int64) ([]int64, error)
}
+17
View File
@@ -0,0 +1,17 @@
package identity
type Role string
const (
RoleAdmin Role = "admin"
RoleUser Role = "user"
)
func RoleFromString(v string) Role {
switch v {
case string(RoleAdmin):
return RoleAdmin
default:
return RoleUser
}
}
+33
View File
@@ -0,0 +1,33 @@
package permission
import "rttys/internal/domain/identity"
type Key string
const (
MeRead Key = "me.read"
AuthWrite Key = "auth.write"
DeviceRead Key = "device.read"
DeviceWrite Key = "device.write"
DeviceGroupRead Key = "device_group.read"
DeviceGroupWrite Key = "device_group.write"
UserGroupRead Key = "user_group.read"
UserGroupWrite Key = "user_group.write"
UserRead Key = "user.read"
UserWrite Key = "user.write"
RelationWrite Key = "relation.write"
)
func DefaultKeysForRole(role identity.Role) []Key {
switch role {
case identity.RoleAdmin:
return []Key{ /* ... */ }
default:
return []Key{ /* ... */ }
}
}
+20
View File
@@ -0,0 +1,20 @@
package permission
import (
"context"
"rttys/internal/domain/identity"
)
type Repository interface {
ListKeysByRole(ctx context.Context, role identity.Role) ([]Key, error)
}
type Service struct {
repo Repository
}
func NewService(repo Repository) *Service { return &Service{repo: repo} }
func (s *Service) ListByRole(ctx context.Context, role identity.Role) ([]Key, error) {
return s.repo.ListKeysByRole(ctx, role)
}
+23
View File
@@ -0,0 +1,23 @@
package user
import (
"rttys/internal/domain/identity"
)
type Status string
const (
StatusActive Status = "active"
StatusDisabled Status = "disabled"
)
type User struct {
ID int64
Username string
Email string
Description string
PasswordHash string
Role identity.Role
Status Status
IsSystem bool
}
+14
View File
@@ -0,0 +1,14 @@
package user
import "context"
type Repository interface {
FindByID(ctx context.Context, id int64) (*User, error)
FindByUsername(ctx context.Context, username string) (*User, error)
FindSystemAdmin(ctx context.Context) (*User, error)
Create(ctx context.Context, u *User) (int64, error)
Update(ctx context.Context, u *User) error
Delete(ctx context.Context, id int64) error
List(ctx context.Context) ([]User, error)
}
+110
View File
@@ -0,0 +1,110 @@
package user
import (
"context"
"errors"
"rttys/internal/domain/identity"
"rttys/internal/pkg/password"
)
var (
ErrUserNotFound = errors.New("user not found")
ErrUserDisabled = errors.New("user disabled")
ErrBadPassword = errors.New("bad password")
)
type Service struct{ repo Repository }
func NewService(repo Repository) *Service { return &Service{repo: repo} }
func (s *Service) Authenticate(ctx context.Context, username, pw string) (*User, error) {
u, err := s.repo.FindByUsername(ctx, username)
if err != nil || u == nil {
return nil, ErrUserNotFound
}
if u.Status == StatusDisabled {
return nil, ErrUserDisabled
}
if !password.VerifyPassword(pw, u.PasswordHash) {
return nil, ErrBadPassword
}
return u, nil
}
func (s *Service) GetByID(ctx context.Context, id int64) (*User, error) {
u, err := s.repo.FindByID(ctx, id)
if err != nil || u == nil {
return nil, ErrUserNotFound
}
if u.Status == StatusDisabled {
return nil, ErrUserDisabled
}
return u, nil
}
// FindByID returns user even if disabled.
func (s *Service) FindByID(ctx context.Context, id int64) (*User, error) {
u, err := s.repo.FindByID(ctx, id)
if err != nil || u == nil {
return nil, ErrUserNotFound
}
return u, nil
}
func (s *Service) GetSystemAdmin(ctx context.Context) (*User, error) {
return s.repo.FindSystemAdmin(ctx)
}
func (s *Service) List(ctx context.Context) ([]User, error) {
return s.repo.List(ctx)
}
// CreateUser creates a user; passwordPlain will be hashed.
func (s *Service) CreateUser(ctx context.Context, username, description, passwordPlain, role, status string) (int64, error) {
hash, err := password.HashPassword(passwordPlain)
if err != nil {
return 0, err
}
u := &User{
Username: username,
Description: description,
PasswordHash: hash,
Role: identity.RoleFromString(role),
Status: Status(status),
}
return s.repo.Create(ctx, u)
}
// UpdateUser updates fields; if passwordPlain is empty, keep existing.
func (s *Service) UpdateUser(ctx context.Context, id int64, username, description, passwordPlain, role, status *string) error {
exist, err := s.repo.FindByID(ctx, id)
if err != nil || exist == nil {
return ErrUserNotFound
}
if username != nil && *username != "" {
exist.Username = *username
}
if description != nil {
exist.Description = *description
}
if role != nil && *role != "" {
exist.Role = identity.RoleFromString(*role)
}
if status != nil && *status != "" {
exist.Status = Status(*status)
}
if passwordPlain != nil && *passwordPlain != "" {
hash, err := password.HashPassword(*passwordPlain)
if err != nil {
return err
}
exist.PasswordHash = hash
}
return s.repo.Update(ctx, exist)
}
func (s *Service) DeleteUser(ctx context.Context, id int64) error {
return s.repo.Delete(ctx, id)
}
+13
View File
@@ -0,0 +1,13 @@
package dto
type LoginReq struct {
Username string `json:"username"`
Password string `json:"password"`
AuthMethod string `json:"authMethod,omitempty"`
}
type LoginResp struct {
Token string `json:"token"`
}
type LogoutResp struct{}
+59
View File
@@ -0,0 +1,59 @@
package dto
import (
"net/http"
"time"
"github.com/gin-gonic/gin"
)
type Meta struct {
TraceID string `json:"traceId"`
TS int64 `json:"ts"`
}
type Envelope[T any] struct {
Ok bool `json:"ok"`
Code string `json:"code"`
Message string `json:"message"`
Data T `json:"data"`
Meta Meta `json:"meta"`
}
const (
CodeOK = "OK"
CodeInvalidArgument = "INVALID_ARGUMENT"
CodeValidationFailed = "VALIDATION_FAILED"
CodeAuthRequired = "AUTH_REQUIRED"
CodeAuthExpired = "AUTH_EXPIRED"
CodeForbidden = "FORBIDDEN"
CodeNotFound = "NOT_FOUND"
CodeConflict = "CONFLICT"
CodeInternalError = "INTERNAL_ERROR"
)
func NowUnix() int64 { return time.Now().Unix() }
// Write always replies with HTTP 200 per spec.
func Write[T any](c *gin.Context, env Envelope[T]) {
c.JSON(http.StatusOK, env)
}
func Ok[T any](traceID string, data T) Envelope[T] {
return Envelope[T]{
Ok: true,
Code: CodeOK,
Data: data,
Meta: Meta{TraceID: traceID, TS: NowUnix()},
}
}
func Err(traceID, code, msg string, data any) Envelope[any] {
return Envelope[any]{
Ok: false,
Code: code,
Message: msg,
Data: data,
Meta: Meta{TraceID: traceID, TS: NowUnix()},
}
}
+28
View File
@@ -0,0 +1,28 @@
package dto
type Device struct {
ID int64 `json:"id"`
Ddns string `json:"ddns"`
Status string `json:"status"`
ConnectedTime int64 `json:"connectedTime"`
IP string `json:"ip"`
Mac string `json:"mac"`
Description string `json:"description"`
Client string `json:"client"`
DeviceGroupID *int64 `json:"deviceGroupId"`
DeviceGroupName string `json:"deviceGroupName"`
}
type ListDevicesResp struct {
Items []Device `json:"items"`
Page int `json:"page"`
PageSize int `json:"pageSize"`
Total int `json:"total"`
}
type MoveDevicesToGroupReq struct {
DeviceIDs []int64 `json:"deviceIds"`
GroupID int64 `json:"groupId"`
}
type MoveDevicesToGroupResp struct{}
+50
View File
@@ -0,0 +1,50 @@
package dto
type DeviceGroupUserGroupRef struct {
UserGroupID int64 `json:"userGroupId"`
UserGroupName string `json:"userGroupName"`
}
type DeviceGroup struct {
ID int64 `json:"id"`
Name string `json:"name"`
DeviceCount int64 `json:"deviceCount"`
Description string `json:"description"`
UserGroupList []DeviceGroupUserGroupRef `json:"userGroupList"`
}
type ListDeviceGroupsResp struct {
Items []DeviceGroup `json:"items"`
}
type DeviceGroupOption struct {
GroupID int64 `json:"groupId"`
Name string `json:"name"`
}
type ListDeviceGroupOptionsResp struct {
Items []DeviceGroupOption `json:"items"`
}
type CreateDeviceGroupReq struct {
Name string `json:"name"`
Description string `json:"description"`
UserGroupIDs []int64 `json:"userGroupIds"`
DeviceIDs []int64 `json:"deviceIds"`
}
type CreateDeviceGroupResp struct {
ID int64 `json:"id"`
}
type UpdateDeviceGroupReq struct {
Name string `json:"name"`
Description string `json:"description"`
UserGroupIDs []int64 `json:"userGroupIds"`
}
type DeleteDeviceGroupResp struct{}
type ModifyDeviceGroupDevicesReq struct {
DeviceIDs []int64 `json:"deviceIds"`
}
+13
View File
@@ -0,0 +1,13 @@
package dto
type MeUser struct {
ID int64 `json:"id"`
Username string `json:"username"`
DisplayName string `json:"displayName"`
Role string `json:"role"`
}
type MeResp struct {
User MeUser `json:"user"`
Permissions []string `json:"permissions"`
}
+25
View File
@@ -0,0 +1,25 @@
package dto
type SetUserGroupsReq struct {
GroupIDs []int64 `json:"groupIds"`
}
type SetUserGroupsResp struct {
UserID int64 `json:"userId"`
GroupIDs []int64 `json:"groupIds"`
}
type SetUserGroupDeviceGroupsReq struct {
DeviceGroupIDs []int64 `json:"deviceGroupIds"`
}
type SetUserGroupDeviceGroupsResp struct {
UserGroupID int64 `json:"userGroupId"`
DeviceGroupIDs []int64 `json:"deviceGroupIds"`
}
type SetDeviceGroupDevicesReq struct {
DeviceIDs []int64 `json:"deviceIds"`
}
type SetDeviceGroupDevicesResp struct {
DeviceGroupID int64 `json:"deviceGroupId"`
DeviceIDs []int64 `json:"deviceIds"`
}
+41
View File
@@ -0,0 +1,41 @@
package dto
type UserGroupRef struct {
UserGroupID int64 `json:"userGroupId"`
UserGroupName string `json:"userGroupName"`
}
type User struct {
ID int64 `json:"id"`
Username string `json:"username"`
Description string `json:"description"`
Role string `json:"role"`
IsSystem bool `json:"isSystem"`
UserGroupList []UserGroupRef `json:"userGroupList"`
}
type ListUsersResp struct {
Items []User `json:"items"`
}
type CreateUserReq struct {
Role string `json:"role"`
Username string `json:"username"`
Description string `json:"description"`
Password string `json:"password"`
Repassword string `json:"repassword"`
UserGroupIDs []int64 `json:"userGroupIds"`
}
type CreateUserResp struct{}
type UpdateUserReq struct {
Role *string `json:"role"`
Username *string `json:"username"`
Description *string `json:"description"`
Password *string `json:"password"`
Repassword *string `json:"repassword"`
UserGroupIDs *[]int64 `json:"userGroupIds"`
}
type DeleteUserResp struct{}
+43
View File
@@ -0,0 +1,43 @@
package dto
type UserGroupDeviceGroupRef struct {
DeviceGroupID int64 `json:"deviceGroupId"`
DeviceGroupName string `json:"deviceGroupName"`
}
type UserGroup struct {
ID int64 `json:"id"`
UserGroup string `json:"userGroup"`
Description string `json:"description"`
UserCount int64 `json:"userCount"`
DeviceGroupList []UserGroupDeviceGroupRef `json:"deviceGroupList"`
}
type ListUserGroupsResp struct {
Items []UserGroup `json:"items"`
}
type UserGroupOption struct {
UserGroupID int64 `json:"userGroupId"`
Name string `json:"name"`
}
type ListUserGroupOptionsResp struct {
Items []UserGroupOption `json:"items"`
}
type CreateUserGroupReq struct {
Name string `json:"name"`
Description string `json:"description"`
}
type CreateUserGroupResp struct {
ID int64 `json:"id"`
}
type UpdateUserGroupReq struct {
Name string `json:"name"`
Description string `json:"description"`
}
type DeleteUserGroupResp struct{}
+101
View File
@@ -0,0 +1,101 @@
package handler
import (
"rttys/internal/pkg/ldap"
"rttys/xconfig"
"strings"
"rttys/internal/domain/user"
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"rttys/internal/pkg/randtoken"
"rttys/internal/store/memory"
"github.com/gin-gonic/gin"
)
type AuthHandler struct {
userSvc *user.Service
sessionStore *memory.SessionStore
}
func NewAuthHandler(userSvc *user.Service, sessionStore *memory.SessionStore) *AuthHandler {
return &AuthHandler{userSvc: userSvc, sessionStore: sessionStore}
}
// POST /api/login
func (h *AuthHandler) Login(c *gin.Context) {
traceID := middleware.GetTraceID(c)
var req dto.LoginReq
if err := c.ShouldBindJSON(&req); err != nil || req.Username == "" || req.Password == "" {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "username/password",
}))
return
}
cfg := xconfig.Must()
var userID int64
// ---- LDAP ----
authMethod := req.AuthMethod
if authMethod == "ldap" {
ok, errorType := ldap.AuthenticateUserWithError(cfg, req.Username, req.Password, authMethod)
if !ok {
if errorType == "authorization" {
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "User not authorized", nil))
} else {
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "Authentication failed", nil))
}
return
}
sysAdmin, err := h.userSvc.GetSystemAdmin(c.Request.Context())
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "System admin not found", nil))
return
}
userID = sysAdmin.ID
} else {
u, err := h.userSvc.Authenticate(c.Request.Context(), req.Username, req.Password)
if err != nil || u == nil {
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "Authentication failed", nil))
return
}
userID = u.ID
}
sid, err := randtoken.New()
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
h.sessionStore.Create(sid, userID)
dto.Write(c, dto.Ok(traceID, dto.LoginResp{
Token: sid,
}))
}
// POST /api/logout
func (h *AuthHandler) Logout(c *gin.Context) {
traceID := middleware.GetTraceID(c)
// 1) Try bearer
authz := strings.TrimSpace(c.GetHeader("Authorization"))
if strings.HasPrefix(strings.ToLower(authz), "bearer ") {
token := strings.TrimSpace(authz[7:])
if token != "" {
h.sessionStore.Delete(token)
}
}
// 2) Try cookie sid
if sid, err := c.Cookie("sid"); err == nil && strings.TrimSpace(sid) != "" {
h.sessionStore.Delete(strings.TrimSpace(sid))
// 清 cookie
c.SetCookie("sid", "", -1, "/", "", false, true)
}
dto.Write(c, dto.Ok(traceID, dto.LogoutResp{}))
}
+281
View File
@@ -0,0 +1,281 @@
package handler
import (
"errors"
"io"
"sort"
"strconv"
"strings"
"rttys/internal/domain/device"
"rttys/internal/domain/identity"
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"rttys/internal/store/sqlite"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
type DeviceHandler struct {
devSvc *device.Service
groupRepo *sqlite.GroupRepo
relationsRepo *sqlite.RelationsRepo
}
func NewDeviceHandler(
devSvc *device.Service,
groupRepo *sqlite.GroupRepo,
relationsRepo *sqlite.RelationsRepo,
) *DeviceHandler {
return &DeviceHandler{
devSvc: devSvc,
groupRepo: groupRepo,
relationsRepo: relationsRepo,
}
}
// GET /api/devices
func (h *DeviceHandler) ListDevices(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
var filterGroupID *int64
if raw := strings.TrimSpace(c.Query("groupId")); raw != "" {
id, err := strconv.ParseInt(raw, 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "groupId",
}))
return
}
filterGroupID = &id
}
isAdmin := p.Role == identity.RoleAdmin
var items []device.Device
if isAdmin {
var err error
if filterGroupID != nil {
items, err = h.devSvc.ListByDeviceGroupIDs(c.Request.Context(), []int64{*filterGroupID})
} else {
items, err = h.devSvc.ListVisible(c.Request.Context(), p.Role, p.UserID)
}
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
} else {
if h.groupRepo == nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
dgIDs, err := h.groupRepo.ListDeviceGroupIDsByUser(c.Request.Context(), p.UserID)
if err != nil || len(dgIDs) == 0 {
dto.Write(c, dto.Ok(traceID, dto.ListDevicesResp{
Items: []dto.Device{},
Page: 1,
PageSize: 0,
Total: 0,
}))
return
}
if filterGroupID != nil {
allowed := false
for _, gid := range dgIDs {
if gid == *filterGroupID {
allowed = true
break
}
}
if !allowed {
dto.Write(c, dto.Ok(traceID, dto.ListDevicesResp{
Items: []dto.Device{},
Page: 1,
PageSize: 0,
Total: 0,
}))
return
}
dgIDs = []int64{*filterGroupID}
}
items, err = h.devSvc.ListByDeviceGroupIDs(c.Request.Context(), dgIDs)
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
}
groupNameByID := map[int64]string{}
if h.groupRepo != nil {
groups, err := h.groupRepo.ListDeviceGroupsVisibleToUser(c.Request.Context(), p.UserID, isAdmin)
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
for _, g := range groups {
groupNameByID[g.ID] = g.Name
}
}
sort.SliceStable(items, func(i, j int) bool {
oi := items[i].Status == device.StatusOnline
oj := items[j].Status == device.StatusOnline
if oi != oj {
return oi
}
return items[i].Ddns < items[j].Ddns
})
out := make([]dto.Device, 0, len(items))
for _, d := range items {
var groupName string
if d.DeviceGroupID != nil {
groupName = groupNameByID[*d.DeviceGroupID]
}
var connectedTime int64
if d.LastSeenAt != nil {
connectedTime = *d.LastSeenAt
}
out = append(out, dto.Device{
ID: d.ID,
Ddns: d.Ddns,
Status: string(d.Status),
ConnectedTime: connectedTime,
IP: d.IP,
Mac: d.Mac,
Description: d.Description,
Client: d.Client,
DeviceGroupID: d.DeviceGroupID,
DeviceGroupName: groupName,
})
}
dto.Write(c, dto.Ok(traceID, dto.ListDevicesResp{
Items: out,
Page: 1,
PageSize: len(out),
Total: len(out),
}))
}
type UpdateDeviceRequest struct {
Description *string `json:"description"`
}
// PUT /api/devices/:id
func (h *DeviceHandler) UpdateDevice(c *gin.Context) {
traceID := middleware.GetTraceID(c)
idStr := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseInt(idStr, 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "id",
}))
return
}
var req UpdateDeviceRequest
if err := c.ShouldBindJSON(&req); err != nil && !errors.Is(err, io.EOF) {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "description",
"error": err.Error(),
}))
return
}
db := sqlite.MustContainer().Gorm.WithContext(c.Request.Context())
var row struct {
Description string `gorm:"column:description"`
}
tx := db.Table("devices").Select("description").Where("id = ?", id).First(&row)
if errors.Is(tx.Error, gorm.ErrRecordNotFound) {
dto.Write(c, dto.Err(traceID, dto.CodeNotFound, "Device not found", nil))
return
}
if tx.Error != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{
"detail": tx.Error.Error(),
}))
return
}
newDesc := row.Description
if req.Description != nil {
newDesc = *req.Description
}
if req.Description != nil {
res := db.Exec(
`UPDATE devices SET description=? WHERE id=?`,
newDesc, id,
)
if res.Error != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{
"detail": res.Error.Error(),
}))
return
}
if res.RowsAffected == 0 {
dto.Write(c, dto.Err(traceID, dto.CodeNotFound, "Device not found", nil))
return
}
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
// DELETE /api/devices/:id
func (h *DeviceHandler) DeleteDevice(c *gin.Context) {
traceID := middleware.GetTraceID(c)
idStr := strings.TrimSpace(c.Param("id"))
id, err := strconv.ParseInt(idStr, 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "id",
}))
return
}
db := sqlite.MustContainer().Gorm.WithContext(c.Request.Context())
res := db.Exec(`DELETE FROM devices WHERE id=?`, id)
if res.Error != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{
"detail": res.Error.Error(),
}))
return
}
if res.RowsAffected == 0 {
dto.Write(c, dto.Err(traceID, dto.CodeNotFound, "Device not found", nil))
return
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
// POST /api/devices/move-to-device-group
func (h *DeviceHandler) MoveToDeviceGroup(c *gin.Context) {
traceID := middleware.GetTraceID(c)
_ = middleware.MustPrincipal(c)
if h.relationsRepo == nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
var req dto.MoveDevicesToGroupReq
if err := c.ShouldBindJSON(&req); err != nil || req.GroupID <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "groupId"}))
return
}
if err := h.relationsRepo.AddDevicesToGroup(c.Request.Context(), req.GroupID, req.DeviceIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
dto.Write(c, dto.Ok(traceID, dto.MoveDevicesToGroupResp{}))
}
+215
View File
@@ -0,0 +1,215 @@
package handler
import (
"rttys/internal/domain/devicegroup"
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"rttys/internal/store/sqlite"
"strconv"
"strings"
"github.com/gin-gonic/gin"
)
type DeviceGroupHandler struct {
groupRepo *sqlite.GroupRepo
relationsRepo *sqlite.RelationsRepo
}
func NewDeviceGroupHandler(groupRepo *sqlite.GroupRepo, relationsRepo *sqlite.RelationsRepo) *DeviceGroupHandler {
return &DeviceGroupHandler{groupRepo: groupRepo, relationsRepo: relationsRepo}
}
// GET /api/device-groups
func (h *DeviceGroupHandler) ListDeviceGroups(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
items, err := h.groupRepo.ListDeviceGroupDetails(c.Request.Context(), p.UserID, string(p.Role) == "admin")
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
out := make([]dto.DeviceGroup, 0, len(items))
for _, it := range items {
userGroups := make([]dto.DeviceGroupUserGroupRef, 0, len(it.UserGroups))
for _, ug := range it.UserGroups {
userGroups = append(userGroups, dto.DeviceGroupUserGroupRef{
UserGroupID: ug.ID,
UserGroupName: ug.Name,
})
}
out = append(out, dto.DeviceGroup{
ID: it.ID,
Name: it.Name,
DeviceCount: it.DeviceCount,
Description: it.Description,
UserGroupList: userGroups,
})
}
dto.Write(c, dto.Ok(traceID, dto.ListDeviceGroupsResp{Items: out}))
}
var _ devicegroup.DeviceGroup
// GET /api/device-groups/options
func (h *DeviceGroupHandler) ListOptions(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
items, err := h.groupRepo.ListDeviceGroupsVisibleToUser(c.Request.Context(), p.UserID, string(p.Role) == "admin")
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
out := make([]dto.DeviceGroupOption, 0, len(items))
for _, it := range items {
out = append(out, dto.DeviceGroupOption{GroupID: it.ID, Name: it.Name})
}
dto.Write(c, dto.Ok(traceID, dto.ListDeviceGroupOptionsResp{Items: out}))
}
func (h *DeviceGroupHandler) Create(c *gin.Context) {
traceID := middleware.GetTraceID(c)
var req dto.CreateDeviceGroupReq
if err := c.ShouldBindJSON(&req); err != nil || req.Name == "" {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "name"}))
return
}
id, err := h.groupRepo.CreateDeviceGroup(c.Request.Context(), req.Name, req.Description)
if err != nil {
if strings.Contains(strings.ToLower(err.Error()), "unique") {
dto.Write(c, dto.Err(traceID, dto.CodeConflict, "Name already exists", nil))
return
}
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
if h.relationsRepo != nil {
if err := h.relationsRepo.SetDeviceGroupUserGroups(c.Request.Context(), id, req.UserGroupIDs); err != nil {
_ = h.groupRepo.DeleteDeviceGroup(c.Request.Context(), id)
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
if err := h.relationsRepo.AddDevicesToGroup(c.Request.Context(), id, req.DeviceIDs); err != nil {
_ = h.groupRepo.DeleteDeviceGroup(c.Request.Context(), id)
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
}
dto.Write(c, dto.Ok(traceID, dto.CreateDeviceGroupResp{ID: id}))
}
func (h *DeviceGroupHandler) Update(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.UpdateDeviceGroupReq
if err := c.ShouldBindJSON(&req); err != nil || req.Name == "" {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "name"}))
return
}
if err := h.groupRepo.UpdateDeviceGroup(c.Request.Context(), id, req.Name, req.Description); err != nil {
if strings.Contains(strings.ToLower(err.Error()), "unique") {
dto.Write(c, dto.Err(traceID, dto.CodeConflict, "Name already exists", nil))
return
}
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
if h.relationsRepo != nil {
if err := h.relationsRepo.SetDeviceGroupUserGroups(c.Request.Context(), id, req.UserGroupIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
func (h *DeviceGroupHandler) Delete(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
if err := h.groupRepo.DeleteDeviceGroup(c.Request.Context(), id); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
dto.Write(c, dto.Ok(traceID, dto.DeleteDeviceGroupResp{}))
}
// POST /api/device-groups/:id/devices
func (h *DeviceGroupHandler) AddDevices(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.ModifyDeviceGroupDevicesReq
if err := c.ShouldBindJSON(&req); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", nil))
return
}
if h.relationsRepo == nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
if err := h.relationsRepo.AddDevicesToGroup(c.Request.Context(), id, req.DeviceIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
// DELETE /api/device-groups/:id/devices
func (h *DeviceGroupHandler) RemoveDevices(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.ModifyDeviceGroupDevicesReq
if err := c.ShouldBindJSON(&req); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", nil))
return
}
if h.relationsRepo == nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
if err := h.relationsRepo.RemoveDevicesFromGroup(c.Request.Context(), id, req.DeviceIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
+28
View File
@@ -0,0 +1,28 @@
package handler
import (
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"github.com/gin-gonic/gin"
)
type MeHandler struct{}
func NewMeHandler() *MeHandler { return &MeHandler{} }
// GET /api/me
func (h *MeHandler) GetMe(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
dto.Write(c, dto.Ok(traceID, dto.MeResp{
User: dto.MeUser{
ID: p.UserID,
Username: p.Username,
DisplayName: p.DisplayName,
Role: string(p.Role),
},
Permissions: p.PermissionKeys,
}))
}
+89
View File
@@ -0,0 +1,89 @@
package handler
import (
"strconv"
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"rttys/internal/store/sqlite"
"github.com/gin-gonic/gin"
)
type RelationsHandler struct{ repo *sqlite.RelationsRepo }
func NewRelationsHandler(repo *sqlite.RelationsRepo) *RelationsHandler {
return &RelationsHandler{repo: repo}
}
func (h *RelationsHandler) SetUserGroups(c *gin.Context) {
traceID := middleware.GetTraceID(c)
_ = middleware.MustPrincipal(c)
userID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || userID <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.SetUserGroupsReq
if err := c.ShouldBindJSON(&req); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", nil))
return
}
if err := h.repo.SetUserGroups(c.Request.Context(), userID, req.GroupIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
dto.Write(c, dto.Ok(traceID, dto.SetUserGroupsResp{UserID: userID, GroupIDs: req.GroupIDs}))
}
func (h *RelationsHandler) SetUserGroupDeviceGroups(c *gin.Context) {
traceID := middleware.GetTraceID(c)
_ = middleware.MustPrincipal(c)
ugID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || ugID <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.SetUserGroupDeviceGroupsReq
if err := c.ShouldBindJSON(&req); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", nil))
return
}
if err := h.repo.SetUserGroupDeviceGroups(c.Request.Context(), ugID, req.DeviceGroupIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
dto.Write(c, dto.Ok(traceID, dto.SetUserGroupDeviceGroupsResp{UserGroupID: ugID, DeviceGroupIDs: req.DeviceGroupIDs}))
}
func (h *RelationsHandler) SetDeviceGroupDevices(c *gin.Context) {
traceID := middleware.GetTraceID(c)
_ = middleware.MustPrincipal(c)
dgID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || dgID <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.SetDeviceGroupDevicesReq
if err := c.ShouldBindJSON(&req); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", nil))
return
}
if err := h.repo.SetDeviceGroupDevices(c.Request.Context(), dgID, req.DeviceIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
dto.Write(c, dto.Ok(traceID, dto.SetDeviceGroupDevicesResp{DeviceGroupID: dgID, DeviceIDs: req.DeviceIDs}))
}
+235
View File
@@ -0,0 +1,235 @@
package handler
import (
"rttys/internal/domain/identity"
"sort"
"strconv"
"strings"
"rttys/internal/domain/user"
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"rttys/internal/store/memory"
"rttys/internal/store/sqlite"
"github.com/gin-gonic/gin"
)
type UserHandler struct {
userSvc *user.Service
groupRepo *sqlite.GroupRepo
relationsRepo *sqlite.RelationsRepo
sessionStore *memory.SessionStore
}
func NewUserHandler(userSvc *user.Service, groupRepo *sqlite.GroupRepo, relationsRepo *sqlite.RelationsRepo, sessionStore *memory.SessionStore) *UserHandler {
return &UserHandler{
userSvc: userSvc,
groupRepo: groupRepo,
relationsRepo: relationsRepo,
sessionStore: sessionStore,
}
}
func (h *UserHandler) ListUsers(c *gin.Context) {
traceID := middleware.GetTraceID(c)
items, err := h.userSvc.List(c.Request.Context())
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
sort.SliceStable(items, func(i, j int) bool {
rank := func(u user.User) int {
if u.IsSystem {
return 0
}
if u.Role == identity.RoleAdmin {
return 1
}
return 2
}
ri := rank(items[i])
rj := rank(items[j])
if ri != rj {
return ri < rj
}
return false
})
userIDs := make([]int64, 0, len(items))
for _, u := range items {
userIDs = append(userIDs, u.ID)
}
var groupsByUserID map[int64][]sqlite.UserGroupBrief
if h.groupRepo != nil {
groupsByUserID, err = h.groupRepo.ListUserGroupsByUserIDs(c.Request.Context(), userIDs)
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
}
out := make([]dto.User, 0, len(items))
for _, u := range items {
groups := make([]dto.UserGroupRef, 0)
if list, ok := groupsByUserID[u.ID]; ok {
for _, g := range list {
groups = append(groups, dto.UserGroupRef{
UserGroupID: g.ID,
UserGroupName: g.Name,
})
}
}
out = append(out, dto.User{
ID: u.ID,
Role: string(u.Role),
Username: u.Username,
Description: u.Description,
IsSystem: u.IsSystem,
UserGroupList: groups,
})
}
dto.Write(c, dto.Ok(traceID, dto.ListUsersResp{Items: out}))
}
func (h *UserHandler) CreateUser(c *gin.Context) {
traceID := middleware.GetTraceID(c)
var req dto.CreateUserReq
if err := c.ShouldBindJSON(&req); err != nil || req.Username == "" || req.Password == "" {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "username/password",
}))
return
}
if req.Repassword != "" && req.Repassword != req.Password {
dto.Write(c, dto.Err(traceID, dto.CodeValidationFailed, "Passwords do not match", map[string]any{
"field": "repassword",
}))
return
}
if req.Role == "" {
req.Role = "user"
}
status := "active"
id, err := h.userSvc.CreateUser(c.Request.Context(), req.Username, req.Description, req.Password, req.Role, status)
if err != nil {
// best-effort conflict detection
if strings.Contains(strings.ToLower(err.Error()), "unique") {
dto.Write(c, dto.Err(traceID, dto.CodeConflict, "Username already exists", nil))
return
}
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
if h.relationsRepo != nil {
if err := h.relationsRepo.SetUserGroups(c.Request.Context(), id, req.UserGroupIDs); err != nil {
_ = h.userSvc.DeleteUser(c.Request.Context(), id)
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
}
dto.Write(c, dto.Ok(traceID, dto.CreateUserResp{}))
}
func (h *UserHandler) UpdateUser(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "id",
}))
return
}
var req dto.UpdateUserReq
if err := c.ShouldBindJSON(&req); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", nil))
return
}
if req.Password != nil && req.Repassword != nil && *req.Password != *req.Repassword {
dto.Write(c, dto.Err(traceID, dto.CodeValidationFailed, "Passwords do not match", map[string]any{
"field": "repassword",
}))
return
}
target, err := h.userSvc.FindByID(c.Request.Context(), id)
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeNotFound, "Not found", nil))
return
}
if target.IsSystem {
req.Username = nil
req.Role = nil
req.Password = nil
req.Repassword = nil
}
if err := h.userSvc.UpdateUser(c.Request.Context(), id, req.Username, req.Description, req.Password, req.Role, nil); err != nil {
if strings.Contains(strings.ToLower(err.Error()), "not found") {
dto.Write(c, dto.Err(traceID, dto.CodeNotFound, "Not found", nil))
return
}
if strings.Contains(strings.ToLower(err.Error()), "unique") {
dto.Write(c, dto.Err(traceID, dto.CodeConflict, "Username already exists", nil))
return
}
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
if h.relationsRepo != nil && req.UserGroupIDs != nil {
if err := h.relationsRepo.SetUserGroups(c.Request.Context(), id, *req.UserGroupIDs); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", map[string]any{"detail": err.Error()}))
return
}
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
func (h *UserHandler) DeleteUser(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{
"field": "id",
}))
return
}
if id == p.UserID {
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "Cannot delete your own account", nil))
return
}
u, err := h.userSvc.FindByID(c.Request.Context(), id)
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeNotFound, "Not found", nil))
return
}
if u.IsSystem {
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "System user cannot be deleted", nil))
return
}
if h.sessionStore != nil {
h.sessionStore.DeleteByUserID(id)
}
if err := h.userSvc.DeleteUser(c.Request.Context(), id); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
dto.Write(c, dto.Ok(traceID, dto.DeleteUserResp{}))
}
+133
View File
@@ -0,0 +1,133 @@
package handler
import (
"rttys/internal/http/dto"
"rttys/internal/http/middleware"
"rttys/internal/store/sqlite"
"strconv"
"strings"
"github.com/gin-gonic/gin"
)
type UserGroupHandler struct {
groupRepo *sqlite.GroupRepo
}
func NewUserGroupHandler(groupRepo *sqlite.GroupRepo) *UserGroupHandler {
return &UserGroupHandler{groupRepo: groupRepo}
}
// GET /api/user-groups
func (h *UserGroupHandler) ListUserGroups(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
items, err := h.groupRepo.ListUserGroupDetails(c.Request.Context(), p.UserID, string(p.Role) == "admin")
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
out := make([]dto.UserGroup, 0, len(items))
for _, it := range items {
deviceGroups := make([]dto.UserGroupDeviceGroupRef, 0, len(it.DeviceGroups))
for _, dg := range it.DeviceGroups {
deviceGroups = append(deviceGroups, dto.UserGroupDeviceGroupRef{
DeviceGroupID: dg.ID,
DeviceGroupName: dg.Name,
})
}
out = append(out, dto.UserGroup{
ID: it.ID,
UserGroup: it.Name,
Description: it.Description,
UserCount: it.UserCount,
DeviceGroupList: deviceGroups,
})
}
dto.Write(c, dto.Ok(traceID, dto.ListUserGroupsResp{Items: out}))
}
// GET /api/user-groups/options
func (h *UserGroupHandler) ListOptions(c *gin.Context) {
traceID := middleware.GetTraceID(c)
p := middleware.MustPrincipal(c)
items, err := h.groupRepo.ListUserGroupsVisibleToUser(c.Request.Context(), p.UserID, string(p.Role) == "admin")
if err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
out := make([]dto.UserGroupOption, 0, len(items))
for _, it := range items {
out = append(out, dto.UserGroupOption{UserGroupID: it.ID, Name: it.Name})
}
dto.Write(c, dto.Ok(traceID, dto.ListUserGroupOptionsResp{Items: out}))
}
func (h *UserGroupHandler) Create(c *gin.Context) {
traceID := middleware.GetTraceID(c)
var req dto.CreateUserGroupReq
if err := c.ShouldBindJSON(&req); err != nil || req.Name == "" {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "name"}))
return
}
id, err := h.groupRepo.CreateUserGroup(c.Request.Context(), req.Name, req.Description)
if err != nil {
if strings.Contains(strings.ToLower(err.Error()), "unique") {
dto.Write(c, dto.Err(traceID, dto.CodeConflict, "Name already exists", nil))
return
}
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
dto.Write(c, dto.Ok(traceID, dto.CreateUserGroupResp{ID: id}))
}
func (h *UserGroupHandler) Update(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
var req dto.UpdateUserGroupReq
if err := c.ShouldBindJSON(&req); err != nil || req.Name == "" {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "name"}))
return
}
if err := h.groupRepo.UpdateUserGroup(c.Request.Context(), id, req.Name, req.Description); err != nil {
if strings.Contains(strings.ToLower(err.Error()), "unique") {
dto.Write(c, dto.Err(traceID, dto.CodeConflict, "Name already exists", nil))
return
}
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
dto.Write(c, dto.Ok(traceID, struct{}{}))
}
func (h *UserGroupHandler) Delete(c *gin.Context) {
traceID := middleware.GetTraceID(c)
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
dto.Write(c, dto.Err(traceID, dto.CodeInvalidArgument, "Invalid argument", map[string]any{"field": "id"}))
return
}
if err := h.groupRepo.DeleteUserGroup(c.Request.Context(), id); err != nil {
dto.Write(c, dto.Err(traceID, dto.CodeInternalError, "Internal error", nil))
return
}
dto.Write(c, dto.Ok(traceID, dto.DeleteUserGroupResp{}))
}
+19
View File
@@ -0,0 +1,19 @@
package handler
import (
"context"
)
type UserGroupAdminHandler struct {
create func(ctx context.Context, name, description string) (int64, error)
update func(ctx context.Context, id int64, name, description string) error
delete func(ctx context.Context, id int64) error
}
func NewUserGroupAdminHandler(
create func(ctx context.Context, name, description string) (int64, error),
update func(ctx context.Context, id int64, name, description string) error,
del func(ctx context.Context, id int64) error,
) *UserGroupAdminHandler {
return &UserGroupAdminHandler{create: create, update: update, delete: del}
}
+130
View File
@@ -0,0 +1,130 @@
package middleware
import (
"net/http"
"rttys/internal/domain/identity"
"strings"
"rttys/internal/domain/permission"
"rttys/internal/domain/user"
"rttys/internal/http/dto"
"rttys/internal/store/memory"
"github.com/gin-gonic/gin"
)
const PrincipalKey = "principal"
type Principal struct {
UserID int64 `json:"userId"`
Username string `json:"username"`
DisplayName string `json:"displayName"`
Role identity.Role `json:"role"`
PermissionKeys []string `json:"permissions"`
}
func MustPrincipal(c *gin.Context) Principal {
v, ok := c.Get(PrincipalKey)
if !ok {
panic("principal missing")
}
return v.(Principal)
}
func Auth(sessionStore *memory.SessionStore, userSvc *user.Service, permSvc *permission.Service) gin.HandlerFunc {
return func(c *gin.Context) {
traceID := GetTraceID(c)
// 1) Prefer Bearer token
token := parseBearer(c.GetHeader("Authorization"))
// 2) Fallback to cookie sid
if token == "" {
sid, err := c.Cookie("sid")
if err == nil {
token = strings.TrimSpace(sid)
}
}
// 3) Fallback to Token header (compat with API docs)
if token == "" {
token = strings.TrimSpace(c.GetHeader("Token"))
}
if token == "" {
dto.Write(c, dto.Err(traceID, dto.CodeAuthRequired, "Please login", nil))
c.Abort()
return
}
sess, ok := sessionStore.Get(token)
if !ok {
dto.Write(c, dto.Err(traceID, dto.CodeAuthExpired, "Session expired", nil))
c.Abort()
return
}
u, err := userSvc.GetByID(c.Request.Context(), sess.UserID)
if err != nil || u == nil {
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "Permission denied", nil))
c.Abort()
return
}
keys, _ := permSvc.ListByRole(c.Request.Context(), u.Role)
perms := make([]string, 0, len(keys))
for _, k := range keys {
perms = append(perms, string(k))
}
displayName := u.Description
if strings.TrimSpace(displayName) == "" {
displayName = u.Username
}
c.Set(PrincipalKey, Principal{
UserID: u.ID,
Username: u.Username,
DisplayName: displayName,
Role: u.Role,
PermissionKeys: perms,
})
c.Next()
}
}
// Require checks capability keys (frontend/back-end single source of truth).
func Require(required permission.Key) gin.HandlerFunc {
return func(c *gin.Context) {
traceID := GetTraceID(c)
p := MustPrincipal(c)
for _, k := range p.PermissionKeys {
if k == string(required) {
c.Next()
return
}
}
dto.Write(c, dto.Err(traceID, dto.CodeForbidden, "Permission denied", map[string]any{
"required": string(required),
}))
c.Abort()
}
}
func parseBearer(v string) string {
v = strings.TrimSpace(v)
if v == "" {
return ""
}
parts := strings.SplitN(v, " ", 2)
if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
return strings.TrimSpace(parts[1])
}
return ""
}
// Write wrapper for gin to keep consistent HTTP 200.
func Write(c *gin.Context, payload any) {
c.JSON(http.StatusOK, payload)
}
+31
View File
@@ -0,0 +1,31 @@
package middleware
import (
"crypto/rand"
"encoding/hex"
"github.com/gin-gonic/gin"
)
const TraceIDKey = "traceId"
func Trace() gin.HandlerFunc {
return func(c *gin.Context) {
b := make([]byte, 8)
_, _ = rand.Read(b)
traceID := hex.EncodeToString(b)
c.Set(TraceIDKey, traceID)
c.Writer.Header().Set("X-Trace-Id", traceID)
c.Next()
}
}
func GetTraceID(c *gin.Context) string {
if v, ok := c.Get(TraceIDKey); ok {
if s, ok := v.(string); ok {
return s
}
}
return ""
}
+165
View File
@@ -0,0 +1,165 @@
package http
import (
"net"
"rttys/internal/domain/device"
"rttys/internal/domain/permission"
"rttys/internal/domain/user"
"rttys/internal/http/dto"
"rttys/internal/http/handler"
"rttys/internal/http/middleware"
"rttys/internal/store/memory"
"rttys/internal/store/sqlite"
"rttys/xconfig"
"strings"
"github.com/gin-gonic/gin"
)
type Deps struct {
UserSvc *user.Service
PermSvc *permission.Service
DevSvc *device.Service
GroupRepo *sqlite.GroupRepo
SessionStore *memory.SessionStore
RelationsRepo *sqlite.RelationsRepo
Cfg *xconfig.Config
CloudVersion string
}
func RegisterAPIRoutes(r *gin.Engine, d Deps) {
cfg := d.Cfg
if cfg == nil {
cfg = xconfig.Must()
}
authH := handler.NewAuthHandler(d.UserSvc, d.SessionStore)
meH := handler.NewMeHandler()
devH := handler.NewDeviceHandler(d.DevSvc, d.GroupRepo, d.RelationsRepo)
dgH := handler.NewDeviceGroupHandler(d.GroupRepo, d.RelationsRepo)
ugH := handler.NewUserGroupHandler(d.GroupRepo)
relH := handler.NewRelationsHandler(d.RelationsRepo)
userH := handler.NewUserHandler(d.UserSvc, d.GroupRepo, d.RelationsRepo, d.SessionStore)
// public
r.GET("/auth-config", func(c *gin.Context) {
traceID := middleware.GetTraceID(c)
data := authConfigResp{
LdapEnabled: cfg.LdapEnabled,
LegacyPassword: cfg.Password != "",
OidcEnabled: cfg.OIDCEnabled,
KVMCloudVersion: d.CloudVersion,
}
dto.Write(c, dto.Ok(traceID, data))
})
// public
r.POST("/api/login", authH.Login)
// authed group
api := r.Group("/api")
api.Use(middleware.Auth(d.SessionStore, d.UserSvc, d.PermSvc))
// device script info (requires login)
api.GET("/script-info", func(c *gin.Context) {
traceID := middleware.GetTraceID(c)
// Get domain info
host := c.Request.Host
hostname, _, err := net.SplitHostPort(host)
if err != nil {
hostname = host // Use host directly if no port
}
chosen := hostname
// -------- Reverse proxy mode: force IP ----------
if cfg.ReverseProxyEnabled {
// Reverse proxy mode: always use configured WebRTC IP
if strings.TrimSpace(cfg.WebrtcIP) != "" {
chosen = strings.TrimSpace(cfg.WebrtcIP)
}
} else {
// -------- 3) Original behavior (unchanged) ----------
// 1) If hostname is domain, keep it
// 2) If hostname is IP and cfg.WebrtcIP is set, use cfg.WebrtcIP
if isIP(hostname) && cfg.WebrtcIP != "" {
chosen = cfg.WebrtcIP
}
}
data := scriptInfoResp{
Hostname: chosen, // reuse the same chosen value
Port: cfg.AddrDev,
Token: cfg.Token,
WebrtcIP: chosen, // same as hostname
WebrtcPort: cfg.WebrtcPort,
WebrtcUsername: cfg.WebrtcUsername,
WebrtcPassword: cfg.WebrtcPassword,
}
dto.Write(c, dto.Ok(traceID, data))
})
// auth
api.POST("/logout", middleware.Require(permission.AuthWrite), authH.Logout)
// me
api.GET("/me", middleware.Require(permission.MeRead), meH.GetMe)
// device scope list
api.GET("/devices", middleware.Require(permission.DeviceRead), devH.ListDevices)
api.POST("/devices/move-to-device-group", middleware.Require(permission.DeviceGroupWrite), devH.MoveToDeviceGroup)
api.PUT("/devices/:id", middleware.Require(permission.DeviceWrite), devH.UpdateDevice)
api.DELETE("/devices/:id", middleware.Require(permission.DeviceWrite), devH.DeleteDevice)
// --- users ---
api.GET("/users", middleware.Require(permission.UserRead), userH.ListUsers)
api.POST("/users", middleware.Require(permission.UserWrite), userH.CreateUser)
api.PUT("/users/:id", middleware.Require(permission.UserWrite), userH.UpdateUser)
api.DELETE("/users/:id", middleware.Require(permission.UserWrite), userH.DeleteUser)
// user groups
api.GET("/user-groups", middleware.Require(permission.UserGroupRead), ugH.ListUserGroups)
api.GET("/user-groups/options", middleware.Require(permission.UserGroupRead), ugH.ListOptions)
api.POST("/user-groups", middleware.Require(permission.UserGroupWrite), ugH.Create)
api.PUT("/user-groups/:id", middleware.Require(permission.UserGroupWrite), ugH.Update)
api.DELETE("/user-groups/:id", middleware.Require(permission.UserGroupWrite), ugH.Delete)
// device groups list
api.GET("/device-groups", middleware.Require(permission.DeviceGroupRead), dgH.ListDeviceGroups)
api.GET("/device-groups/options", middleware.Require(permission.DeviceGroupRead), dgH.ListOptions)
api.POST("/device-groups", middleware.Require(permission.DeviceGroupWrite), dgH.Create)
api.PUT("/device-groups/:id", middleware.Require(permission.DeviceGroupWrite), dgH.Update)
api.DELETE("/device-groups/:id", middleware.Require(permission.DeviceGroupWrite), dgH.Delete)
api.POST("/device-groups/:id/devices", middleware.Require(permission.DeviceGroupWrite), dgH.AddDevices)
api.DELETE("/device-groups/:id/devices", middleware.Require(permission.DeviceGroupWrite), dgH.RemoveDevices)
// Relations (cover / set)
api.PUT("/users/:id/user-groups", middleware.Require(permission.UserWrite), relH.SetUserGroups)
api.PUT("/user-groups/:id/device-groups", middleware.Require(permission.UserGroupWrite), relH.SetUserGroupDeviceGroups)
api.PUT("/device-groups/:id/devices", middleware.Require(permission.DeviceGroupWrite), relH.SetDeviceGroupDevices)
}
type authConfigResp struct {
LdapEnabled bool `json:"ldapEnabled"`
LegacyPassword bool `json:"legacyPassword"`
OidcEnabled bool `json:"oidcEnabled"`
KVMCloudVersion string `json:"kvmCloudVersion"`
}
type scriptInfoResp struct {
Hostname string `json:"hostname"`
Port string `json:"port"`
Token string `json:"token"`
WebrtcIP string `json:"webrtcIP"`
WebrtcPort string `json:"webrtcPort"`
WebrtcUsername string `json:"webrtcUsername"`
WebrtcPassword string `json:"webrtcPassword"`
}
func isIP(addr string) bool {
return net.ParseIP(addr) != nil
}
+37
View File
@@ -0,0 +1,37 @@
package legacy
import (
"context"
"rttys/internal/store/sqlite"
"rttys/model"
)
func SaveOrUpdateDeviceMeta(deviceID, mac, description, ip string) error {
repo := sqlite.MustContainer().DeviceMeta
return repo.SaveOrUpdate(context.Background(), deviceID, mac, description, ip)
}
func UpdateDeviceClient(deviceID, client string) error {
repo := sqlite.MustContainer().DeviceMeta
return repo.UpdateClient(context.Background(), deviceID, client)
}
func GetDeviceMetaByDeviceID(deviceID string) (*model.DeviceMeta, error) {
repo := sqlite.MustContainer().DeviceMeta
return repo.GetByDeviceID(context.Background(), deviceID)
}
func UpdateDeviceDescriptionIfEmpty(deviceID, description string) error {
repo := sqlite.MustContainer().DeviceMeta
return repo.UpdateDescriptionIfEmpty(context.Background(), deviceID, description)
}
func DeleteDeviceMetaByDeviceID(deviceID string) error {
repo := sqlite.MustContainer().DeviceMeta
return repo.DeleteByDeviceID(context.Background(), deviceID)
}
func MarkDeviceOffline(deviceID string) error {
repo := sqlite.MustContainer().DeviceMeta
return repo.MarkOffline(context.Background(), deviceID)
}
+376
View File
@@ -0,0 +1,376 @@
/*
* @Author: CU-Jon
* @Date: 2025-09-26 13:28:12 EDT
* @LastEditors: CU-Jon
* @LastEditTime: 2025-09-26 14:02:57 EDT
* @FilePath: \glkvm-cloud\ldap.go
* @Description: LDAP认证模块 (LDAP authentication module)
*/
package ldap
import (
"crypto/tls"
"fmt"
"rttys/xconfig"
"strings"
"time"
"github.com/go-ldap/ldap/v3"
"github.com/rs/zerolog/log"
)
// LDAP认证器结构体 (LDAP authenticator struct)
type LDAPAuthenticator struct {
config *xconfig.Config
}
// 创建新的LDAP认证器 (Create new LDAP authenticator)
func NewLDAPAuthenticator(config *xconfig.Config) *LDAPAuthenticator {
return &LDAPAuthenticator{config: config}
}
// 执行用户LDAP认证 (Perform LDAP authentication for a user)
func (l *LDAPAuthenticator) Authenticate(username, password string) (bool, error) {
if !l.config.LdapEnabled {
return false, fmt.Errorf("LDAP authentication is disabled")
}
if username == "" || password == "" {
return false, fmt.Errorf("username and password are required")
}
// 连接到LDAP服务器 (Connect to LDAP server)
conn, err := l.connect()
if err != nil {
return false, fmt.Errorf("failed to connect to LDAP server: %v", err)
}
defer conn.Close()
// 使用服务账户进行绑定和搜索 (Use service account for binding and searching)
if l.config.LdapBindDN == "" || l.config.LdapBindPassword == "" {
return false, fmt.Errorf("service account credentials are required for LDAP authentication - BindDN empty: %v, BindPassword empty: %v", l.config.LdapBindDN == "", l.config.LdapBindPassword == "")
}
err = conn.Bind(l.config.LdapBindDN, l.config.LdapBindPassword)
if err != nil {
return false, fmt.Errorf("service account bind failed: %v", err)
} // 使用服务账户搜索用户 (Use service account to search for user)
userDN, err := l.findUserDN(conn, username)
if err != nil {
return false, fmt.Errorf("user search failed: %v", err)
}
// 找到用户,现在用用户凭证验证密码 (Found user, now validate password with user credentials)
err = conn.Bind(userDN, password)
if err != nil {
return false, fmt.Errorf("password validation failed: %v", err)
}
// 重新绑定为服务账户以进行授权检查 (Rebind as service account for authorization check)
err = conn.Bind(l.config.LdapBindDN, l.config.LdapBindPassword)
if err != nil {
return false, fmt.Errorf("failed to rebind as service account for authorization: %v", err)
}
// 检查用户授权 (Check user authorization)
authorized, err := l.checkAuthorization(conn, userDN, username)
if err != nil {
return false, fmt.Errorf("authorization check failed: %v", err)
}
if !authorized {
return false, fmt.Errorf("user not authorized")
}
return true, nil
}
// 建立到LDAP服务器的连接 (Establish connection to LDAP server)
func (l *LDAPAuthenticator) connect() (*ldap.Conn, error) {
address := fmt.Sprintf("%s:%d", l.config.LdapServer, l.config.LdapPort)
var conn *ldap.Conn
var err error
if l.config.LdapUseTLS {
// TLS配置 (TLS configuration)
tlsConfig := &tls.Config{
ServerName: l.config.LdapServer,
InsecureSkipVerify: true, // 跳过证书验证以避免自签名证书问题 (Skip certificate verification to avoid self-signed certificate issues)
}
if l.config.LdapPort == 636 {
// 使用LDAPS (直接TLS连接) (Use LDAPS - direct TLS connection)
conn, err = ldap.DialTLS("tcp", address, tlsConfig)
} else {
// 使用StartTLS (先连接再升级到TLS) (Use StartTLS - connect first then upgrade to TLS)
conn, err = ldap.Dial("tcp", address)
if err == nil {
err = conn.StartTLS(tlsConfig)
}
}
} else {
// 使用普通连接 (Use plain connection)
conn, err = ldap.Dial("tcp", address)
}
if err != nil {
return nil, err
}
// 设置超时时间 (Set timeout)
conn.SetTimeout(10 * time.Second)
return conn, nil
}
// 基于用户名搜索用户DN (Search for user DN based on username)
func (l *LDAPAuthenticator) findUserDN(conn *ldap.Conn, username string) (string, error) {
// 准备搜索过滤器 (Prepare search filter)
filter := fmt.Sprintf(l.config.LdapUserFilter, username)
if l.config.LdapUserFilter == "" {
filter = fmt.Sprintf("(uid=%s)", username)
}
// 执行搜索 (Perform search)
searchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, // 无大小限制 (No size limit)
0, // 无时间限制 (No time limit)
false,
filter,
[]string{"dn"},
nil,
)
sr, err := conn.Search(searchRequest)
if err != nil {
return "", err
}
if len(sr.Entries) == 0 {
return "", fmt.Errorf("user not found")
}
if len(sr.Entries) > 1 {
return "", fmt.Errorf("multiple users found")
}
return sr.Entries[0].DN, nil
}
// 基于组或用户列表检查用户是否授权 (Check if user is authorized based on groups or users list)
func (l *LDAPAuthenticator) checkAuthorization(conn *ldap.Conn, userDN, username string) (bool, error) {
// 如果没有配置限制,则允许所有已认证用户 (If no restrictions are configured, allow all authenticated users)
if l.config.LdapAllowedGroups == "" && l.config.LdapAllowedUsers == "" {
return true, nil
}
// 检查允许的用户列表 (Check allowed users list)
if l.config.LdapAllowedUsers != "" {
allowedUsers := strings.Split(strings.TrimSpace(l.config.LdapAllowedUsers), ",")
for _, allowedUser := range allowedUsers {
if strings.TrimSpace(allowedUser) == username {
return true, nil
}
}
}
// 检查允许的组 (Check allowed groups)
if l.config.LdapAllowedGroups != "" {
return l.checkGroupMembership(conn, userDN, username)
}
return false, nil
}
// 检查用户是否属于任何允许的组 (Check if user belongs to any of the allowed groups)
func (l *LDAPAuthenticator) checkGroupMembership(conn *ldap.Conn, userDN, username string) (bool, error) {
allowedGroups := strings.Split(strings.TrimSpace(l.config.LdapAllowedGroups), ",")
for _, group := range allowedGroups {
group = strings.TrimSpace(group)
if group == "" {
continue
}
// 搜索组成员关系 - 尝试不同的常见LDAP组结构 (Search for group membership - try different common LDAP group structures)
isMember, err := l.isGroupMember(conn, userDN, username, group)
if err != nil {
log.Warn().Msgf("Error checking group membership for %s in %s: %v", username, group, err)
continue
}
if isMember {
return true, nil
}
}
return false, nil
}
// 检查用户是否是指定组的成员 (Check if user is a member of the specified group)
func (l *LDAPAuthenticator) isGroupMember(conn *ldap.Conn, userDN, username, groupName string) (bool, error) {
// 首先查找用户的实际DN,因为我们可能使用了UPN格式进行认证 (First find the user's actual DN, as we may have used UPN format for authentication)
actualUserDN, err := l.findActualUserDN(conn, username)
if err != nil {
actualUserDN = userDN // 回退到原始DN (Fallback to original DN)
}
// 尝试不同的常见组搜索模式 (Try different common group search patterns)
// 模式1:通过CN搜索组并检查成员属性 (Pattern 1: Search for group by CN and check member attribute)
groupFilter := fmt.Sprintf("(cn=%s)", groupName)
groupSearchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, 0, false,
groupFilter,
[]string{"member", "memberUid", "uniqueMember"},
nil,
)
sr, err := conn.Search(groupSearchRequest)
if err != nil {
log.Warn().Msgf("Group search failed: %v", err)
return false, err
}
for _, entry := range sr.Entries {
members := entry.GetAttributeValues("member")
memberUids := entry.GetAttributeValues("memberUid")
uniqueMembers := entry.GetAttributeValues("uniqueMember")
// 检查member属性(完整DN) (Check member attribute - full DN)
for _, member := range members {
if member == userDN || member == actualUserDN {
return true, nil
}
// 也检查是否member DN包含用户名 (Also check if member DN contains the username)
if strings.Contains(strings.ToLower(member), strings.ToLower("cn="+username)) {
return true, nil
}
}
// 检查memberUid属性(仅用户名) (Check memberUid attribute - username only)
for _, memberUid := range memberUids {
if memberUid == username {
return true, nil
}
}
// 检查uniqueMember属性(完整DN) (Check uniqueMember attribute - full DN)
for _, uniqueMember := range uniqueMembers {
if uniqueMember == userDN {
return true, nil
}
}
}
// 模式2:通过用户名搜索用户并检查memberOf属性 (Pattern 2: Search for user by username and check memberOf attribute)
// 使用配置的用户过滤器或默认的uid过滤器 (Use configured user filter or default uid filter)
userFilter := fmt.Sprintf(l.config.LdapUserFilter, username)
if l.config.LdapUserFilter == "" {
userFilter = fmt.Sprintf("(uid=%s)", username)
}
userSearchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, 0, false,
userFilter,
[]string{"memberOf", "distinguishedName"},
nil,
)
sr, err = conn.Search(userSearchRequest)
if err != nil {
log.Warn().Msgf("User search for memberOf failed: %v", err)
} else {
for _, entry := range sr.Entries {
memberOfValues := entry.GetAttributeValues("memberOf")
for _, memberOf := range memberOfValues {
if strings.Contains(strings.ToLower(memberOf), strings.ToLower("cn="+groupName)) {
return true, nil
}
}
}
}
return false, nil
}
// 查找用户的实际DN (Find the user's actual DN)
func (l *LDAPAuthenticator) findActualUserDN(conn *ldap.Conn, username string) (string, error) {
// 使用配置的用户过滤器搜索用户 (Search for user using configured user filter)
userFilter := fmt.Sprintf(l.config.LdapUserFilter, username)
if l.config.LdapUserFilter == "" {
userFilter = fmt.Sprintf("(uid=%s)", username)
}
userSearchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, 0, false,
userFilter,
[]string{"distinguishedName"},
nil,
)
sr, err := conn.Search(userSearchRequest)
if err != nil {
return "", err
}
if len(sr.Entries) == 0 {
return "", fmt.Errorf("user not found")
}
if len(sr.Entries) > 1 {
return "", fmt.Errorf("multiple users found")
}
return sr.Entries[0].DN, nil
}
// 执行用户认证,支持LDAP和传统密码认证 (Perform user authentication with LDAP and legacy password support)
func AuthenticateUser(cfg *xconfig.Config, username, password, authMethod string) bool {
success, _ := AuthenticateUserWithError(cfg, username, password, authMethod)
return success
}
// 执行用户认证并返回错误类型,支持LDAP和传统密码认证 (Perform user authentication with error type, supporting LDAP and legacy password authentication)
func AuthenticateUserWithError(cfg *xconfig.Config, username, password, authMethod string) (bool, string) {
// 处理LDAP认证 (Handle LDAP authentication)
if cfg.LdapEnabled && authMethod == "ldap" && username != "" {
ldapAuth := NewLDAPAuthenticator(cfg)
success, err := ldapAuth.Authenticate(username, password)
if err != nil {
log.Error().Msgf("LDAP authentication error: %v", err)
// 检查错误类型以区分认证和授权错误 (Check error type to distinguish between authentication and authorization errors)
if strings.Contains(err.Error(), "user not authorized") {
return false, "authorization"
}
return false, "authentication"
}
return success, ""
}
if authMethod == "legacy" || authMethod == "" {
if cfg.Password == password {
return true, ""
}
return false, "authentication"
}
return false, "authentication"
}
+22
View File
@@ -0,0 +1,22 @@
package password
import "golang.org/x/crypto/bcrypt"
const bcryptCost = 12
// HashPassword hashes a plaintext password using bcrypt.
func HashPassword(pw string) (string, error) {
hash, err := bcrypt.GenerateFromPassword([]byte(pw), bcryptCost)
if err != nil {
return "", err
}
return string(hash), nil
}
// VerifyPassword checks a plaintext password against a stored bcrypt hash.
func VerifyPassword(pw, hash string) bool {
if hash == "" {
return false
}
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(pw)) == nil
}
+16
View File
@@ -0,0 +1,16 @@
package randtoken
import (
"crypto/rand"
"encoding/base64"
)
// New returns a high-entropy random token (URL-safe).
// 32 bytes => 256-bit.
func New() (string, error) {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(b), nil
}
+16 -16
View File
@@ -1,4 +1,4 @@
package main
package proxy
import (
"net"
@@ -17,7 +17,7 @@ type HostInfo struct {
XFPort string // X-Forwarded-Port (raw)
}
func getHostInfoFromRequest(req *http.Request) HostInfo {
func GetHostInfoFromRequest(req *http.Request) HostInfo {
hi := HostInfo{
RawHost: req.Host,
XFHost: req.Header.Get("X-Forwarded-Host"),
@@ -61,17 +61,17 @@ func getHostInfoFromRequest(req *http.Request) HostInfo {
return hi
}
// isIPHost checks whether host is an IP address.
func isIPHost(host string) bool {
// IsIPHost checks whether host is an IP address.
func IsIPHost(host string) bool {
ip := net.ParseIP(strings.TrimSpace(host))
return ip != nil
}
// domainAllowed checks whether host is allowed.
// DomainAllowed checks whether host is allowed.
// Allow:
// - exact match: base
// - subdomain: *.base
func domainAllowed(host, base string) bool {
func DomainAllowed(host, base string) bool {
host = strings.ToLower(strings.TrimSuffix(strings.TrimSpace(host), "."))
base = strings.ToLower(strings.TrimSuffix(strings.TrimSpace(base), "."))
@@ -84,7 +84,7 @@ func domainAllowed(host, base string) bool {
return strings.HasSuffix(host, "."+base)
}
// buildRedirectHost removes the first label of the hostname and prepends devid.
// BuildRedirectHost removes the first label of the hostname and prepends devid.
// Rules:
// - "www.example.com" -> "devid.example.com"
// - "www.l1.example.com" -> "devid.l1.example.com"
@@ -93,7 +93,7 @@ func domainAllowed(host, base string) bool {
// - Single label / abnormal cases -> "devid." + hostname (fallback)
//
// The input hostname must be a pure hostname without port.
func buildRedirectHost(hostname, devid string) string {
func BuildRedirectHost(hostname, devid string) string {
// Allow FQDN with trailing dot like "example.com."
hostname = strings.TrimSuffix(hostname, ".")
@@ -112,7 +112,7 @@ func buildRedirectHost(hostname, devid string) string {
case 0:
return devid // extreme case: just return devid
case 1:
// Single label (e.g., "localhost") — keep original as suffix
// Single label (e.g., "localhost") keep original as suffix
return devid + "." + labels[0]
default:
// >=2: drop the leftmost label
@@ -121,7 +121,7 @@ func buildRedirectHost(hostname, devid string) string {
}
}
func joinHostPortIfNeeded(host, scheme, port string) string {
func JoinHostPortIfNeeded(host, scheme, port string) string {
if port == "" {
return host
}
@@ -132,7 +132,7 @@ func joinHostPortIfNeeded(host, scheme, port string) string {
return net.JoinHostPort(host, port)
}
func buildRedirectLocation(scheme, hostPort, path, sid string) string {
func BuildRedirectLocation(scheme, hostPort, path, sid string) string {
if path == "" {
path = "/"
}
@@ -142,16 +142,16 @@ func buildRedirectLocation(scheme, hostPort, path, sid string) string {
Path: path,
}
q := u.Query()
q.Set("sid", sid)
q.Set("rttysid", sid)
u.RawQuery = q.Encode()
return u.String()
}
// getRequestHostInfo extracts domain(host), port and scheme(proto) from request headers.
// GetRequestHostInfo extracts domain(host), port and scheme(proto) from request headers.
// Priority:
// 1) X-Forwarded-Host / X-Forwarded-Proto / X-Forwarded-Port (reverse proxy)
// 2) Host header / TLS info
func getRequestHostInfo(req *http.Request) (host string, port string, proto string) {
func GetRequestHostInfo(req *http.Request) (host string, port string, proto string) {
// 1) Reverse-proxy headers
xfh := strings.TrimSpace(req.Header.Get("X-Forwarded-Host"))
xfp := strings.TrimSpace(req.Header.Get("X-Forwarded-Proto"))
@@ -202,13 +202,13 @@ func getRequestHostInfo(req *http.Request) (host string, port string, proto stri
return host, port, proto
}
// extractDeviceIDFromHost extracts deviceId from hostname.
// ExtractDeviceIDFromHost extracts deviceId from hostname.
// Rules:
// - IP address -> ("", false)
// - lv99862.example.com -> ("lv99862", true)
// - lv99862.l1.example.com -> ("lv99862", true)
// - localhost / single label -> ("localhost", true)
func extractDeviceIDFromHost(host string) (string, bool) {
func ExtractDeviceIDFromHost(host string) (string, bool) {
host = strings.TrimSpace(host)
if host == "" {
return "", false
+407
View File
@@ -0,0 +1,407 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"context"
"crypto/tls"
"fmt"
"io/fs"
"net"
"net/http"
"path"
"rttys/internal/domain/device"
"rttys/internal/domain/permission"
"rttys/internal/domain/user"
httpx "rttys/internal/http"
"rttys/internal/http/middleware"
"rttys/internal/pkg/password"
"rttys/internal/proxy"
"rttys/internal/store/memory"
"rttys/internal/store/sqlite"
"rttys/ui"
"rttys/xconfig"
"strings"
"time"
"github.com/gin-contrib/cors"
"github.com/gin-gonic/gin"
"github.com/rs/zerolog/log"
"gorm.io/gorm"
)
type AppContainer struct {
DB *sqlite.AppDB
DeviceMetaRepo *sqlite.DeviceMetaRepo
UserSvc *user.Service
}
var sessionStore *memory.SessionStore
const defaultDBPath = "/home/database/glkvm-cloud.db"
func InitAppContainer(r *gin.Engine) (*AppContainer, error) {
ctx := context.Background()
cfg := xconfig.Must()
// --- DB ---
appDB, err := sqlite.Open(ctx, sqlite.Options{
DSN: defaultDBPath,
MaxOpenConns: 1,
MaxIdleConns: 1,
LogSQL: true,
})
if err != nil {
log.Fatal().Err(err).Msg("open sqlite failed")
}
deviceMetaRepo := sqlite.NewDeviceMetaRepo(appDB.Gorm())
if err := sqlite.InitSchema(ctx, appDB.SQL(), "/home/database/schema.sql"); err != nil {
log.Fatal().Err(err).Msg("init schema failed")
}
if err := ensureAdminUser(ctx, appDB.Gorm(), cfg.AdminName, cfg.Password); err != nil {
log.Fatal().Err(err).Msg("ensure admin user failed")
}
// --- Repos & Services ---
userRepo := sqlite.NewUserRepo(appDB.Gorm())
groupRepo := sqlite.NewGroupRepo(appDB.Gorm())
deviceRepo := sqlite.NewDeviceRepo(appDB.Gorm())
relationsRepo := sqlite.NewRelationsRepo(appDB.Gorm())
userSvc := user.NewService(userRepo)
devSvc := device.NewService(deviceRepo, groupRepo)
permRepo := memory.NewPermissionRepo() // permissions stay in-memory
permSvc := permission.NewService(permRepo)
sessionStore = memory.NewSessionStore(cfg.AuthSessionTTL)
httpx.RegisterAPIRoutes(r, httpx.Deps{
UserSvc: userSvc,
PermSvc: permSvc,
DevSvc: devSvc,
GroupRepo: groupRepo,
SessionStore: sessionStore,
RelationsRepo: relationsRepo,
Cfg: cfg,
CloudVersion: KVMCloudVersion,
})
c := &AppContainer{
DB: appDB,
DeviceMetaRepo: deviceMetaRepo,
UserSvc: userSvc,
}
return c, nil
}
func ensureAdminUser(ctx context.Context, db *gorm.DB, adminName, plainPassword string) error {
if db == nil {
return fmt.Errorf("db is nil")
}
hash, err := password.HashPassword(plainPassword)
if err != nil {
return err
}
// First, rename the existing system admin user to the configured name (if changed).
// This handles the case where the admin username was previously "admin" (or another name)
// and the user now wants a different username via RTTYS_ADMIN_NAME.
if err := db.WithContext(ctx).Exec(
`UPDATE users SET username = ? WHERE is_system = 1 AND role = 'admin' AND username != ?`,
adminName, adminName,
).Error; err != nil {
return fmt.Errorf("rename system admin user: %w", err)
}
// Upsert: create the admin user if not exists, or update password/role/status.
// On conflict, also set description to 'System Administrator' if it is currently empty.
return db.WithContext(ctx).Exec(
`INSERT INTO users (username, description, password_hash, role, status, is_system)
VALUES (?, 'System Administrator', ?, 'admin', 'active', 1)
ON CONFLICT(username) DO UPDATE SET
password_hash=excluded.password_hash,
role='admin',
status='active',
is_system=1,
description=CASE WHEN (description IS NULL OR description = '') THEN 'System Administrator' ELSE description END`,
adminName, hash,
).Error
}
func (srv *RttyServer) ListenAPI() error {
cfg := &srv.cfg
gin.SetMode(gin.ReleaseMode)
r := gin.New()
r.Use(gin.Recovery())
r.Use(middleware.Trace())
r.Use(func(c *gin.Context) {
hi := proxy.GetHostInfoFromRequest(c.Request)
host := hi.Host
allowedHost := cfg.WebUIHost
// If WebUIHost is configured, enforce host validation
if allowedHost != "" && !proxy.IsIPHost(host) {
if !proxy.DomainAllowed(host, allowedHost) {
html := generateErrorHTML("invalid")
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(html))
c.Abort()
return
}
}
c.Next()
})
if cfg.AllowOrigins {
log.Debug().Msg("Allow all origins")
r.Use(cors.Default())
}
authorized := r.Group("/", func(c *gin.Context) {
if !cfg.LocalAuth && isLocalRequest(c) {
return
}
if !httpAuth(cfg, c) {
c.AbortWithStatus(http.StatusUnauthorized)
return
}
})
authorized.GET("/connect/:devid", func(c *gin.Context) {
if !callUserHookUrl(cfg, c) {
c.Status(http.StatusForbidden)
return
}
if c.GetHeader("Upgrade") != "websocket" {
group := c.Query("group")
devid := c.Param("devid")
if dev := srv.GetDevice(group, devid); dev == nil {
c.Redirect(http.StatusFound, "/error/offline")
return
}
url := "/rtty/" + devid
if group != "" {
url += "?group=" + group
}
c.Redirect(http.StatusFound, url)
} else {
handleUserConnection(srv, c)
}
})
authorized.POST("/cmd/:devid", func(c *gin.Context) {
if !callUserHookUrl(cfg, c) {
c.Status(http.StatusForbidden)
return
}
cmdInfo := &CommandReqInfo{}
err := c.BindJSON(&cmdInfo)
if err != nil || cmdInfo.Cmd == "" || cmdInfo.Username == "" {
cmdErrResp(c, rttyCmdErrInvalid)
return
}
dev := srv.GetDevice(c.Query("group"), c.Param("devid"))
if dev == nil {
cmdErrResp(c, rttyCmdErrOffline)
return
}
dev.handleCmdReq(c, cmdInfo)
})
authorized.Any("/web/:devid/:proto/:addr/*path", func(c *gin.Context) {
httpProxyRedirect(srv, c, "")
})
container, err := InitAppContainer(r)
if err != nil {
return err
}
defer container.DB.Close()
sqlite.SetContainer(&sqlite.Container{
Gorm: container.DB.Gorm(),
DeviceMeta: sqlite.NewDeviceMetaRepo(container.DB.Gorm()),
})
// ===== 添加OIDC路由 =====
RegisterOIDCRoutes(r, cfg, container.UserSvc)
fs, err := fs.Sub(ui.StaticFS, "dist")
if err != nil {
return err
}
root := http.FS(fs)
fh := http.FileServer(root)
r.NoRoute(func(c *gin.Context) {
if strings.HasPrefix(c.Request.URL.Path, "/api/") {
c.JSON(http.StatusNotFound, gin.H{"code": 404, "msg": "not found"})
return
}
upath := path.Clean(c.Request.URL.Path)
if strings.HasSuffix(upath, ".js") || strings.HasSuffix(upath, ".css") {
if strings.Contains(c.Request.Header.Get("Accept-Encoding"), "gzip") {
f, err := root.Open(upath + ".gz")
if err == nil {
f.Close()
c.Request.URL.Path += ".gz"
if strings.HasSuffix(upath, ".js") {
c.Writer.Header().Set("Content-Type", "application/javascript")
} else if strings.HasSuffix(upath, ".css") {
c.Writer.Header().Set("Content-Type", "text/css")
}
c.Writer.Header().Set("Content-Encoding", "gzip")
}
}
} else if upath != "/" {
f, err := root.Open(upath)
if err != nil {
c.Request.URL.Path = "/"
r.HandleContext(c)
return
}
defer f.Close()
}
fh.ServeHTTP(c.Writer, c.Request)
})
ln, err := net.Listen("tcp", cfg.AddrUser)
if err != nil {
return err
}
defer ln.Close()
// If we're behind a reverse proxy (TLS terminated by nginx), never enable TLS here.
enableTLS := !cfg.ReverseProxyEnabled && cfg.SslCert != "" && cfg.SslKey != ""
if enableTLS {
crt, err := tls.LoadX509KeyPair(cfg.SslCert, cfg.SslKey)
if err != nil {
log.Fatal().Msg(err.Error())
}
tlsConfig := &tls.Config{Certificates: []tls.Certificate{crt}}
ln = tls.NewListener(ln, tlsConfig)
}
log.Info().Msgf("Listen users on: %s", ln.Addr().(*net.TCPAddr))
return r.RunListener(ln)
}
func callUserHookUrl(cfg *xconfig.Config, c *gin.Context) bool {
if cfg.UserHookUrl == "" {
return true
}
upath := c.Request.URL.RawPath
// Create HTTP request with original headers
req, err := http.NewRequest("GET", cfg.UserHookUrl, nil)
if err != nil {
log.Error().Err(err).Msgf("create hook request for \"%s\" fail", upath)
return false
}
// Copy all headers from original request
for key, values := range c.Request.Header {
lowerKey := strings.ToLower(key)
if lowerKey == "upgrade" || lowerKey == "connection" || lowerKey == "accept-encoding" {
continue
}
for _, value := range values {
req.Header.Add(key, value)
}
}
// Add custom headers for hook identification
req.Header.Set("X-Rttys-Hook", "true")
req.Header.Set("X-Original-Method", c.Request.Method)
req.Header.Set("X-Original-URL", c.Request.URL.String())
cli := &http.Client{
Timeout: 3 * time.Second,
}
resp, err := cli.Do(req)
if err != nil {
log.Error().Err(err).Msgf("call user hook url for \"%s\" fail", upath)
return false
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
log.Error().Msgf("call user hook url for \"%s\", StatusCode: %d", upath, resp.StatusCode)
return false
}
return true
}
func isLocalRequest(c *gin.Context) bool {
addr, _ := net.ResolveTCPAddr("tcp", c.Request.RemoteAddr)
return addr.IP.IsLoopback()
}
func httpAuth(cfg *xconfig.Config, c *gin.Context) bool {
if !cfg.LocalAuth && isLocalRequest(c) {
return true
}
// Keep legacy behavior: if password is not set, no auth required
if cfg.Password == "" {
return true
}
sid, err := c.Cookie("sid")
if err != nil || strings.TrimSpace(sid) == "" {
return false
}
sid = strings.TrimSpace(sid)
// New session-based auth
_, ok := sessionStore.Get(sid)
return ok
}
+86
View File
@@ -0,0 +1,86 @@
package server
import (
"encoding/json"
"os"
"runtime"
xlog "rttys/log"
"rttys/xconfig"
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
)
func RunFromEnv() error {
defaultLogPath := "/var/log/rttys.log"
if runtime.GOOS == "windows" {
defaultLogPath = "rttys.log"
}
defer LogPanic()
cfg := xconfig.Config{
AddrDev: ":5912",
AddrUser: ":5913",
LocalAuth: true,
LogPath: defaultLogPath,
LogLevel: "info",
}
confPath := os.Getenv("RTTYS_CONF")
if confPath == "" {
if _, err := os.Stat("rttys.conf"); err == nil {
confPath = "rttys.conf"
}
}
if err := cfg.Load(confPath); err != nil {
return err
}
xlog.SetPath(cfg.LogPath)
switch cfg.LogLevel {
case "debug":
zerolog.SetGlobalLevel(zerolog.DebugLevel)
case "warn":
zerolog.SetGlobalLevel(zerolog.WarnLevel)
case "error":
zerolog.SetGlobalLevel(zerolog.ErrorLevel)
default:
zerolog.SetGlobalLevel(zerolog.InfoLevel)
}
if cfg.Verbose {
xlog.Verbose()
}
log.Info().Msg("Go Version: " + runtime.Version())
log.Info().Msgf("Go OS/Arch: %s/%s", runtime.GOOS, runtime.GOARCH)
log.Info().Msg("Rttys Version: " + RttysVersion)
if GitCommit != "" {
log.Info().Msg("Git Commit: " + GitCommit)
}
if BuildTime != "" {
log.Info().Msg("Build Time: " + BuildTime)
}
if runtime.GOOS != "windows" {
go StartSignalHandler()
}
{
importJSON, _ := json.MarshalIndent(cfg, "", " ")
log.Info().Msg("==== Loaded Configuration ====")
log.Info().Msg(string(importJSON))
log.Info().Msg("==============================")
}
xconfig.InitGlobal(&cfg)
srv := New(cfg)
return srv.Run()
}
+148 -148
View File
@@ -1,148 +1,148 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
import (
"context"
"encoding/json"
"net/http"
"rttys/utils"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/rs/zerolog/log"
"github.com/valyala/bytebufferpool"
)
type CommandReq struct {
cancel context.CancelFunc
acked bool
c *gin.Context
}
type CommandReqInfo struct {
Cmd string `json:"cmd"`
Username string `json:"username"`
Params []string `json:"params"`
}
type CommandRespInfo struct {
Token string `json:"token"`
Attrs json.RawMessage `json:"attrs"`
}
const (
rttyCmdErrInvalid = 1001
rttyCmdErrOffline = 1002
rttyCmdErrTimeout = 1003
)
var cmdErrMsg = map[int]string{
rttyCmdErrInvalid: "invalid format",
rttyCmdErrOffline: "device offline",
rttyCmdErrTimeout: "timeout",
}
func (dev *Device) handleCmdReq(c *gin.Context, info *CommandReqInfo) {
ctx, cancel := context.WithCancel(dev.ctx)
defer cancel()
req := &CommandReq{
cancel: cancel,
c: c,
}
token := utils.GenUniqueID()
msg := bytebufferpool.Get()
defer bytebufferpool.Put(msg)
BpWriteCString(msg, info.Username)
BpWriteCString(msg, info.Cmd)
BpWriteCString(msg, token)
msg.WriteByte(byte(len(info.Params)))
for _, param := range info.Params {
BpWriteCString(msg, param)
}
log.Debug().Msgf("send cmd request for device '%s', token '%s'", dev.id, token)
err := dev.WriteMsg(msgTypeCmd, "", msg.Bytes())
if err != nil {
cmdErrResp(c, rttyCmdErrOffline)
return
}
waitTime := CommandTimeout
wait := c.Query("wait")
if wait != "" {
waitTime, _ = strconv.Atoi(wait)
}
if waitTime == 0 {
c.Status(http.StatusOK)
return
}
dev.commands.Store(token, req)
if waitTime < 0 || waitTime > CommandTimeout {
waitTime = CommandTimeout
}
tmr := time.NewTimer(time.Second * time.Duration(waitTime))
log.Debug().Msgf("wait for cmd response for device '%s', token '%s', waitTime %ds", dev.id, token, waitTime)
select {
case <-tmr.C:
cmdErrResp(c, rttyCmdErrTimeout)
case <-ctx.Done():
}
dev.commands.Delete(token)
if !req.acked {
cmdErrResp(c, rttyCmdErrOffline)
}
log.Debug().Msgf("handle cmd request for device '%s', token '%s' done", dev.id, token)
}
func cmdErrResp(c *gin.Context, err int) {
c.JSON(http.StatusOK, gin.H{
"err": err,
"msg": cmdErrMsg[err],
})
}
func BpWriteCString(bb *bytebufferpool.ByteBuffer, s string) {
bb.WriteString(s)
bb.WriteByte(0)
}
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"context"
"encoding/json"
"net/http"
"rttys/utils"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/rs/zerolog/log"
"github.com/valyala/bytebufferpool"
)
type CommandReq struct {
cancel context.CancelFunc
acked bool
c *gin.Context
}
type CommandReqInfo struct {
Cmd string `json:"cmd"`
Username string `json:"username"`
Params []string `json:"params"`
}
type CommandRespInfo struct {
Token string `json:"token"`
Attrs json.RawMessage `json:"attrs"`
}
const (
rttyCmdErrInvalid = 1001
rttyCmdErrOffline = 1002
rttyCmdErrTimeout = 1003
)
var cmdErrMsg = map[int]string{
rttyCmdErrInvalid: "invalid format",
rttyCmdErrOffline: "device offline",
rttyCmdErrTimeout: "timeout",
}
func (dev *Device) handleCmdReq(c *gin.Context, info *CommandReqInfo) {
ctx, cancel := context.WithCancel(dev.ctx)
defer cancel()
req := &CommandReq{
cancel: cancel,
c: c,
}
token := utils.GenUniqueID()
msg := bytebufferpool.Get()
defer bytebufferpool.Put(msg)
BpWriteCString(msg, info.Username)
BpWriteCString(msg, info.Cmd)
BpWriteCString(msg, token)
msg.WriteByte(byte(len(info.Params)))
for _, param := range info.Params {
BpWriteCString(msg, param)
}
log.Debug().Msgf("send cmd request for device '%s', token '%s'", dev.id, token)
err := dev.WriteMsg(msgTypeCmd, "", msg.Bytes())
if err != nil {
cmdErrResp(c, rttyCmdErrOffline)
return
}
waitTime := CommandTimeout
wait := c.Query("wait")
if wait != "" {
waitTime, _ = strconv.Atoi(wait)
}
if waitTime == 0 {
c.Status(http.StatusOK)
return
}
dev.commands.Store(token, req)
if waitTime < 0 || waitTime > CommandTimeout {
waitTime = CommandTimeout
}
tmr := time.NewTimer(time.Second * time.Duration(waitTime))
log.Debug().Msgf("wait for cmd response for device '%s', token '%s', waitTime %ds", dev.id, token, waitTime)
select {
case <-tmr.C:
cmdErrResp(c, rttyCmdErrTimeout)
case <-ctx.Done():
}
dev.commands.Delete(token)
if !req.acked {
cmdErrResp(c, rttyCmdErrOffline)
}
log.Debug().Msgf("handle cmd request for device '%s', token '%s' done", dev.id, token)
}
func cmdErrResp(c *gin.Context, err int) {
c.JSON(http.StatusOK, gin.H{
"err": err,
"msg": cmdErrMsg[err],
})
}
func BpWriteCString(bb *bytebufferpool.ByteBuffer, s string) {
bb.WriteString(s)
bb.WriteByte(0)
}
+706
View File
@@ -0,0 +1,706 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"bufio"
"bytes"
"context"
"crypto/tls"
"encoding/binary"
"fmt"
"io"
"net"
"net/http"
"rttys/internal/legacy"
"strings"
"sync"
"time"
"rttys/utils"
"github.com/gorilla/websocket"
jsoniter "github.com/json-iterator/go"
"github.com/rs/zerolog/log"
"github.com/valyala/bytebufferpool"
)
type DeviceInfo struct {
ID string `json:"id"`
Mac string `json:"mac"`
Connected uint32 `json:"connected"`
Uptime uint32 `json:"uptime"`
Desc string `json:"description"`
Proto uint8 `json:"proto"`
IPaddr string `json:"ipaddr"`
}
type Device struct {
group string
id string
proto uint8
desc string
timestamp int64
uptime uint32
token string
heartbeat time.Duration
clientInfoMu sync.RWMutex
clientInfo []byte
users sync.Map
pending sync.Map
commands sync.Map
https sync.Map
conn net.Conn
br *bufio.Reader
readBuf []byte
close sync.Once
ctx context.Context
cancel context.CancelFunc
}
const (
msgTypeRegister = byte(iota)
msgTypeLogin
msgTypeLogout
msgTypeTermData
msgTypeWinsize
msgTypeCmd
msgTypeHeartbeat
msgTypeFile
msgTypeHttp
msgTypeAck
)
// Custom extension message types (keep out of upstream range to avoid conflicts).
const (
msgTypeDeviceInfo = byte(0xF0)
)
const (
msgTypeFileSend = byte(iota)
msgTypeFileRecv
msgTypeFileInfo
msgTypeFileData
msgTypeFileAck
msgTypeFileAbort
)
const (
msgRegAttrHeartbeat = iota
msgRegAttrDevid
msgRegAttrDescription
msgRegAttrToken
msgRegAttrGroup
)
const (
msgHeartbeatAttrUptime = iota
)
const (
devRegErrUnsupportedProto = iota + 1
devRegErrInvalidToken
devRegErrHookFailed
devRegErrIdConflicting
)
const (
RttyProtoRequired uint8 = 3
WaitRegistTimeout = 5 * time.Second
DefaultHeartbeat = 5 * time.Second
TermLoginTimeout = 5 * time.Second
CommandTimeout = 30
MaxDeviceInfoSize = 8 * 1024
)
var DevRegErrMsg = map[byte]string{
0: "Success",
devRegErrUnsupportedProto: "Unsupported protocol",
devRegErrInvalidToken: "Invalid token",
devRegErrHookFailed: "Hook failed",
devRegErrIdConflicting: "ID conflict",
}
var DeviceMsgHandlers = map[byte]func(*Device, []byte) error{
msgTypeHeartbeat: handleHeartbeatMsg,
msgTypeLogin: handleLoginMsg,
msgTypeLogout: handleLogoutMsg,
msgTypeTermData: handleTermDataMsg,
msgTypeFile: handleFileMsg,
msgTypeCmd: handleCmdMsg,
msgTypeHttp: handleHttpMsg,
msgTypeDeviceInfo: handleDeviceInfoMsg,
}
func (srv *RttyServer) ListenDevices() {
cfg := &srv.cfg
ln, err := net.Listen("tcp", cfg.AddrDev)
if err != nil {
log.Fatal().Msg(err.Error())
}
defer ln.Close()
if cfg.SslCert != "" && cfg.SslKey != "" {
crt, err := tls.LoadX509KeyPair(cfg.SslCert, cfg.SslKey)
if err != nil {
log.Fatal().Msg(err.Error())
}
tlsConfig := &tls.Config{
GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) {
// 忽略 SNI,始终返回唯一证书
return &crt, nil
},
}
ln = tls.NewListener(ln, tlsConfig)
}
log.Info().Msgf("Listen devices on: %s", ln.Addr().(*net.TCPAddr))
for {
conn, err := ln.Accept()
if err != nil {
log.Error().Msg(err.Error())
continue
}
go handleDeviceConnection(srv, conn)
}
}
func handleDeviceConnection(srv *RttyServer, conn net.Conn) {
defer LogPanic()
dev := &Device{
conn: conn,
heartbeat: DefaultHeartbeat,
timestamp: time.Now().Unix(),
br: bufio.NewReader(conn),
}
defer dev.Close(srv)
dev.ctx, dev.cancel = context.WithCancel(context.Background())
log.Debug().Msgf("new device '%s' connected", conn.RemoteAddr())
conn.SetReadDeadline(time.Now().Add(WaitRegistTimeout))
typ, data, err := dev.ReadMsg()
if err != nil {
log.Error().Msgf("read register msg fail: %v", err)
return
}
if typ != msgTypeRegister {
log.Error().Msg("register msg expected first")
return
}
if !dev.ParseRegister(data) {
log.Error().Msg("invalid device info")
return
}
code := dev.Register(srv)
err = dev.WriteMsg(msgTypeRegister, "", append([]byte{code}, DevRegErrMsg[code]...))
if err != nil {
log.Printf("send register to device '%s' fail: %v", dev.id, err)
return
}
if code != 0 {
return
}
deviceRemoteIP := ""
if addr, ok := dev.conn.RemoteAddr().(*net.TCPAddr); ok {
deviceRemoteIP = addr.IP.String()
} else if host, _, err := net.SplitHostPort(dev.conn.RemoteAddr().String()); err == nil {
deviceRemoteIP = host
}
log.Info().Msgf("device '%s' registered, group '%s' proto %d, heartbeat %v, remoteIP '%s'",
dev.id, dev.group, dev.proto, dev.heartbeat, deviceRemoteIP)
// 2. Load existing metadata by device_id
description := ""
meta, err := legacy.GetDeviceMetaByDeviceID(dev.id)
if err == nil && meta != nil {
description = meta.Description
}
if err := legacy.SaveOrUpdateDeviceMeta(
dev.id,
dev.desc, // device register mac info with desc filed
description,
deviceRemoteIP,
); err != nil {
return
}
for {
conn.SetReadDeadline(time.Now().Add(dev.heartbeat * 3 / 2))
typ, data, err = dev.ReadMsg()
if err != nil {
if err != io.EOF {
log.Error().Msgf("read msg from device '%s' fail: %v", dev.id, err)
}
return
}
log.Debug().Msgf("device msg %s from device %s", msgTypeName(typ), dev.id)
handler, ok := DeviceMsgHandlers[typ]
if !ok {
log.Error().Msgf("unexpected message '%s' from device '%s'", msgTypeName(typ), dev.id)
return
}
err = handler(dev, data)
if err != nil {
log.Error().Msg(err.Error())
return
}
}
}
func msgTypeName(typ byte) string {
switch typ {
case msgTypeRegister:
return "register"
case msgTypeLogin:
return "login"
case msgTypeLogout:
return "logout"
case msgTypeTermData:
return "termdata"
case msgTypeWinsize:
return "winsize"
case msgTypeCmd:
return "cmd"
case msgTypeHeartbeat:
return "heartbeat"
case msgTypeFile:
return "file"
case msgTypeHttp:
return "http"
case msgTypeAck:
return "ack"
case msgTypeDeviceInfo:
return "deviceinfo"
default:
return fmt.Sprintf("unknown(%d)", typ)
}
}
func (dev *Device) ReadMsg() (byte, []byte, error) {
head := make([]byte, 3)
br := dev.br
_, err := io.ReadFull(br, head)
if err != nil {
return 0, nil, err
}
typ := head[0]
msgLen := binary.BigEndian.Uint16(head[1:])
if cap(dev.readBuf) < int(msgLen) {
dev.readBuf = make([]byte, msgLen)
} else {
dev.readBuf = dev.readBuf[:msgLen]
}
_, err = io.ReadFull(br, dev.readBuf)
if err != nil {
return 0, nil, err
}
return typ, dev.readBuf, nil
}
func (dev *Device) WriteMsg(typ byte, sid string, data []byte) error {
bb := bytebufferpool.Get()
defer bytebufferpool.Put(bb)
b := []byte{typ, 0, 0}
binary.BigEndian.PutUint16(b[1:], uint16(len(sid)+len(data)))
bb.Write(b)
bb.WriteString(sid)
bb.Write(data)
_, err := bb.WriteTo(dev.conn)
return err
}
func (dev *Device) WriteFileMsg(typ byte, sid string, fileType byte, data []byte) error {
bb := bytebufferpool.Get()
defer bytebufferpool.Put(bb)
bb.WriteByte(fileType)
bb.Write(data)
return dev.WriteMsg(typ, sid, bb.Bytes())
}
func (dev *Device) Close(srv *RttyServer) {
dev.close.Do(func() {
log.Error().Msgf("device '%s' disconnected", dev.id)
srv.DelDevice(dev)
if dev.id != "" {
_ = legacy.MarkDeviceOffline(dev.id)
}
dev.cancel()
dev.conn.Close()
})
}
func (dev *Device) ParseRegister(b []byte) bool {
if len(b) < 1 {
return false
}
dev.proto = b[0]
if dev.proto > 4 {
attrs := utils.ParseTLV(b[1:])
if attrs == nil {
return false
}
for typ, val := range attrs {
switch typ {
case msgRegAttrHeartbeat:
dev.heartbeat = time.Duration(val[0]) * time.Second
case msgRegAttrDevid:
dev.id = string(val)
case msgRegAttrDescription:
dev.desc = string(val)
case msgRegAttrToken:
dev.token = string(val)
case msgRegAttrGroup:
dev.group = string(val)
}
}
return true
}
b = b[1:]
fields := bytes.Split(b, []byte{0})
if len(fields) < 3 {
return false
}
dev.id = string(fields[0])
dev.desc = string(fields[1])
dev.token = string(fields[2])
return true
}
func (dev *Device) Register(srv *RttyServer) byte {
cfg := &srv.cfg
if dev.proto < RttyProtoRequired {
log.Error().Msgf("minimum proto required %d, found %d for device '%s'", RttyProtoRequired, dev.proto, dev.id)
return devRegErrHookFailed
}
log.Info().Msgf("cfg.Token:%s,dev.token:%s", cfg.Token, dev.token)
if cfg.Token != "" && dev.token != cfg.Token {
log.Error().Msgf("invalid token for device '%s'", dev.id)
return devRegErrInvalidToken
}
devHookUrl := cfg.DevHookUrl
if devHookUrl != "" {
cli := &http.Client{
Timeout: 3 * time.Second,
}
data := fmt.Sprintf(`{"group":"%s", "devid":"%s", "token":"%s"}`, dev.group, dev.id, dev.token)
resp, err := cli.Post(devHookUrl, "application/json", strings.NewReader(data))
if err != nil {
log.Error().Msgf("call device hook url fail for device %s: %v", dev.id, err)
return devRegErrHookFailed
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
log.Error().Msgf("call device hook url for device '%s', StatusCode: %d", dev.id, resp.StatusCode)
return devRegErrHookFailed
}
}
if !srv.AddDevice(dev) {
return devRegErrIdConflicting
}
return 0
}
func (dev *Device) setClientInfo(data []byte) {
dev.clientInfoMu.Lock()
dev.clientInfo = append(dev.clientInfo[:0], data...)
dev.clientInfoMu.Unlock()
}
func handleDeviceInfoMsg(dev *Device, data []byte) error {
if len(data) == 0 {
log.Warn().Msgf("device '%s' sent empty client info", dev.id)
return nil
}
if len(data) > MaxDeviceInfoSize {
log.Warn().Msgf("device '%s' client info too large: %d bytes", dev.id, len(data))
return nil
}
var payload struct {
Client string `json:"client"`
OS string `json:"os"`
Hostname string `json:"hostname"`
LocalIP string `json:"local_ip"`
}
if err := jsoniter.Unmarshal(data, &payload); err != nil {
log.Warn().Msgf("device '%s' client info invalid json: %v", dev.id, err)
return nil
}
dev.setClientInfo(data)
if payload.Client != "" {
if err := legacy.UpdateDeviceClient(dev.id, payload.Client); err != nil {
log.Warn().Err(err).Msgf("device '%s' update client info failed", dev.id)
}
}
// Auto-fill description with os/hostname/local_ip if not already set by user
if payload.OS != "" || payload.Hostname != "" {
parts := make([]string, 0, 3)
if payload.OS != "" {
parts = append(parts, payload.OS)
}
if payload.Hostname != "" {
parts = append(parts, payload.Hostname)
}
if payload.LocalIP != "" {
parts = append(parts, payload.LocalIP)
}
desc := strings.Join(parts, " / ")
if err := legacy.UpdateDeviceDescriptionIfEmpty(dev.id, desc); err != nil {
log.Warn().Err(err).Msgf("device '%s' auto-fill description failed", dev.id)
}
}
log.Info().Msgf("device '%s' client info: %s", dev.id, string(data))
log.Debug().Msgf("device '%s' client info updated", dev.id)
return nil
}
func handleHeartbeatMsg(dev *Device, data []byte) error {
if !parseHeartbeat(dev, data) {
return fmt.Errorf("invalid heartbeat msg from device '%s'", dev.id)
}
return dev.WriteMsg(msgTypeHeartbeat, "", nil)
}
func parseHeartbeat(dev *Device, data []byte) bool {
if dev.proto > 4 {
attrs := utils.ParseTLV(data)
if attrs == nil {
return false
}
for typ, val := range attrs {
switch typ {
case msgHeartbeatAttrUptime:
dev.uptime = binary.BigEndian.Uint32(val)
}
}
} else {
if len(data) < 4 {
return false
}
dev.uptime = binary.BigEndian.Uint32(data[:4])
}
return true
}
func handleLogoutMsg(dev *Device, data []byte) error {
if len(data) < 32 {
return fmt.Errorf("invalid logout msg from device '%s'", dev.id)
}
sid := string(data[:32])
if val, loaded := dev.users.LoadAndDelete(sid); loaded {
user := val.(*User)
user.Close()
}
return nil
}
func handleLoginMsg(dev *Device, data []byte) error {
if len(data) < 33 {
return fmt.Errorf("invalid login msg from device '%s'", dev.id)
}
sid := string(data[:32])
code := data[32]
if val, loaded := dev.pending.LoadAndDelete(sid); loaded {
user := val.(*User)
ok := code == 0
errCode := 0
if ok {
log.Debug().Msgf("login session '%s' for device '%s' success", sid, dev.id)
dev.users.Store(sid, user)
} else {
errCode = LoginErrorBusy
log.Error().Msgf("login session '%s' for device '%s' fail, due to device busy", sid, dev.id)
}
if errCode == 0 {
user.WriteMsg(websocket.TextMessage, []byte(fmt.Appendf(nil, `{"type":"login"}`)))
} else {
user.SendCloseMsg(LoginErrorBusy, "device busy")
}
user.pending <- ok
}
return nil
}
func handleTermDataMsg(dev *Device, data []byte) error {
if len(data) < 32 {
return fmt.Errorf("invalid term data msg from device '%s'", dev.id)
}
sid := string(data[:32])
if val, ok := dev.users.Load(sid); ok {
user := val.(*User)
data[31] = 0
user.WriteMsg(websocket.BinaryMessage, data[31:])
}
return nil
}
func handleFileMsg(dev *Device, data []byte) error {
if len(data) < 33 {
return fmt.Errorf("invalid file msg from device '%s'", dev.id)
}
sid := string(data[:32])
typ := data[32]
if val, ok := dev.users.Load(sid); ok {
user := val.(*User)
switch typ {
case msgTypeFileSend:
user.WriteMsg(websocket.TextMessage,
fmt.Appendf(nil, `{"type":"sendfile", "name": "%s"}`, string(data[33:])))
case msgTypeFileRecv:
user.WriteMsg(websocket.TextMessage, []byte(`{"type":"recvfile"}`))
case msgTypeFileData:
data[32] = 1
user.WriteMsg(websocket.BinaryMessage, data[32:])
case msgTypeFileAck:
user.WriteMsg(websocket.TextMessage, []byte(`{"type":"fileAck"}`))
case msgTypeFileAbort:
user.WriteMsg(websocket.BinaryMessage, []byte{1})
}
}
return nil
}
func handleHttpMsg(dev *Device, data []byte) error {
if len(data) < 18 {
return fmt.Errorf("invalid http msg from device '%s'", dev.id)
}
addr := data[:18]
data = data[18:]
if c, ok := dev.https.Load(string(addr)); ok {
c := c.(net.Conn)
if len(data) == 0 {
c.Close()
} else {
c.Write(data)
}
}
return nil
}
func handleCmdMsg(dev *Device, data []byte) error {
info := &CommandRespInfo{}
err := jsoniter.Unmarshal(data, info)
if err != nil {
return fmt.Errorf("parse command resp info error: %v", err)
}
var attrs map[string]any
err = jsoniter.Unmarshal(info.Attrs, &attrs)
if err != nil {
return fmt.Errorf("parse command resp attrs error: %v", err)
}
attrs["devid"] = dev.id
if val, ok := dev.commands.Load(info.Token); ok {
req := val.(*CommandReq)
req.acked = true
req.c.JSON(http.StatusOK, attrs)
req.cancel()
}
return nil
}
+735
View File
@@ -0,0 +1,735 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"bufio"
"context"
"crypto/tls"
"encoding/binary"
"errors"
"fmt"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"rttys/internal/proxy"
"rttys/utils"
"github.com/gin-gonic/gin"
"github.com/rs/zerolog/log"
"github.com/valyala/bytebufferpool"
)
type HttpProxySession struct {
expire atomic.Int64
ctx context.Context
cancel context.CancelFunc
devid string
group string
destaddr string
https bool
}
var httpProxySessions = sync.Map{}
const httpProxySessionsExpire = 15 * time.Minute
func (ses *HttpProxySession) Expire() {
ses.expire.Store(time.Now().Add(httpProxySessionsExpire).Unix())
}
func (ses *HttpProxySession) String() string {
return fmt.Sprintf("{devid: %s, group: %s, destaddr: %s, https: %v}",
ses.devid, ses.group, ses.destaddr, ses.https)
}
func (srv *RttyServer) ListenHttpProxy() {
cfg := &srv.cfg
if cfg.AddrHttpProxy != "" {
addr, err := net.ResolveTCPAddr("tcp", cfg.AddrHttpProxy)
if err != nil {
log.Warn().Msg("invalid http proxy addr: " + err.Error())
} else {
srv.httpProxyPort = addr.Port
}
}
ln, err := net.Listen("tcp", cfg.AddrHttpProxy)
if err != nil {
log.Fatal().Msg(err.Error())
}
defer ln.Close()
// In reverse proxy mode (TLS terminated by nginx), never enable TLS here.
enableTLS := !cfg.ReverseProxyEnabled && cfg.SslCert != "" && cfg.SslKey != ""
if enableTLS {
crt, err := tls.LoadX509KeyPair(cfg.SslCert, cfg.SslKey)
if err != nil {
log.Fatal().Msg(err.Error())
}
tlsConfig := &tls.Config{Certificates: []tls.Certificate{crt}}
ln = tls.NewListener(ln, tlsConfig)
}
srv.httpProxyPort = ln.Addr().(*net.TCPAddr).Port
log.Info().Msgf("Listen http proxy on: %s", ln.Addr().(*net.TCPAddr))
go httpProxySessionsClean()
for {
c, err := ln.Accept()
if err != nil {
log.Error().Msg(err.Error())
continue
}
go doHttpProxy(srv, c)
}
}
func httpProxySessionsClean() {
for {
time.Sleep(time.Second * 30)
httpProxySessions.Range(func(key, value any) bool {
ses := value.(*HttpProxySession)
if time.Now().Unix() > ses.expire.Load() {
log.Debug().Msgf("Http proxy session '%s' expired", key)
ses.cancel()
httpProxySessions.Delete(key)
}
return true
})
}
}
func doHttpProxy(srv *RttyServer, c net.Conn) {
defer LogPanic()
defer c.Close()
br := bufio.NewReader(c)
req, err := http.ReadRequest(br)
if err != nil {
return
}
domain, port, proto := proxy.GetRequestHostInfo(req)
log.Debug().Msgf("http proxy incoming host=%s port=%s proto=%s uri=%s",
domain, port, proto, req.URL.String())
devID, ok := proxy.ExtractDeviceIDFromHost(domain)
if ok {
log.Debug().Msgf("parsed deviceId from host: %s", devID)
} else {
log.Debug().Msgf("host is IP or invalid, skip deviceId parsing")
}
queryParams := req.URL.Query()
name := queryParams.Get("rttysid")
if name != "" {
location := "/"
Write302WithCookie(c, location, "rtty-http-sid", name)
return
}
cookie, err := req.Cookie("rtty-http-sid")
if err != nil {
log.Debug().Msgf(`not found cookie "rtty-http-sid"`)
sendHTTPErrorResponse(c, "invalid")
return
}
sid := cookie.Value
sesVal, ok := httpProxySessions.Load(sid)
if !ok {
log.Debug().Msgf(`not found httpProxySession "%s"`, sid)
sendHTTPErrorResponse(c, "unauthorized")
return
}
ses := sesVal.(*HttpProxySession)
dev := srv.GetDevice(ses.group, ses.devid)
if dev == nil {
log.Debug().Msgf(`device "%s" group "%s" offline`, ses.devid, ses.group)
sendHTTPErrorResponse(c, "offline")
return
}
// 3) match hostDevID vs session devid, and optionally lookup by hostDevID
if devID != "" {
match := devID == ses.devid
log.Debug().Msgf(
"http proxy devid check: hostDevID=%s sessionDevid=%s match=%v hostDevFound=%v sid=%s group=%s",
devID, ses.devid, match, domain, sid, ses.group,
)
// If you want, you can also log when mismatch happens
if !match {
log.Info().Msgf(
"http proxy devid mismatch: hostDevID=%s sessionDevid=%s sid=%s group=%s host=%s uri=%s",
devID, ses.devid, sid, ses.group, domain, req.URL.String(),
)
sendHTTPErrorResponse(c, "invalid")
}
} else {
log.Debug().Msgf(
"http proxy devid check skipped: no hostDevID (host=%s) sid=%s group=%s sessionDevid=%s",
domain, sid, ses.group, ses.devid,
)
}
hostHeaderRewrite := ses.destaddr
destAddr := genDestAddr(hostHeaderRewrite)
srcAddr := tcpAddr2Bytes(c.RemoteAddr().(*net.TCPAddr))
ctx, cancel := context.WithCancel(ses.ctx)
defer cancel()
go func() {
<-ctx.Done()
c.Close()
log.Debug().Msgf("http proxy conn closed: %s", ses)
dev.https.Delete(string(srcAddr))
sendHttpReq(dev, ses.https, srcAddr[:], destAddr, nil)
}()
log.Debug().Msgf("new http proxy conn: %s", ses)
dev.https.Store(string(srcAddr), c)
hpw := &HttpProxyWriter{destAddr, srcAddr, hostHeaderRewrite, dev, ses.https}
req.Host = hostHeaderRewrite
hpw.WriteRequest(req)
if req.Header.Get("Upgrade") == "websocket" {
b := make([]byte, 4096)
for {
n, err := c.Read(b)
if err != nil {
return
}
sendHttpReq(dev, ses.https, srcAddr, destAddr, b[:n])
ses.Expire()
}
} else {
for {
req, err := http.ReadRequest(br)
if err != nil {
return
}
hpw.WriteRequest(req)
ses.Expire()
}
}
}
func httpProxyRedirect(srv *RttyServer, c *gin.Context, group string) {
cfg := &srv.cfg
devid := c.Param("devid")
proto := c.Param("proto")
addr := c.Param("addr")
rawPath := c.Param("path")
log.Info().Msgf("httpProxyRedirect devid: %s, proto: %s, addr: %s, path: %s", devid, proto, addr, rawPath)
if !callUserHookUrl(cfg, c) {
c.Status(http.StatusForbidden)
return
}
log.Debug().Msgf("httpProxyRedirect devid: %s, proto: %s, addr: %s, path: %s", devid, proto, addr, rawPath)
_, _, err := httpProxyVaildAddr(addr)
if err != nil {
log.Debug().Msgf("invalid addr: %s", addr)
c.Status(http.StatusBadRequest)
return
}
path, err := url.Parse(rawPath)
if err != nil {
log.Debug().Msgf("invalid path: %s", rawPath)
c.Status(http.StatusBadRequest)
return
}
dev := srv.GetDevice(group, devid)
if dev == nil {
c.Redirect(http.StatusFound, "/error/offline")
return
}
location := c.Request.Header.Get("HttpProxyRedir")
log.Info().Msgf("HttpProxyRedir location: %s, devid: %s", location, devid)
if location == "" {
location = cfg.HttpProxyRedirURL
if location != "" {
log.Debug().Msgf("use HttpProxyRedirURL from config: %s, devid: %s", location, devid)
}
} else {
log.Debug().Msgf("use HttpProxyRedir from HTTP header: %s, devid: %s", location, devid)
}
if location == "" {
host, _, err := net.SplitHostPort(c.Request.Host)
if err != nil {
host = c.Request.Host
}
location = "http://" + host
if srv.httpProxyPort != 80 {
location += fmt.Sprintf(":%d", srv.httpProxyPort)
}
}
location += path.Path
if path.RawQuery != "" {
location += "&" + path.RawQuery
}
sid, err := c.Cookie("rtty-http-sid")
log.Info().Msgf("rtty-http-sid: %s", sid)
if err == nil {
if v, loaded := httpProxySessions.LoadAndDelete(sid); loaded {
s := v.(*HttpProxySession)
s.cancel()
log.Debug().Msgf(`del old httpProxySession "%s" for device "%s"`, sid, devid)
}
}
sid = utils.GenUniqueID()
log.Info().Msgf("rtty-http-sid: %s", sid)
ctx, cancel := context.WithCancel(dev.ctx)
ses := &HttpProxySession{
ctx: ctx,
cancel: cancel,
devid: devid,
group: group,
destaddr: addr,
https: proto == "https",
}
ses.Expire()
httpProxySessions.Store(sid, ses)
log.Debug().Msgf(`new httpProxySession "%s" for device "%s"`, sid, devid)
domain := c.Request.Header.Get("HttpProxyRedirDomain")
if domain == "" {
domain = cfg.HttpProxyRedirDomain
if domain != "" {
log.Debug().Msgf("set cookie domain from config: %s, devid: %s", domain, devid)
}
} else {
log.Debug().Msgf("set cookie domain from HTTP header: %s, devid: %s", domain, devid)
}
// Get domain info
host := c.Request.Host
hostname, _, err := net.SplitHostPort(host)
if err != nil {
hostname = host
}
log.Info().Msgf("hostname: %s", hostname)
ip := net.ParseIP(hostname)
isIP := ip != nil
if isIP {
location = fmt.Sprintf("https://%s%s?rttysid=%s", hostname, cfg.AddrHttpProxy, sid)
log.Info().Msgf("Using IP redirect: %s", location)
} else {
redirHost := proxy.BuildRedirectHost(hostname, devid)
// Keep original behavior when NOT in reverse proxy mode
if !cfg.ReverseProxyEnabled {
location = fmt.Sprintf("https://%s%s?rttysid=%s", redirHost, cfg.AddrHttpProxy, sid)
log.Info().Msgf("Using domain redirect: %s", location)
} else {
// ---- verify forwarded headers from reverse proxy ----
rawHost := c.GetHeader("Host")
xfHost := c.GetHeader("X-Forwarded-Host")
xfProto := c.GetHeader("X-Forwarded-Proto")
xfPort := c.GetHeader("X-Forwarded-Port")
xRealIP := c.GetHeader("X-Real-IP")
xFF := c.GetHeader("X-Forwarded-For")
log.Info().Msgf(
"reverse-proxy info: method=%s uri=%s host=%q tls=%v remoteIP=%q",
c.Request.Method,
c.Request.URL.String(),
rawHost,
c.Request.TLS != nil,
c.ClientIP(),
)
log.Info().Msgf(
"reverse-proxy headers: Host=%q X-Forwarded-Host=%q X-Forwarded-Proto=%q X-Forwarded-Port=%q X-Real-IP=%q X-Forwarded-For=%q",
rawHost, xfHost, xfProto, xfPort, xRealIP, xFF,
)
// -------------------------------------------------
// Proxy mode:
// 1) If DEVICE_ENDPOINT_HOST is configured, use it directly
// 2) Otherwise, fallback to forwarded-header logic
// -------------------------------------------------
// 0) scheme: follow reverse proxy
scheme := ""
if v := strings.TrimSpace(c.GetHeader("X-Forwarded-Proto")); v != "" {
scheme = strings.ToLower(strings.Split(v, ",")[0])
} else if c.Request.TLS != nil {
scheme = "https"
} else {
scheme = "http"
}
// [A] Prefer explicit DEVICE_ENDPOINT_HOST if set
if v := strings.TrimSpace(cfg.DeviceEndpointHost); v != "" {
endpoint := v // already normalized when reading env: host[:port] only
baseHost := endpoint
port := ""
if h, p, err := net.SplitHostPort(endpoint); err == nil {
baseHost = h
port = p
}
// Build device host: <deviceId>.<baseHost>
// NOTE: DEVICE_ENDPOINT_HOST is a base domain (host[:port]) for device access,
baseHost = strings.TrimSuffix(strings.TrimSpace(baseHost), ".")
deviceHost := devid
if baseHost != "" {
deviceHost = devid + "." + baseHost
}
hostPort := proxy.JoinHostPortIfNeeded(deviceHost, scheme, port)
redirectPath := c.Request.URL.Path
location = proxy.BuildRedirectLocation(scheme, hostPort, redirectPath, sid)
log.Info().Msgf("Using domain redirect (proxy mode, DEVICE_ENDPOINT_HOST): %s", location)
} else {
// 1) external port: prefer the one user actually accessed
port := ""
if fp := strings.TrimSpace(c.GetHeader("X-Forwarded-Port")); fp != "" {
port = strings.TrimSpace(strings.Split(fp, ",")[0])
} else if fh := strings.TrimSpace(c.GetHeader("X-Forwarded-Host")); fh != "" {
fh = strings.TrimSpace(strings.Split(fh, ",")[0])
if _, p, err := net.SplitHostPort(fh); err == nil && p != "" {
port = p
}
}
log.Info().Msgf("port: %s", port)
// 3) Build host: in proxy mode redirect domain to be redirHost
hostPort := proxy.JoinHostPortIfNeeded(redirHost, scheme, port)
redirectPath := c.Request.URL.Path
location = proxy.BuildRedirectLocation(scheme, hostPort, redirectPath, sid)
log.Info().Msgf("Using domain redirect (proxy mode): %s", location)
}
}
}
log.Info().Msgf("Final redirect location: %s", location)
c.Redirect(http.StatusFound, location)
}
func sendHttpReq(dev *Device, https bool, srcAddr []byte, destAddr []byte, data []byte) {
bb := bytebufferpool.Get()
defer bytebufferpool.Put(bb)
if dev.proto > 3 {
if https {
bb.WriteByte(1)
} else {
bb.WriteByte(0)
}
}
bb.Write(srcAddr)
bb.Write(destAddr)
bb.Write(data)
dev.WriteMsg(msgTypeHttp, "", bb.Bytes())
}
func genDestAddr(addr string) []byte {
destIP, destPort, err := httpProxyVaildAddr(addr)
if err != nil {
return nil
}
b := make([]byte, 6)
copy(b, destIP)
binary.BigEndian.PutUint16(b[4:], destPort)
return b
}
func tcpAddr2Bytes(addr *net.TCPAddr) []byte {
b := make([]byte, 18)
binary.BigEndian.PutUint16(b[:2], uint16(addr.Port))
copy(b[2:], addr.IP)
return b
}
func httpProxyVaildAddr(addr string) (net.IP, uint16, error) {
ips, ports, err := net.SplitHostPort(addr)
if err != nil {
ips = addr
ports = "80"
}
ip := net.ParseIP(ips)
if ip == nil {
return nil, 0, errors.New("invalid IPv4 Addr")
}
ip = ip.To4()
if ip == nil {
return nil, 0, errors.New("invalid IPv4 Addr")
}
port, _ := strconv.Atoi(ports)
return ip, uint16(port), nil
}
type HttpProxyWriter struct {
destAddr []byte
srcAddr []byte
hostHeaderRewrite string
dev *Device
https bool
}
func (rw *HttpProxyWriter) Write(p []byte) (n int, err error) {
sendHttpReq(rw.dev, rw.https, rw.srcAddr, rw.destAddr, p)
return len(p), nil
}
func (rw *HttpProxyWriter) WriteRequest(req *http.Request) {
req.Host = rw.hostHeaderRewrite
req.Write(rw)
}
func generateErrorHTML(errorType string) string {
return fmt.Sprintf(
`<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>RTTY</title>
<style>
body {
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, "Helvetica Neue", Arial, sans-serif;
background-color: #555;
line-height: 1.6;
}
.error-container {
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
min-height: 60vh;
text-align: center;
}
.error-icon {
margin-bottom: 2rem;
animation: fadeIn 0.8s ease-in-out;
}
.error-icon svg {
width: 90px;
height: 90px;
fill: #f56565;
}
.error-content {
max-width: 700px;
animation: slideUp 0.8s ease-out 0.2s both;
}
.error-title {
font-size: 1.8rem;
font-weight: 600;
color: #7a8fb0;
margin-bottom: 1rem;
line-height: 1.2;
}
.error-message {
font-size: 1rem;
color: #b6c1d3;
margin-bottom: 2rem;
line-height: 1.6;
text-align: left;
}
@keyframes fadeIn {
from {
opacity: 0;
transform: scale(0.8);
}
to {
opacity: 1;
transform: scale(1);
}
}
@keyframes slideUp {
from {
opacity: 0;
transform: translateY(20px);
}
to {
opacity: 1;
transform: translateY(0);
}
}
</style>
</head>
<body>
<div class="error-container">
<div class="error-icon">
<svg viewBox="0 0 24 24">
<path d="M1 21h22L12 2 1 21zm12-3h-2v-2h2v2zm0-4h-2v-4h2v4z"/>
</svg>
</div>
<div class="error-content">
<h2 class="error-title" id="errorTitle"></h2>
<p class="error-message" id="errorMessage"></p>
</div>
</div>
<script>
const translations = {
en: {
'Device Unavailable': 'Device Unavailable',
'Invalid Request': 'Invalid Request',
'Unauthorized Access': 'Unauthorized Access',
'Device offline message': 'The device is currently offline. Please check the device status and try again.',
'Invalid request message': 'The request is invalid or malformed',
'Unauthorized request message': 'You are not authorized to access this resource. Please check your session and try again.'
},
'zh-CN': {
'Device Unavailable': '璁惧涓嶅彲鐢?,
'Invalid Request': '鏃犳晥璇锋眰',
'Unauthorized Access': '鏈巿鏉冭闂?,
'Device offline message': '璁惧褰撳墠绂荤嚎锛岃妫€鏌ヨ澶囩姸鎬佸悗閲嶈瘯銆?,
'Invalid request message': '璇锋眰鏃犳晥鎴栨牸寮忛敊璇?,
'Unauthorized request message': '鎮ㄦ棤鏉冭闂璧勬簮銆傝妫€鏌ユ偍鐨勪細璇濆苟閲嶈瘯銆?
}
};
function t(key, lang) {
return translations[lang][key] || translations.en[key] || key;
}
function updateContent() {
const errorType = '%s';
const lang = navigator.language === 'zh-CN' ? 'zh-CN' : 'en';
let title = '', message = '';
switch (errorType) {
case 'offline':
title = t('Device Unavailable', lang);
message = t('Device offline message', lang);
break;
case 'invalid':
title = t('Invalid Request', lang);
message = t('Invalid request message', lang);
break;
case 'unauthorized':
title = t('Unauthorized Access', lang);
message = t('Unauthorized request message', lang);
break;
}
document.getElementById('errorTitle').textContent = title;
document.getElementById('errorMessage').textContent = message;
// Update page title
if (title) {
document.title = title + ' - RTTY';
} else {
document.title = 'Error - RTTY';
}
}
// Initialize page on load
document.addEventListener('DOMContentLoaded', updateContent);
</script>
</body>
</html>`, errorType)
}
func sendHTTPErrorResponse(conn net.Conn, errorType string) {
htmlContent := generateErrorHTML(errorType)
response := "HTTP/1.1 200 OK\r\n"
response += "Content-Type: text/html; charset=utf-8\r\n"
response += fmt.Sprintf("Content-Length: %d\r\n", len(htmlContent))
response += "Connection: close\r\n"
response += "\r\n"
response += htmlContent
conn.Write([]byte(response))
}
func Write302WithCookie(conn net.Conn, location, cookieName, cookieValue string) {
cookie := fmt.Sprintf("%s=%s; Path=/; HttpOnly", cookieName, cookieValue)
response := fmt.Sprintf(
"HTTP/1.1 302 Found\r\n"+
"Location: %s\r\n"+
"Set-Cookie: %s\r\n"+
"Content-Length: 0\r\n"+
"Connection: close\r\n"+
"\r\n",
location, cookie,
)
_, _ = conn.Write([]byte(response))
}
+41
View File
@@ -0,0 +1,41 @@
package server
import (
"net/http"
"rttys/internal/proxy"
)
type HostInfo = proxy.HostInfo
func getHostInfoFromRequest(req *http.Request) HostInfo {
return proxy.GetHostInfoFromRequest(req)
}
func isIPHost(host string) bool {
return proxy.IsIPHost(host)
}
func domainAllowed(host, base string) bool {
return proxy.DomainAllowed(host, base)
}
func buildRedirectHost(hostname, devid string) string {
return proxy.BuildRedirectHost(hostname, devid)
}
func joinHostPortIfNeeded(host, scheme, port string) string {
return proxy.JoinHostPortIfNeeded(host, scheme, port)
}
func buildRedirectLocation(scheme, hostPort, path, sid string) string {
return proxy.BuildRedirectLocation(scheme, hostPort, path, sid)
}
func getRequestHostInfo(req *http.Request) (host string, port string, proto string) {
return proxy.GetRequestHostInfo(req)
}
func extractDeviceIDFromHost(host string) (string, bool) {
return proxy.ExtractDeviceIDFromHost(host)
}
+27 -18
View File
@@ -1,4 +1,4 @@
package main
package server
import (
"context"
@@ -6,7 +6,6 @@ import (
"encoding/json"
"fmt"
oidc "github.com/coreos/go-oidc/v3/oidc"
"github.com/fanjindong/go-cache"
"github.com/gin-gonic/gin"
"github.com/gorilla/sessions"
"github.com/rs/zerolog/log"
@@ -14,7 +13,9 @@ import (
"math/rand"
"net/http"
"net/url"
"rttys/utils"
"rttys/internal/domain/user"
"rttys/internal/pkg/randtoken"
"rttys/xconfig"
"strings"
"time"
)
@@ -27,7 +28,7 @@ var (
)
// Register OIDC routes
func RegisterOIDCRoutes(r *gin.Engine, cfg *Config) {
func RegisterOIDCRoutes(r *gin.Engine, cfg *xconfig.Config, userSvc *user.Service) {
if !cfg.OIDCEnabled {
return
}
@@ -69,11 +70,11 @@ func RegisterOIDCRoutes(r *gin.Engine, cfg *Config) {
// OIDC auth routes (public, no existing auth required)
r.GET("/auth/oidc/login", oidcLoginHandler(cfg))
r.GET("/auth/oidc/callback", oidcCallbackHandler(cfg))
r.GET("/auth/oidc/callback", oidcCallbackHandler(cfg, userSvc))
}
// Start OIDC login
func oidcLoginHandler(cfg *Config) gin.HandlerFunc {
func oidcLoginHandler(cfg *xconfig.Config) gin.HandlerFunc {
return func(c *gin.Context) {
// Generate state and nonce
@@ -109,7 +110,7 @@ func oidcLoginHandler(cfg *Config) gin.HandlerFunc {
}
// Handle OIDC callback
func oidcCallbackHandler(cfg *Config) gin.HandlerFunc {
func oidcCallbackHandler(cfg *xconfig.Config, userSvc *user.Service) gin.HandlerFunc {
return func(c *gin.Context) {
// Get session
session, err := oauthStore.Get(c.Request, "oidc-session")
@@ -155,7 +156,7 @@ func oidcCallbackHandler(cfg *Config) gin.HandlerFunc {
var userEmail string
var userName string
// Standard OIDC – verify and parse ID token
// Standard OIDC verify and parse ID token
rawIDToken, ok := tokens["id_token"].(string)
if !ok {
log.Error().Msg("No ID token in response")
@@ -231,15 +232,23 @@ func oidcCallbackHandler(cfg *Config) gin.HandlerFunc {
return
}
// Create application session
sid := utils.GenUniqueID()
httpSessions.Set(sid, gin.H{
"email": userEmail,
"name": userName,
"oidc": true,
}, cache.WithEx(httpSessionExpire))
// ==== Create application session (new session_store, same as LDAP) ====
sid, err := randtoken.New() // randtoken.New()
if err != nil {
log.Error().Err(err).Msg("Failed to create session token")
c.Redirect(http.StatusFound, "/?error=internal_error")
return
}
c.SetCookie("sid", sid, 0, "", "", cfg.SslCert != "", true)
sysAdmin, err := userSvc.GetSystemAdmin(c.Request.Context())
if err != nil {
log.Error().Err(err).Msg("Failed to find system admin user")
c.Redirect(http.StatusFound, "/?error=internal_error")
return
}
sessionStore.Create(sid, sysAdmin.ID)
c.SetCookie("sid", sid, 0, "/", "", cfg.SslCert != "", false)
// Clean up OAuth session
session.Options.MaxAge = -1
@@ -251,7 +260,7 @@ func oidcCallbackHandler(cfg *Config) gin.HandlerFunc {
}
// Exchange authorization code for tokens
func exchangeCodeForTokens(cfg *Config, code string) (map[string]interface{}, error) {
func exchangeCodeForTokens(cfg *xconfig.Config, code string) (map[string]interface{}, error) {
data := url.Values{}
data.Set("code", code)
data.Set("client_id", cfg.OIDCGenericClientID)
@@ -297,7 +306,7 @@ func generateRandomString(length int) string {
return base64.URLEncoding.EncodeToString(b)[:length]
}
func isOIDCUserAllowed(cfg *Config, claims map[string]interface{}) bool {
func isOIDCUserAllowed(cfg *xconfig.Config, claims map[string]interface{}) bool {
email, _ := claims["email"].(string)
sub, _ := claims["sub"].(string)
preferredUsername, _ := claims["preferred_username"].(string)
+20
View File
@@ -0,0 +1,20 @@
package server
import (
"os"
"runtime/debug"
"github.com/rs/zerolog/log"
)
func LogPanic() {
if r := recover(); r != nil {
SaveCrashLog(r, debug.Stack())
os.Exit(2)
}
}
func SaveCrashLog(p any, stack []byte) {
log.Error().Msgf("%v", p)
log.Error().Msg(string(stack))
}
+202
View File
@@ -0,0 +1,202 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"context"
"flag"
"io"
"net/http"
"net/http/cookiejar"
"net/url"
"strings"
"sync"
"testing"
"time"
"rttys/xconfig"
"github.com/gorilla/websocket"
"github.com/rs/zerolog/log"
)
func TestRttysStress(t *testing.T) {
duration := 10 * time.Minute
timeoutFlag := flag.Lookup("test.timeout")
if timeoutFlag != nil {
duration = timeoutFlag.Value.(flag.Getter).Get().(time.Duration)
}
cfg := xconfig.Config{
AddrDev: ":5912",
AddrUser: ":5913",
}
srv := &RttyServer{cfg: cfg}
go func() {
err := srv.Run()
if err != nil {
log.Fatal().Msg(err.Error())
}
}()
ctx, cancel := context.WithTimeout(context.Background(), duration-time.Second*2)
defer cancel()
time.Sleep(time.Millisecond * 100)
log.Info().Msg("Waiting for devices to connect for testing...")
devices := &sync.Map{}
for {
select {
case <-ctx.Done():
log.Info().Msg("Test timeout, exiting...")
return
default:
time.Sleep(time.Second * 1)
srv.groups.Range(func(key, value any) bool {
group := key.(string)
g := value.(*DeviceGroup)
g.devices.Range(func(key, value any) bool {
dev := value.(*Device)
if _, loaded := devices.LoadOrStore(dev.id, group+dev.id); !loaded {
go runDeviceTest(ctx, devices, group, dev.id)
}
return true
})
return true
})
}
}
}
func runDeviceTest(ctx context.Context, devices *sync.Map, group, devID string) {
ctx, cancel := context.WithCancel(ctx)
defer func() {
time.Sleep(time.Second)
cancel()
devices.Delete(group + devID)
}()
go runHttpTest(ctx, group, devID)
wg := &sync.WaitGroup{}
for range 7 {
wg.Add(1)
go runWebSocketTest(ctx, group, devID, wg)
}
wg.Wait()
}
func runWebSocketTest(ctx context.Context, group, devID string, wg *sync.WaitGroup) {
conn, _, err := websocket.DefaultDialer.Dial("ws://127.0.0.1:5913/connect/"+devID+"?group="+group, nil)
if err != nil {
log.Fatal().Msg(err.Error())
}
defer conn.Close()
defer wg.Done()
go func() {
<-ctx.Done()
conn.Close()
}()
go func() {
msg := []byte{0}
msg = append(msg, []byte("ttttttttttttttttttttttttttttt\n")...)
msg = append(msg, []byte("ttttttttttttttttttttttttttttt\n")...)
for {
err = conn.WriteMessage(websocket.BinaryMessage, msg)
if err != nil {
return
}
time.Sleep(time.Millisecond * 20)
}
}()
for {
_, _, err := conn.ReadMessage()
if err != nil {
return
}
}
}
func runHttpTest(ctx context.Context, group, devID string) {
for {
select {
case <-ctx.Done():
return
default:
runHttpTestOnce(ctx, group, devID)
}
}
}
func runHttpTestOnce(ctx context.Context, group, devID string) {
addr := ""
if group == "" {
addr = "http://127.0.0.1:5913/web/"
} else {
addr = "http://127.0.0.1:5913/web2/" + group + "/"
}
addr += devID + "/http/" + encodeURIComponent("127.0.0.1:80/")
jar, _ := cookiejar.New(nil)
client := &http.Client{
Jar: jar,
}
request, _ := http.NewRequestWithContext(ctx, "GET", addr, nil)
for range 10 {
res, err := client.Do(request)
if err != nil {
log.Info().Msg(err.Error())
return
}
defer res.Body.Close()
io.ReadAll(res.Body)
time.Sleep(10 * time.Millisecond)
}
}
func encodeURIComponent(str string) string {
r := url.QueryEscape(str)
r = strings.ReplaceAll(r, "+", "%20")
return r
}
+141 -107
View File
@@ -1,144 +1,178 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"context"
"net"
"net/http"
"rttys/internal/store/sqlite"
"rttys/xconfig"
"strings"
"sync"
"sync/atomic"
"github.com/rs/zerolog/log"
)
type RttyServer struct {
mu sync.RWMutex
groups sync.Map
cfg Config
cfg xconfig.Config
httpProxyPort int
}
type DeviceGroup struct {
devices sync.Map
count atomic.Int32
}
func New(cfg xconfig.Config) *RttyServer {
return &RttyServer{cfg: cfg}
}
func (srv *RttyServer) Run() error {
log.Debug().Msgf("%+v", srv.cfg)
if err := markAllDevicesOffline(); err != nil {
log.Warn().Err(err).Msg("mark all devices offline failed")
}
if srv.cfg.PprofAddr != "" {
go srv.ListenPprof()
}
log.Info().Msgf("SslCert: %s,SslKey: %s", srv.cfg.SslCert, srv.cfg.SslKey)
log.Info().Msgf("SslCert: %s,SslKey: %s", srv.cfg.SslCert, srv.cfg.SslKey)
go srv.ListenDevices()
go srv.ListenHttpProxy()
return srv.ListenAPI()
}
func (srv *RttyServer) ListenPprof() {
ln, err := net.Listen("tcp", srv.cfg.PprofAddr)
func markAllDevicesOffline() error {
db, err := sqlite.Open(context.Background(), sqlite.Options{
DSN: defaultDBPath,
MaxOpenConns: 1,
MaxIdleConns: 1,
LogSQL: false,
})
if err != nil {
log.Error().Err(err).Msgf("Failed to start pprof server")
return
return err
}
defer ln.Close()
defer db.Close()
addr := ln.Addr().(*net.TCPAddr)
log.Info().Msgf("Starting pprof server on: %s", addr)
host := addr.IP.String()
if host == "0.0.0.0" || host == "::" {
host = "localhost"
}
log.Info().Msgf("Access pprof at: http://%s:%d/debug/pprof/", host, addr.Port)
err = http.Serve(ln, nil)
if err != nil {
log.Error().Err(err).Msgf("pprof server failed")
}
}
func (srv *RttyServer) GetDevice(group, id string) *Device {
srv.mu.RLock()
defer srv.mu.RUnlock()
g := srv.GetGroup(group, false)
if g == nil {
return nil
}
if v, ok := g.devices.Load(id); ok {
return v.(*Device)
}
return nil
}
func (srv *RttyServer) AddDevice(dev *Device) bool {
srv.mu.Lock()
defer srv.mu.Unlock()
g := srv.GetGroup(dev.group, true)
if _, loaded := g.devices.LoadOrStore(dev.id, dev); loaded {
return false
}
g.count.Add(1)
return true
}
func (srv *RttyServer) DelDevice(dev *Device) {
srv.mu.Lock()
defer srv.mu.Unlock()
g := srv.GetGroup(dev.group, false)
if g == nil {
return
}
if deleted := g.devices.CompareAndDelete(dev.id, dev); deleted {
if g.count.Add(-1) == 0 {
srv.groups.Delete(dev.group)
}
}
}
func (srv *RttyServer) GetGroup(group string, create bool) *DeviceGroup {
if create {
val, _ := srv.groups.LoadOrStore(group, &DeviceGroup{})
return val.(*DeviceGroup)
} else {
val, ok := srv.groups.Load(group)
if !ok {
res := db.Gorm().Exec(`UPDATE devices SET status='offline' WHERE status='online'`)
if res.Error != nil {
if strings.Contains(res.Error.Error(), "no such table") {
return nil
}
return val.(*DeviceGroup)
return res.Error
}
return nil
}
func (srv *RttyServer) ListenPprof() {
ln, err := net.Listen("tcp", srv.cfg.PprofAddr)
if err != nil {
log.Error().Err(err).Msgf("Failed to start pprof server")
return
}
defer ln.Close()
addr := ln.Addr().(*net.TCPAddr)
log.Info().Msgf("Starting pprof server on: %s", addr)
host := addr.IP.String()
if host == "0.0.0.0" || host == "::" {
host = "localhost"
}
log.Info().Msgf("Access pprof at: http://%s:%d/debug/pprof/", host, addr.Port)
err = http.Serve(ln, nil)
if err != nil {
log.Error().Err(err).Msgf("pprof server failed")
}
}
func (srv *RttyServer) GetDevice(group, id string) *Device {
srv.mu.RLock()
defer srv.mu.RUnlock()
g := srv.GetGroup(group, false)
if g == nil {
return nil
}
if v, ok := g.devices.Load(id); ok {
return v.(*Device)
}
return nil
}
func (srv *RttyServer) AddDevice(dev *Device) bool {
srv.mu.Lock()
defer srv.mu.Unlock()
g := srv.GetGroup(dev.group, true)
if _, loaded := g.devices.LoadOrStore(dev.id, dev); loaded {
return false
}
g.count.Add(1)
return true
}
func (srv *RttyServer) DelDevice(dev *Device) {
srv.mu.Lock()
defer srv.mu.Unlock()
g := srv.GetGroup(dev.group, false)
if g == nil {
return
}
if deleted := g.devices.CompareAndDelete(dev.id, dev); deleted {
if g.count.Add(-1) == 0 {
srv.groups.Delete(dev.group)
}
}
}
func (srv *RttyServer) GetGroup(group string, create bool) *DeviceGroup {
if create {
val, _ := srv.groups.LoadOrStore(group, &DeviceGroup{})
return val.(*DeviceGroup)
} else {
val, ok := srv.groups.Load(group)
if !ok {
return nil
}
return val.(*DeviceGroup)
}
}
+6
View File
@@ -0,0 +1,6 @@
package server
// StartSignalHandler enables runtime signal handling (noop on Windows).
func StartSignalHandler() {
signalHandle()
}
@@ -1,30 +1,30 @@
//go:build !windows
// +build !windows
package main
import (
"os"
"os/signal"
"syscall"
xlog "rttys/log"
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
)
func signalHandle() {
c := make(chan os.Signal, 1)
signal.Notify(c, syscall.SIGUSR1)
for s := range c {
switch s {
case syscall.SIGUSR1:
xlog.Verbose()
zerolog.SetGlobalLevel(zerolog.DebugLevel)
log.Debug().Msg("Debug mode enabled")
}
}
}
//go:build !windows
// +build !windows
package server
import (
"os"
"os/signal"
"syscall"
xlog "rttys/log"
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
)
func signalHandle() {
c := make(chan os.Signal, 1)
signal.Notify(c, syscall.SIGUSR1)
for s := range c {
switch s {
case syscall.SIGUSR1:
xlog.Verbose()
zerolog.SetGlobalLevel(zerolog.DebugLevel)
log.Debug().Msg("Debug mode enabled")
}
}
}
@@ -1,12 +1,12 @@
//go:build windows
// +build windows
package main
import (
"github.com/rs/zerolog/log"
)
func signalHandle() {
log.Debug().Msg("Signal handling not supported on Windows")
}
//go:build windows
// +build windows
package server
import (
"github.com/rs/zerolog/log"
)
func signalHandle() {
log.Debug().Msg("Signal handling not supported on Windows")
}
+238
View File
@@ -0,0 +1,238 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package server
import (
"context"
"encoding/binary"
"net/http"
"sync"
"sync/atomic"
"time"
"rttys/utils"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
jsoniter "github.com/json-iterator/go"
"github.com/rs/zerolog/log"
)
type User struct {
conn *websocket.Conn
sid string
dev *Device
pending chan bool
close sync.Once
closed atomic.Bool
}
type UserMsg struct {
Type string `json:"type"`
Cols uint16 `json:"cols"`
Rows uint16 `json:"rows"`
Ack uint16 `json:"ack"`
Size uint32 `json:"size"`
Name string `json:"name"`
}
const (
LoginErrorOffline = 4000
LoginErrorBusy = 4001
LoginErrorTimeout = 4002
)
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
}
func handleUserConnection(srv *RttyServer, c *gin.Context) {
defer LogPanic()
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
log.Error().Err(err).Msg("upgrade to websocket failed")
return
}
devid := c.Param("devid")
if devid == "" {
log.Error().Msg("device ID is required")
conn.Close()
return
}
user := &User{conn: conn}
dev := srv.GetDevice(c.Query("group"), devid)
if dev == nil {
user.SendCloseMsg(LoginErrorOffline, "device not found")
conn.Close()
return
}
sid := utils.GenUniqueID()
user.sid = sid
user.dev = dev
user.pending = make(chan bool, 1)
dev.pending.Store(sid, user)
defer user.Close()
if err := dev.WriteMsg(msgTypeLogin, sid, nil); err != nil {
log.Error().Msgf("send login msg to device %s fail: %v", dev.id, err)
return
}
ctx, cancel := context.WithCancel(dev.ctx)
go func() {
<-ctx.Done()
user.Close()
}()
defer cancel()
if !waitForLogin(user, dev, ctx, sid) {
return
}
for {
msgType, data, err := conn.ReadMessage()
if err != nil {
if !user.closed.Load() {
closeError, ok := err.(*websocket.CloseError)
if !ok || (closeError.Code != websocket.CloseGoingAway &&
closeError.Code != websocket.CloseAbnormalClosure &&
closeError.Code != websocket.CloseNormalClosure) {
log.Error().Msgf("user read fail: %v", err)
}
}
return
}
if msgType == websocket.BinaryMessage {
if len(data) < 1 {
log.Error().Msgf("invalid msg from user")
return
}
typ := msgTypeTermData
if data[0] == 1 {
typ = msgTypeFile
}
err = dev.WriteMsg(typ, sid, data[1:])
} else {
msg := &UserMsg{}
err = jsoniter.Unmarshal(data, msg)
if err != nil {
log.Error().Msgf("invalid msg from user")
return
}
switch msg.Type {
case "winsize":
b := make([]byte, 4)
binary.BigEndian.PutUint16(b, msg.Cols)
binary.BigEndian.PutUint16(b[2:], msg.Rows)
err = dev.WriteMsg(msgTypeWinsize, sid, b)
case "ack":
b := make([]byte, 2)
binary.BigEndian.PutUint16(b, msg.Ack)
err = dev.WriteMsg(msgTypeAck, sid, b)
case "fileInfo":
b := make([]byte, 4+len(msg.Name))
binary.BigEndian.PutUint32(b, msg.Size)
copy(b[4:], []byte(msg.Name))
err = dev.WriteFileMsg(msgTypeFile, sid, msgTypeFileInfo, b)
case "fileCanceled":
err = dev.WriteFileMsg(msgTypeFile, sid, msgTypeFileAbort, nil)
case "fileAck":
err = dev.WriteFileMsg(msgTypeFile, sid, msgTypeFileAck, nil)
}
}
if err != nil {
log.Error().Msgf("write msg to device '%s' fail: %v", dev.id, err)
return
}
}
}
func (user *User) SendCloseMsg(code int, text string) {
user.conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(code, text), time.Now().Add(time.Second))
}
func (user *User) Close() {
user.close.Do(func() {
dev := user.dev
sid := user.sid
user.closed.Store(true)
if _, loaded := dev.users.LoadAndDelete(sid); loaded {
dev.WriteMsg(msgTypeLogout, sid, nil)
}
dev.pending.Delete(sid)
user.conn.Close()
log.Debug().Msgf("user with session '%s' closed", sid)
})
}
func (user *User) WriteMsg(typ int, data []byte) error {
return user.conn.WriteMessage(typ, data)
}
func waitForLogin(user *User, dev *Device, ctx context.Context, sid string) bool {
for {
select {
case <-ctx.Done():
return false
case ok := <-user.pending:
return ok
case <-time.After(TermLoginTimeout):
if _, loaded := dev.pending.LoadAndDelete(sid); loaded {
log.Error().Msgf("login timeout for session %s of device %s", sid, dev.id)
user.SendCloseMsg(LoginErrorTimeout, "login timeout")
return false
}
}
}
}
+9
View File
@@ -0,0 +1,9 @@
package server
const RttysVersion = "5.2.0"
const KVMCloudVersion = "v2.2.0"
var (
GitCommit = ""
BuildTime = ""
)
+42
View File
@@ -0,0 +1,42 @@
package memory
import (
"context"
"rttys/internal/domain/identity"
"rttys/internal/domain/permission"
)
type PermissionRepo struct {
roleToKeys map[identity.Role][]permission.Key
}
func NewPermissionRepo() *PermissionRepo {
// Role defaults are aligned with the design docs:
// - admin: all enabled permissions
// - user : read (and basic auth/me)
return &PermissionRepo{
roleToKeys: map[identity.Role][]permission.Key{
identity.RoleAdmin: {
permission.MeRead, permission.AuthWrite,
permission.DeviceRead, permission.DeviceWrite,
permission.DeviceGroupRead, permission.DeviceGroupWrite,
permission.UserGroupRead, permission.UserGroupWrite,
permission.UserRead, permission.UserWrite,
permission.RelationWrite,
},
identity.RoleUser: {
permission.MeRead, permission.AuthWrite,
permission.DeviceRead,
permission.DeviceGroupRead, permission.UserGroupRead,
},
},
}
}
func (r *PermissionRepo) ListKeysByRole(ctx context.Context, role identity.Role) ([]permission.Key, error) {
keys := r.roleToKeys[role]
out := make([]permission.Key, 0, len(keys))
out = append(out, keys...)
return out, nil
}
+80
View File
@@ -0,0 +1,80 @@
package memory
import (
"sync"
"time"
)
type Session struct {
Token string
UserID int64
ExpiresAt time.Time
}
type SessionStore struct {
mu sync.RWMutex
ttl time.Duration
data map[string]Session
}
func NewSessionStore(ttl time.Duration) *SessionStore {
s := &SessionStore{
ttl: ttl,
data: map[string]Session{},
}
go s.gcLoop()
return s
}
func (s *SessionStore) Create(token string, userID int64) Session {
sess := Session{Token: token, UserID: userID, ExpiresAt: time.Now().Add(s.ttl)}
s.mu.Lock()
s.data[token] = sess
s.mu.Unlock()
return sess
}
func (s *SessionStore) Get(token string) (Session, bool) {
s.mu.RLock()
sess, ok := s.data[token]
s.mu.RUnlock()
if !ok {
return Session{}, false
}
if time.Now().After(sess.ExpiresAt) {
s.Delete(token)
return Session{}, false
}
return sess, true
}
func (s *SessionStore) Delete(token string) {
s.mu.Lock()
delete(s.data, token)
s.mu.Unlock()
}
func (s *SessionStore) DeleteByUserID(userID int64) {
s.mu.Lock()
for k, v := range s.data {
if v.UserID == userID {
delete(s.data, k)
}
}
s.mu.Unlock()
}
func (s *SessionStore) gcLoop() {
t := time.NewTicker(2 * time.Minute)
defer t.Stop()
for range t.C {
now := time.Now()
s.mu.Lock()
for k, v := range s.data {
if now.After(v.ExpiresAt) {
delete(s.data, k)
}
}
s.mu.Unlock()
}
}
+31
View File
@@ -0,0 +1,31 @@
package sqlite
import (
"gorm.io/gorm"
"sync"
)
type Container struct {
Gorm *gorm.DB
DeviceMeta *DeviceMetaRepo
}
var (
gContainer *Container
mu sync.RWMutex
)
func SetContainer(c *Container) {
mu.Lock()
defer mu.Unlock()
gContainer = c
}
func MustContainer() *Container {
mu.RLock()
defer mu.RUnlock()
if gContainer == nil {
panic("sqlite container not initialized")
}
return gContainer
}
+88
View File
@@ -0,0 +1,88 @@
package sqlite
import (
"context"
"gorm.io/gorm"
"rttys/internal/domain/device"
)
type DeviceRepo struct{ db *gorm.DB }
func NewDeviceRepo(db *gorm.DB) *DeviceRepo { return &DeviceRepo{db: db} }
// 用于 DB 行映射
type deviceRow struct {
ID int64 `gorm:"column:id"`
Ddns string `gorm:"column:ddns"`
Mac string `gorm:"column:mac"`
Name string `gorm:"column:name"`
Description string `gorm:"column:description"`
IP string `gorm:"column:ip"`
Client string `gorm:"column:client"`
DeviceGroupID *int64 `gorm:"column:device_group_id"` // NULL => nil
Status string `gorm:"column:status"`
LastSeenAt *int64 `gorm:"column:last_seen_at"` // NULL => nil
}
func (deviceRow) TableName() string { return "devices" }
func (r *DeviceRepo) ListAll(ctx context.Context) ([]device.Device, error) {
var rows []deviceRow
err := r.db.WithContext(ctx).
Order("id").
Find(&rows).Error
if err != nil {
return nil, err
}
out := make([]device.Device, 0, len(rows))
for _, row := range rows {
out = append(out, device.Device{
ID: row.ID,
Ddns: row.Ddns,
Mac: row.Mac,
Name: row.Name,
Description: row.Description,
IP: row.IP,
Client: row.Client,
DeviceGroupID: row.DeviceGroupID,
Status: device.Status(row.Status),
LastSeenAt: row.LastSeenAt,
})
}
return out, nil
}
func (r *DeviceRepo) ListByDeviceGroupIDs(ctx context.Context, groupIDs []int64) ([]device.Device, error) {
if len(groupIDs) == 0 {
return []device.Device{}, nil
}
var rows []deviceRow
err := r.db.WithContext(ctx).
Where("device_group_id IN ?", groupIDs).
Order("id").
Find(&rows).Error
if err != nil {
return nil, err
}
out := make([]device.Device, 0, len(rows))
for _, row := range rows {
out = append(out, device.Device{
ID: row.ID,
Ddns: row.Ddns,
Mac: row.Mac,
Name: row.Name,
Description: row.Description,
IP: row.IP,
Client: row.Client,
DeviceGroupID: row.DeviceGroupID,
Status: device.Status(row.Status),
LastSeenAt: row.LastSeenAt,
})
}
return out, nil
}
+127
View File
@@ -0,0 +1,127 @@
package sqlite
import (
"context"
"errors"
"fmt"
"rttys/model"
"gorm.io/gorm"
"rttys/utils"
)
type DeviceMetaRepo struct {
db *gorm.DB
}
func NewDeviceMetaRepo(db *gorm.DB) *DeviceMetaRepo {
return &DeviceMetaRepo{db: db}
}
func (r *DeviceMetaRepo) SaveOrUpdate(ctx context.Context, deviceID, mac, description, ip string) error {
if r.db == nil {
return fmt.Errorf("gorm db is nil")
}
normMac := utils.NormalizeMac(mac)
return r.db.WithContext(ctx).Exec(
`INSERT INTO devices (ddns, mac, description, ip, status, last_seen_at)
VALUES (?, ?, ?, ?, 'online', unixepoch())
ON CONFLICT(ddns) DO UPDATE SET
mac=excluded.mac,
description=excluded.description,
ip=excluded.ip,
status='online',
last_seen_at=unixepoch()`,
deviceID, normMac, description, ip,
).Error
}
func (r *DeviceMetaRepo) UpdateClient(ctx context.Context, deviceID, client string) error {
if r.db == nil {
return fmt.Errorf("gorm db is nil")
}
return r.db.WithContext(ctx).Exec(
`UPDATE devices SET client=? WHERE ddns=?`,
client, deviceID,
).Error
}
func (r *DeviceMetaRepo) UpdateDescriptionIfEmpty(ctx context.Context, deviceID, description string) error {
if r.db == nil {
return fmt.Errorf("gorm db is nil")
}
return r.db.WithContext(ctx).Exec(
`UPDATE devices SET description=? WHERE ddns=? AND (description IS NULL OR description='')`,
description, deviceID,
).Error
}
func (r *DeviceMetaRepo) GetByDeviceID(ctx context.Context, deviceID string) (*model.DeviceMeta, error) {
var meta model.DeviceMeta
err := r.db.WithContext(ctx).Where("ddns = ?", deviceID).First(&meta).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
return &meta, err
}
func (r *DeviceMetaRepo) GetByMac(ctx context.Context, mac string) (*model.DeviceMeta, error) {
normMac := utils.NormalizeMac(mac)
var meta model.DeviceMeta
err := r.db.WithContext(ctx).Where("mac = ?", normMac).First(&meta).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
return &meta, err
}
func (r *DeviceMetaRepo) List(ctx context.Context, keyword string) ([]model.DeviceMeta, error) {
var list []model.DeviceMeta
q := r.db.WithContext(ctx).Model(&model.DeviceMeta{})
if keyword != "" {
normMac := utils.NormalizeMac(keyword)
likeDesc := "%" + keyword + "%"
q = q.Where("ddns = ? OR mac = ? OR description LIKE ? OR ip = ?", keyword, normMac, likeDesc, keyword)
}
if err := q.Order("id ASC").Find(&list).Error; err != nil {
return nil, err
}
return list, nil
}
func (r *DeviceMetaRepo) ListByDeviceIDs(ctx context.Context, deviceIDs []string) ([]model.DeviceMeta, error) {
if len(deviceIDs) == 0 {
return []model.DeviceMeta{}, nil
}
var list []model.DeviceMeta
if err := r.db.WithContext(ctx).
Where("ddns IN ?", deviceIDs).
Find(&list).Error; err != nil {
return nil, err
}
return list, nil
}
func (r *DeviceMetaRepo) DeleteByDeviceID(ctx context.Context, deviceID string) error {
res := r.db.WithContext(ctx).Where("ddns = ?", deviceID).Delete(&model.DeviceMeta{})
if res.Error != nil {
return res.Error
}
if res.RowsAffected == 0 {
return gorm.ErrRecordNotFound
}
return nil
}
func (r *DeviceMetaRepo) MarkOffline(ctx context.Context, deviceID string) error {
return r.db.WithContext(ctx).Exec(
`UPDATE devices SET status='offline' WHERE ddns=?`,
deviceID,
).Error
}
+389
View File
@@ -0,0 +1,389 @@
package sqlite
import (
"context"
"gorm.io/gorm"
"rttys/internal/domain/devicegroup"
"rttys/internal/domain/group"
)
type GroupRepo struct{ db *gorm.DB }
func NewGroupRepo(db *gorm.DB) *GroupRepo { return &GroupRepo{db: db} }
// --- row mapping (only for scan) ---
type userGroupRow struct {
ID int64 `gorm:"column:id"`
Name string `gorm:"column:name"`
Description string `gorm:"column:description"`
}
func (userGroupRow) TableName() string { return "user_groups" }
type deviceGroupRow struct {
ID int64 `gorm:"column:id"`
Name string `gorm:"column:name"`
Description string `gorm:"column:description"`
}
func (deviceGroupRow) TableName() string { return "device_groups" }
type UserGroupBrief struct {
ID int64 `gorm:"column:id"`
Name string `gorm:"column:name"`
}
type DeviceGroupBrief struct {
ID int64 `gorm:"column:id"`
Name string `gorm:"column:name"`
}
type idCountRow struct {
ID int64 `gorm:"column:id"`
Cnt int64 `gorm:"column:cnt"`
}
type deviceGroupUserGroupRow struct {
DeviceGroupID int64 `gorm:"column:device_group_id"`
UserGroupID int64 `gorm:"column:user_group_id"`
UserGroupName string `gorm:"column:user_group_name"`
}
type userGroupDeviceGroupRow struct {
UserGroupID int64 `gorm:"column:user_group_id"`
DeviceGroupID int64 `gorm:"column:device_group_id"`
DeviceGroupName string `gorm:"column:device_group_name"`
}
type userGroupMemberRow struct {
UserID int64 `gorm:"column:user_id"`
UserGroupID int64 `gorm:"column:user_group_id"`
UserGroupName string `gorm:"column:user_group_name"`
}
type DeviceGroupDetail struct {
ID int64
Name string
Description string
DeviceCount int64
UserGroups []UserGroupBrief
}
type UserGroupDetail struct {
ID int64
Name string
Description string
UserCount int64
DeviceGroups []DeviceGroupBrief
}
// -----------------------------
// Visible query
// -----------------------------
func (r *GroupRepo) ListUserGroupsVisibleToUser(ctx context.Context, userID int64, isAdmin bool) ([]group.UserGroup, error) {
if isAdmin {
var rows []userGroupRow
err := r.db.WithContext(ctx).
Raw(`SELECT id,name,description FROM user_groups ORDER BY id`).
Scan(&rows).Error
if err != nil {
return nil, err
}
return mapUserGroups(rows), nil
}
var rows []userGroupRow
err := r.db.WithContext(ctx).
Raw(`
SELECT ug.id,ug.name,ug.description
FROM user_groups ug
JOIN user_group_members ugm ON ugm.group_id=ug.id
WHERE ugm.user_id=?
ORDER BY ug.id`, userID).
Scan(&rows).Error
if err != nil {
return nil, err
}
return mapUserGroups(rows), nil
}
func (r *GroupRepo) ListDeviceGroupIDsByUser(ctx context.Context, userID int64) ([]int64, error) {
var ids []int64
err := r.db.WithContext(ctx).
Raw(`
SELECT DISTINCT l.device_group_id
FROM user_group_members ugm
JOIN user_group_device_group_links l ON l.user_group_id=ugm.group_id
WHERE ugm.user_id=?
ORDER BY l.device_group_id`, userID).
Scan(&ids).Error
if err != nil {
return nil, err
}
return ids, nil
}
func (r *GroupRepo) ListDeviceGroupsVisibleToUser(ctx context.Context, userID int64, isAdmin bool) ([]devicegroup.DeviceGroup, error) {
if isAdmin {
var rows []deviceGroupRow
err := r.db.WithContext(ctx).
Raw(`SELECT id,name,description FROM device_groups ORDER BY id`).
Scan(&rows).Error
if err != nil {
return nil, err
}
return mapDeviceGroups(rows), nil
}
var rows []deviceGroupRow
err := r.db.WithContext(ctx).
Raw(`
SELECT DISTINCT dg.id,dg.name,dg.description
FROM device_groups dg
JOIN user_group_device_group_links l ON l.device_group_id=dg.id
JOIN user_group_members ugm ON ugm.group_id=l.user_group_id
WHERE ugm.user_id=?
ORDER BY dg.id`, userID).
Scan(&rows).Error
if err != nil {
return nil, err
}
return mapDeviceGroups(rows), nil
}
// -----------------------------
// CRUD - use Exec (keeps behavior close to original)
// -----------------------------
func (r *GroupRepo) CreateUserGroup(ctx context.Context, name, description string) (int64, error) {
row := userGroupRow{
Name: name,
Description: description,
}
if err := r.db.WithContext(ctx).Create(&row).Error; err != nil {
return 0, err
}
return row.ID, nil
}
func (r *GroupRepo) UpdateUserGroup(ctx context.Context, id int64, name, description string) error {
return r.db.WithContext(ctx).Exec(
`UPDATE user_groups SET name=?, description=? WHERE id=?`,
name, description, id,
).Error
}
func (r *GroupRepo) DeleteUserGroup(ctx context.Context, id int64) error {
return r.db.WithContext(ctx).Exec(
`DELETE FROM user_groups WHERE id=?`,
id,
).Error
}
func (r *GroupRepo) CreateDeviceGroup(ctx context.Context, name, description string) (int64, error) {
row := deviceGroupRow{
Name: name,
Description: description,
}
if err := r.db.WithContext(ctx).Create(&row).Error; err != nil {
return 0, err
}
return row.ID, nil
}
func (r *GroupRepo) UpdateDeviceGroup(ctx context.Context, id int64, name, description string) error {
return r.db.WithContext(ctx).Exec(
`UPDATE device_groups SET name=?, description=? WHERE id=?`,
name, description, id,
).Error
}
func (r *GroupRepo) DeleteDeviceGroup(ctx context.Context, id int64) error {
return r.db.WithContext(ctx).Exec(
`DELETE FROM device_groups WHERE id=?`,
id,
).Error
}
// -----------------------------
// mapping helpers
// -----------------------------
func mapUserGroups(rows []userGroupRow) []group.UserGroup {
out := make([]group.UserGroup, 0, len(rows))
for _, r := range rows {
out = append(out, group.UserGroup{
ID: r.ID,
Name: r.Name,
Description: r.Description,
})
}
return out
}
func mapDeviceGroups(rows []deviceGroupRow) []devicegroup.DeviceGroup {
out := make([]devicegroup.DeviceGroup, 0, len(rows))
for _, r := range rows {
out = append(out, devicegroup.DeviceGroup{
ID: r.ID,
Name: r.Name,
Description: r.Description,
})
}
return out
}
// ----------------------------------------------------
// Extended listing helpers (counts + relations)
// ----------------------------------------------------
func (r *GroupRepo) ListDeviceGroupDetails(ctx context.Context, userID int64, isAdmin bool) ([]DeviceGroupDetail, error) {
base, err := r.ListDeviceGroupsVisibleToUser(ctx, userID, isAdmin)
if err != nil {
return nil, err
}
if len(base) == 0 {
return []DeviceGroupDetail{}, nil
}
ids := make([]int64, 0, len(base))
detailMap := make(map[int64]*DeviceGroupDetail, len(base))
for _, g := range base {
ids = append(ids, g.ID)
detailMap[g.ID] = &DeviceGroupDetail{
ID: g.ID,
Name: g.Name,
Description: g.Description,
DeviceCount: 0,
UserGroups: []UserGroupBrief{},
}
}
var countRows []idCountRow
if err := r.db.WithContext(ctx).
Raw(`SELECT device_group_id AS id, COUNT(1) AS cnt
FROM devices WHERE device_group_id IN ? GROUP BY device_group_id`, ids).
Scan(&countRows).Error; err != nil {
return nil, err
}
for _, row := range countRows {
if d, ok := detailMap[row.ID]; ok {
d.DeviceCount = row.Cnt
}
}
var linkRows []deviceGroupUserGroupRow
if err := r.db.WithContext(ctx).
Raw(`SELECT l.device_group_id AS device_group_id,
ug.id AS user_group_id,
ug.name AS user_group_name
FROM user_group_device_group_links l
JOIN user_groups ug ON ug.id=l.user_group_id
WHERE l.device_group_id IN ?
ORDER BY l.device_group_id, ug.id`, ids).
Scan(&linkRows).Error; err != nil {
return nil, err
}
for _, row := range linkRows {
if d, ok := detailMap[row.DeviceGroupID]; ok {
d.UserGroups = append(d.UserGroups, UserGroupBrief{ID: row.UserGroupID, Name: row.UserGroupName})
}
}
out := make([]DeviceGroupDetail, 0, len(base))
for _, g := range base {
out = append(out, *detailMap[g.ID])
}
return out, nil
}
func (r *GroupRepo) ListUserGroupDetails(ctx context.Context, userID int64, isAdmin bool) ([]UserGroupDetail, error) {
base, err := r.ListUserGroupsVisibleToUser(ctx, userID, isAdmin)
if err != nil {
return nil, err
}
if len(base) == 0 {
return []UserGroupDetail{}, nil
}
ids := make([]int64, 0, len(base))
detailMap := make(map[int64]*UserGroupDetail, len(base))
for _, g := range base {
ids = append(ids, g.ID)
detailMap[g.ID] = &UserGroupDetail{
ID: g.ID,
Name: g.Name,
Description: g.Description,
UserCount: 0,
DeviceGroups: []DeviceGroupBrief{},
}
}
var countRows []idCountRow
if err := r.db.WithContext(ctx).
Raw(`SELECT group_id AS id, COUNT(1) AS cnt
FROM user_group_members WHERE group_id IN ? GROUP BY group_id`, ids).
Scan(&countRows).Error; err != nil {
return nil, err
}
for _, row := range countRows {
if d, ok := detailMap[row.ID]; ok {
d.UserCount = row.Cnt
}
}
var linkRows []userGroupDeviceGroupRow
if err := r.db.WithContext(ctx).
Raw(`SELECT l.user_group_id AS user_group_id,
dg.id AS device_group_id,
dg.name AS device_group_name
FROM user_group_device_group_links l
JOIN device_groups dg ON dg.id=l.device_group_id
WHERE l.user_group_id IN ?
ORDER BY l.user_group_id, dg.id`, ids).
Scan(&linkRows).Error; err != nil {
return nil, err
}
for _, row := range linkRows {
if d, ok := detailMap[row.UserGroupID]; ok {
d.DeviceGroups = append(d.DeviceGroups, DeviceGroupBrief{ID: row.DeviceGroupID, Name: row.DeviceGroupName})
}
}
out := make([]UserGroupDetail, 0, len(base))
for _, g := range base {
out = append(out, *detailMap[g.ID])
}
return out, nil
}
func (r *GroupRepo) ListUserGroupsByUserIDs(ctx context.Context, userIDs []int64) (map[int64][]UserGroupBrief, error) {
out := make(map[int64][]UserGroupBrief)
if len(userIDs) == 0 {
return out, nil
}
var rows []userGroupMemberRow
if err := r.db.WithContext(ctx).
Raw(`SELECT ugm.user_id AS user_id,
ug.id AS user_group_id,
ug.name AS user_group_name
FROM user_group_members ugm
JOIN user_groups ug ON ug.id=ugm.group_id
WHERE ugm.user_id IN ?
ORDER BY ugm.user_id, ug.id`, userIDs).
Scan(&rows).Error; err != nil {
return nil, err
}
for _, row := range rows {
out[row.UserID] = append(out[row.UserID], UserGroupBrief{ID: row.UserGroupID, Name: row.UserGroupName})
}
return out, nil
}
+221
View File
@@ -0,0 +1,221 @@
package sqlite
import (
"context"
"fmt"
"strings"
"gorm.io/gorm"
)
type RelationsRepo struct{ db *gorm.DB }
func NewRelationsRepo(db *gorm.DB) *RelationsRepo {
return &RelationsRepo{db: db}
}
// ----------------------------------------------------
// user <-> user_groups (cover / set)
// ----------------------------------------------------
func (r *RelationsRepo) SetUserGroups(ctx context.Context, userID int64, groupIDs []int64) error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
// 1) delete old relations
if err := tx.Exec(
`DELETE FROM user_group_members WHERE user_id=?`,
userID,
).Error; err != nil {
return err
}
// 2) insert new relations
if len(groupIDs) == 0 {
return nil
}
vals := make([]string, 0, len(groupIDs))
args := make([]any, 0, len(groupIDs)*2)
for _, gid := range groupIDs {
vals = append(vals, "(?,?)")
args = append(args, userID, gid)
}
q := fmt.Sprintf(
`INSERT INTO user_group_members(user_id, group_id) VALUES %s`,
strings.Join(vals, ","),
)
return tx.Exec(q, args...).Error
})
}
// ----------------------------------------------------
// user_group <-> device_groups (cover / set)
// ----------------------------------------------------
func (r *RelationsRepo) SetUserGroupDeviceGroups(ctx context.Context, userGroupID int64, deviceGroupIDs []int64) error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Exec(
`DELETE FROM user_group_device_group_links WHERE user_group_id=?`,
userGroupID,
).Error; err != nil {
return err
}
if len(deviceGroupIDs) == 0 {
return nil
}
vals := make([]string, 0, len(deviceGroupIDs))
args := make([]any, 0, len(deviceGroupIDs)*2)
for _, dg := range deviceGroupIDs {
vals = append(vals, "(?,?)")
args = append(args, userGroupID, dg)
}
q := fmt.Sprintf(
`INSERT INTO user_group_device_group_links(user_group_id, device_group_id) VALUES %s`,
strings.Join(vals, ","),
)
return tx.Exec(q, args...).Error
})
}
// ----------------------------------------------------
// device_group <-> devices (cover / set, one device -> one group)
// ----------------------------------------------------
func (r *RelationsRepo) SetDeviceGroupDevices(ctx context.Context, deviceGroupID int64, deviceIDs []int64) error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
// case 1: empty list => clear all devices in this group
if len(deviceIDs) == 0 {
return tx.Exec(
`UPDATE devices SET device_group_id=NULL WHERE device_group_id=?`,
deviceGroupID,
).Error
}
// 1) remove devices that are no longer in the group
ph := make([]string, 0, len(deviceIDs))
args := make([]any, 0, len(deviceIDs)+1)
args = append(args, deviceGroupID)
for _, id := range deviceIDs {
ph = append(ph, "?")
args = append(args, id)
}
qRemove := fmt.Sprintf(
`UPDATE devices SET device_group_id=NULL
WHERE device_group_id=?
AND id NOT IN (%s)`,
strings.Join(ph, ","),
)
if err := tx.Exec(qRemove, args...).Error; err != nil {
return err
}
// 2) assign devices to this group (overwrite old group)
ph2 := make([]string, 0, len(deviceIDs))
args2 := make([]any, 0, len(deviceIDs)+1)
args2 = append(args2, deviceGroupID)
for _, id := range deviceIDs {
ph2 = append(ph2, "?")
args2 = append(args2, id)
}
qAssign := fmt.Sprintf(
`UPDATE devices SET device_group_id=?
WHERE id IN (%s)`,
strings.Join(ph2, ","),
)
res := tx.Exec(qAssign, args2...)
if res.Error != nil {
return res.Error
}
return nil
})
}
// ----------------------------------------------------
// device_group <-> user_groups (cover / set)
// ----------------------------------------------------
func (r *RelationsRepo) SetDeviceGroupUserGroups(ctx context.Context, deviceGroupID int64, userGroupIDs []int64) error {
return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Exec(
`DELETE FROM user_group_device_group_links WHERE device_group_id=?`,
deviceGroupID,
).Error; err != nil {
return err
}
if len(userGroupIDs) == 0 {
return nil
}
vals := make([]string, 0, len(userGroupIDs))
args := make([]any, 0, len(userGroupIDs)*2)
for _, ugID := range userGroupIDs {
vals = append(vals, "(?,?)")
args = append(args, ugID, deviceGroupID)
}
q := fmt.Sprintf(
`INSERT INTO user_group_device_group_links(user_group_id, device_group_id) VALUES %s`,
strings.Join(vals, ","),
)
return tx.Exec(q, args...).Error
})
}
// ----------------------------------------------------
// device_group <-> devices (add/remove)
// ----------------------------------------------------
func (r *RelationsRepo) AddDevicesToGroup(ctx context.Context, deviceGroupID int64, deviceIDs []int64) error {
if len(deviceIDs) == 0 {
return nil
}
ph := make([]string, 0, len(deviceIDs))
args := make([]any, 0, len(deviceIDs)+1)
args = append(args, deviceGroupID)
for _, id := range deviceIDs {
ph = append(ph, "?")
args = append(args, id)
}
q := fmt.Sprintf(
`UPDATE devices SET device_group_id=? WHERE id IN (%s)`,
strings.Join(ph, ","),
)
return r.db.WithContext(ctx).Exec(q, args...).Error
}
func (r *RelationsRepo) RemoveDevicesFromGroup(ctx context.Context, deviceGroupID int64, deviceIDs []int64) error {
if len(deviceIDs) == 0 {
return nil
}
ph := make([]string, 0, len(deviceIDs))
args := make([]any, 0, len(deviceIDs)+1)
args = append(args, deviceGroupID)
for _, id := range deviceIDs {
ph = append(ph, "?")
args = append(args, id)
}
q := fmt.Sprintf(
`UPDATE devices SET device_group_id=NULL WHERE device_group_id=? AND id IN (%s)`,
strings.Join(ph, ","),
)
return r.db.WithContext(ctx).Exec(q, args...).Error
}
+100
View File
@@ -0,0 +1,100 @@
PRAGMA foreign_keys = ON;
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT NOT NULL UNIQUE,
email TEXT,
description TEXT NOT NULL DEFAULT '',
password_hash TEXT NOT NULL,
role TEXT NOT NULL CHECK (role IN ('admin','user')),
status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active','disabled')),
is_system INTEGER NOT NULL DEFAULT 0,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE INDEX IF NOT EXISTS idx_users_role ON users(role);
CREATE INDEX IF NOT EXISTS idx_users_status ON users(status);
CREATE TABLE IF NOT EXISTS user_groups (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
description TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE TABLE IF NOT EXISTS user_group_members (
user_id INTEGER NOT NULL,
group_id INTEGER NOT NULL,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
PRIMARY KEY (user_id, group_id),
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE,
FOREIGN KEY (group_id) REFERENCES user_groups(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_ugm_group_id ON user_group_members(group_id);
CREATE TABLE IF NOT EXISTS device_groups (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL UNIQUE,
description TEXT NOT NULL DEFAULT '',
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch())
);
CREATE TABLE IF NOT EXISTS user_group_device_group_links (
user_group_id INTEGER NOT NULL,
device_group_id INTEGER NOT NULL,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
PRIMARY KEY (user_group_id, device_group_id),
FOREIGN KEY (user_group_id) REFERENCES user_groups(id) ON DELETE CASCADE,
FOREIGN KEY (device_group_id) REFERENCES device_groups(id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_ug_dg_device_group_id
ON user_group_device_group_links(device_group_id);
CREATE TABLE IF NOT EXISTS devices (
id INTEGER PRIMARY KEY AUTOINCREMENT,
ddns TEXT NOT NULL UNIQUE,
mac TEXT NOT NULL UNIQUE,
name TEXT NOT NULL DEFAULT '',
description TEXT NOT NULL DEFAULT '',
ip TEXT NOT NULL DEFAULT '',
client TEXT NOT NULL DEFAULT '',
device_group_id INTEGER NULL,
status TEXT NOT NULL DEFAULT 'online' CHECK (status IN ('online','offline','disabled')),
last_seen_at INTEGER NULL,
created_at INTEGER NOT NULL DEFAULT (unixepoch()),
updated_at INTEGER NOT NULL DEFAULT (unixepoch()),
FOREIGN KEY (device_group_id) REFERENCES device_groups(id) ON DELETE SET NULL
);
CREATE INDEX IF NOT EXISTS idx_devices_group_id ON devices(device_group_id);
CREATE INDEX IF NOT EXISTS idx_devices_status ON devices(status);
CREATE INDEX IF NOT EXISTS idx_devices_last_seen ON devices(last_seen_at);
CREATE TRIGGER IF NOT EXISTS trg_users_updated_at
AFTER UPDATE ON users
FOR EACH ROW
BEGIN
UPDATE users SET updated_at = unixepoch() WHERE id = OLD.id;
END;
CREATE TRIGGER IF NOT EXISTS trg_user_groups_updated_at
AFTER UPDATE ON user_groups
FOR EACH ROW
BEGIN
UPDATE user_groups SET updated_at = unixepoch() WHERE id = OLD.id;
END;
CREATE TRIGGER IF NOT EXISTS trg_device_groups_updated_at
AFTER UPDATE ON device_groups
FOR EACH ROW
BEGIN
UPDATE device_groups SET updated_at = unixepoch() WHERE id = OLD.id;
END;
CREATE TRIGGER IF NOT EXISTS trg_devices_updated_at
AFTER UPDATE ON devices
FOR EACH ROW
BEGIN
UPDATE devices SET updated_at = unixepoch() WHERE id = OLD.id;
END;
+113
View File
@@ -0,0 +1,113 @@
package sqlite
import (
"context"
"database/sql"
"fmt"
"os"
"strings"
gormsqlite "github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
type AppDB struct {
gorm *gorm.DB
sql *sql.DB
dsn string
}
func (a *AppDB) Gorm() *gorm.DB { return a.gorm }
func (a *AppDB) SQL() *sql.DB { return a.sql }
func (a *AppDB) Close() error {
if a.sql == nil {
return nil
}
return a.sql.Close()
}
type Options struct {
DSN string // e.g. "/home/database/glkvm-cloud.db"
MaxOpenConns int
MaxIdleConns int
LogSQL bool
}
// Open opens sqlite via GORM (glebarez/sqlite) and exposes both *gorm.DB and *sql.DB.
func Open(ctx context.Context, opt Options) (*AppDB, error) {
if opt.DSN == "" {
return nil, fmt.Errorf("sqlite: DSN is empty")
}
if opt.MaxOpenConns == 0 {
opt.MaxOpenConns = 1
}
if opt.MaxIdleConns == 0 {
opt.MaxIdleConns = 1
}
gormCfg := &gorm.Config{}
if opt.LogSQL {
gormCfg.Logger = logger.Default.LogMode(logger.Info)
}
gdb, err := gorm.Open(gormsqlite.Open(opt.DSN), gormCfg)
if err != nil {
return nil, err
}
raw, err := gdb.DB()
if err != nil {
return nil, err
}
raw.SetMaxOpenConns(opt.MaxOpenConns)
raw.SetMaxIdleConns(opt.MaxIdleConns)
return &AppDB{
gorm: gdb,
sql: raw,
dsn: opt.DSN,
}, nil
}
func InitSchema(ctx context.Context, db *sql.DB, schemaPath string) error {
b, err := os.ReadFile(schemaPath)
if err != nil {
return err
}
if _, err = db.ExecContext(ctx, string(b)); err != nil {
return err
}
if err := ensureDeviceClientColumn(ctx, db); err != nil {
return err
}
return ensureUserIsSystemColumn(ctx, db)
}
func ensureDeviceClientColumn(ctx context.Context, db *sql.DB) error {
if db == nil {
return nil
}
_, err := db.ExecContext(ctx, `ALTER TABLE devices ADD COLUMN client TEXT NOT NULL DEFAULT ''`)
if err == nil {
return nil
}
if strings.Contains(err.Error(), "duplicate column name") {
return nil
}
return err
}
func ensureUserIsSystemColumn(ctx context.Context, db *sql.DB) error {
if db == nil {
return nil
}
_, err := db.ExecContext(ctx, `ALTER TABLE users ADD COLUMN is_system INTEGER NOT NULL DEFAULT 0`)
if err == nil {
return nil
}
if strings.Contains(err.Error(), "duplicate column name") {
return nil
}
return err
}
+173
View File
@@ -0,0 +1,173 @@
package sqlite
import (
"context"
"errors"
"strings"
"gorm.io/gorm"
"rttys/internal/domain/identity"
"rttys/internal/domain/user"
)
type UserRepo struct{ db *gorm.DB }
func NewUserRepo(db *gorm.DB) *UserRepo { return &UserRepo{db: db} }
// 映射用的行结构
type userRow struct {
ID int64 `gorm:"column:id"`
Username string `gorm:"column:username"`
Email string `gorm:"column:email"`
Description string `gorm:"column:description"`
PasswordHash string `gorm:"column:password_hash"`
Role string `gorm:"column:role"`
Status string `gorm:"column:status"`
IsSystem bool `gorm:"column:is_system"`
}
func (userRow) TableName() string { return "users" }
func (r *UserRepo) FindByID(ctx context.Context, id int64) (*user.User, error) {
var row userRow
err := r.db.WithContext(ctx).
Where("id = ?", id).
Take(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("not found")
}
if err != nil {
return nil, err
}
u := &user.User{
ID: row.ID,
Username: row.Username,
Email: row.Email,
Description: row.Description,
PasswordHash: row.PasswordHash,
Role: identity.Role(row.Role),
Status: user.Status(row.Status),
IsSystem: row.IsSystem,
}
return u, nil
}
func (r *UserRepo) FindByUsername(ctx context.Context, username string) (*user.User, error) {
var row userRow
err := r.db.WithContext(ctx).
Where("username = ?", username).
Take(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("not found")
}
if err != nil {
return nil, err
}
return &user.User{
ID: row.ID,
Username: row.Username,
Email: row.Email,
Description: row.Description,
PasswordHash: row.PasswordHash,
Role: identity.Role(row.Role),
Status: user.Status(row.Status),
IsSystem: row.IsSystem,
}, nil
}
func (r *UserRepo) FindSystemAdmin(ctx context.Context) (*user.User, error) {
var row userRow
err := r.db.WithContext(ctx).
Where("is_system = ? AND status = ?", true, "active").
Take(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("system admin not found")
}
if err != nil {
return nil, err
}
return &user.User{
ID: row.ID,
Username: row.Username,
Email: row.Email,
Description: row.Description,
PasswordHash: row.PasswordHash,
Role: identity.Role(row.Role),
Status: user.Status(row.Status),
IsSystem: row.IsSystem,
}, nil
}
func (r *UserRepo) List(ctx context.Context) ([]user.User, error) {
var rows []userRow
if err := r.db.WithContext(ctx).Order("id").Find(&rows).Error; err != nil {
return nil, err
}
out := make([]user.User, 0, len(rows))
for _, row := range rows {
out = append(out, user.User{
ID: row.ID,
Username: row.Username,
Email: row.Email,
Description: row.Description,
PasswordHash: row.PasswordHash,
Role: identity.Role(row.Role),
Status: user.Status(row.Status),
IsSystem: row.IsSystem,
})
}
return out, nil
}
func (r *UserRepo) Create(ctx context.Context, u *user.User) (int64, error) {
row := userRow{
Username: u.Username,
Email: u.Email,
Description: u.Description,
PasswordHash: u.PasswordHash,
Role: string(u.Role),
Status: string(u.Status),
IsSystem: u.IsSystem,
}
if err := r.db.WithContext(ctx).Create(&row).Error; err != nil {
return 0, err
}
return row.ID, nil
}
func (r *UserRepo) Update(ctx context.Context, u *user.User) error {
// 用 Updates 可以避免全量 Save 带来的误更新
return r.db.WithContext(ctx).
Model(&userRow{}).
Where("id = ?", u.ID).
Updates(map[string]any{
"username": u.Username,
"email": u.Email,
"description": u.Description,
"password_hash": u.PasswordHash,
"role": string(u.Role),
"status": string(u.Status),
"is_system": u.IsSystem,
}).Error
}
func (r *UserRepo) Delete(ctx context.Context, id int64) error {
return r.db.WithContext(ctx).
Exec("DELETE FROM users WHERE id = ?", id).Error
}
func IsUniqueViolation(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique constraint failed")
}
-376
View File
@@ -1,376 +0,0 @@
/*
* @Author: CU-Jon
* @Date: 2025-09-26 13:28:12 EDT
* @LastEditors: CU-Jon
* @LastEditTime: 2025-09-26 14:02:57 EDT
* @FilePath: \glkvm-cloud\ldap.go
* @Description: LDAP认证模块 (LDAP authentication module)
*/
package main
import (
"crypto/tls"
"fmt"
"strings"
"time"
"github.com/go-ldap/ldap/v3"
"github.com/rs/zerolog/log"
)
// LDAP认证器结构体 (LDAP authenticator struct)
type LDAPAuthenticator struct {
config *Config
}
// 创建新的LDAP认证器 (Create new LDAP authenticator)
func NewLDAPAuthenticator(config *Config) *LDAPAuthenticator {
return &LDAPAuthenticator{config: config}
}
// 执行用户LDAP认证 (Perform LDAP authentication for a user)
func (l *LDAPAuthenticator) Authenticate(username, password string) (bool, error) {
if !l.config.LdapEnabled {
return false, fmt.Errorf("LDAP authentication is disabled")
}
if username == "" || password == "" {
return false, fmt.Errorf("username and password are required")
}
// 连接到LDAP服务器 (Connect to LDAP server)
conn, err := l.connect()
if err != nil {
return false, fmt.Errorf("failed to connect to LDAP server: %v", err)
}
defer conn.Close()
// 使用服务账户进行绑定和搜索 (Use service account for binding and searching)
if l.config.LdapBindDN == "" || l.config.LdapBindPassword == "" {
return false, fmt.Errorf("service account credentials are required for LDAP authentication - BindDN empty: %v, BindPassword empty: %v", l.config.LdapBindDN == "", l.config.LdapBindPassword == "")
}
err = conn.Bind(l.config.LdapBindDN, l.config.LdapBindPassword)
if err != nil {
return false, fmt.Errorf("service account bind failed: %v", err)
} // 使用服务账户搜索用户 (Use service account to search for user)
userDN, err := l.findUserDN(conn, username)
if err != nil {
return false, fmt.Errorf("user search failed: %v", err)
}
// 找到用户,现在用用户凭证验证密码 (Found user, now validate password with user credentials)
err = conn.Bind(userDN, password)
if err != nil {
return false, fmt.Errorf("password validation failed: %v", err)
}
// 重新绑定为服务账户以进行授权检查 (Rebind as service account for authorization check)
err = conn.Bind(l.config.LdapBindDN, l.config.LdapBindPassword)
if err != nil {
return false, fmt.Errorf("failed to rebind as service account for authorization: %v", err)
}
// 检查用户授权 (Check user authorization)
authorized, err := l.checkAuthorization(conn, userDN, username)
if err != nil {
return false, fmt.Errorf("authorization check failed: %v", err)
}
if !authorized {
return false, fmt.Errorf("user not authorized")
}
return true, nil
}
// 建立到LDAP服务器的连接 (Establish connection to LDAP server)
func (l *LDAPAuthenticator) connect() (*ldap.Conn, error) {
address := fmt.Sprintf("%s:%d", l.config.LdapServer, l.config.LdapPort)
var conn *ldap.Conn
var err error
if l.config.LdapUseTLS {
// TLS配置 (TLS configuration)
tlsConfig := &tls.Config{
ServerName: l.config.LdapServer,
InsecureSkipVerify: true, // 跳过证书验证以避免自签名证书问题 (Skip certificate verification to avoid self-signed certificate issues)
}
if l.config.LdapPort == 636 {
// 使用LDAPS (直接TLS连接) (Use LDAPS - direct TLS connection)
conn, err = ldap.DialTLS("tcp", address, tlsConfig)
} else {
// 使用StartTLS (先连接再升级到TLS) (Use StartTLS - connect first then upgrade to TLS)
conn, err = ldap.Dial("tcp", address)
if err == nil {
err = conn.StartTLS(tlsConfig)
}
}
} else {
// 使用普通连接 (Use plain connection)
conn, err = ldap.Dial("tcp", address)
}
if err != nil {
return nil, err
}
// 设置超时时间 (Set timeout)
conn.SetTimeout(10 * time.Second)
return conn, nil
}
// 基于用户名搜索用户DN (Search for user DN based on username)
func (l *LDAPAuthenticator) findUserDN(conn *ldap.Conn, username string) (string, error) {
// 准备搜索过滤器 (Prepare search filter)
filter := fmt.Sprintf(l.config.LdapUserFilter, username)
if l.config.LdapUserFilter == "" {
filter = fmt.Sprintf("(uid=%s)", username)
}
// 执行搜索 (Perform search)
searchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, // 无大小限制 (No size limit)
0, // 无时间限制 (No time limit)
false,
filter,
[]string{"dn"},
nil,
)
sr, err := conn.Search(searchRequest)
if err != nil {
return "", err
}
if len(sr.Entries) == 0 {
return "", fmt.Errorf("user not found")
}
if len(sr.Entries) > 1 {
return "", fmt.Errorf("multiple users found")
}
return sr.Entries[0].DN, nil
}
// 基于组或用户列表检查用户是否授权 (Check if user is authorized based on groups or users list)
func (l *LDAPAuthenticator) checkAuthorization(conn *ldap.Conn, userDN, username string) (bool, error) {
// 如果没有配置限制,则允许所有已认证用户 (If no restrictions are configured, allow all authenticated users)
if l.config.LdapAllowedGroups == "" && l.config.LdapAllowedUsers == "" {
return true, nil
}
// 检查允许的用户列表 (Check allowed users list)
if l.config.LdapAllowedUsers != "" {
allowedUsers := strings.Split(strings.TrimSpace(l.config.LdapAllowedUsers), ",")
for _, allowedUser := range allowedUsers {
if strings.TrimSpace(allowedUser) == username {
return true, nil
}
}
}
// 检查允许的组 (Check allowed groups)
if l.config.LdapAllowedGroups != "" {
return l.checkGroupMembership(conn, userDN, username)
}
return false, nil
}
// 检查用户是否属于任何允许的组 (Check if user belongs to any of the allowed groups)
func (l *LDAPAuthenticator) checkGroupMembership(conn *ldap.Conn, userDN, username string) (bool, error) {
allowedGroups := strings.Split(strings.TrimSpace(l.config.LdapAllowedGroups), ",")
for _, group := range allowedGroups {
group = strings.TrimSpace(group)
if group == "" {
continue
}
// 搜索组成员关系 - 尝试不同的常见LDAP组结构 (Search for group membership - try different common LDAP group structures)
isMember, err := l.isGroupMember(conn, userDN, username, group)
if err != nil {
log.Warn().Msgf("Error checking group membership for %s in %s: %v", username, group, err)
continue
}
if isMember {
return true, nil
}
}
return false, nil
}
// 检查用户是否是指定组的成员 (Check if user is a member of the specified group)
func (l *LDAPAuthenticator) isGroupMember(conn *ldap.Conn, userDN, username, groupName string) (bool, error) {
// 首先查找用户的实际DN,因为我们可能使用了UPN格式进行认证 (First find the user's actual DN, as we may have used UPN format for authentication)
actualUserDN, err := l.findActualUserDN(conn, username)
if err != nil {
actualUserDN = userDN // 回退到原始DN (Fallback to original DN)
}
// 尝试不同的常见组搜索模式 (Try different common group search patterns)
// 模式1:通过CN搜索组并检查成员属性 (Pattern 1: Search for group by CN and check member attribute)
groupFilter := fmt.Sprintf("(cn=%s)", groupName)
groupSearchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, 0, false,
groupFilter,
[]string{"member", "memberUid", "uniqueMember"},
nil,
)
sr, err := conn.Search(groupSearchRequest)
if err != nil {
log.Warn().Msgf("Group search failed: %v", err)
return false, err
}
for _, entry := range sr.Entries {
members := entry.GetAttributeValues("member")
memberUids := entry.GetAttributeValues("memberUid")
uniqueMembers := entry.GetAttributeValues("uniqueMember")
// 检查member属性(完整DN) (Check member attribute - full DN)
for _, member := range members {
if member == userDN || member == actualUserDN {
return true, nil
}
// 也检查是否member DN包含用户名 (Also check if member DN contains the username)
if strings.Contains(strings.ToLower(member), strings.ToLower("cn="+username)) {
return true, nil
}
}
// 检查memberUid属性(仅用户名) (Check memberUid attribute - username only)
for _, memberUid := range memberUids {
if memberUid == username {
return true, nil
}
}
// 检查uniqueMember属性(完整DN) (Check uniqueMember attribute - full DN)
for _, uniqueMember := range uniqueMembers {
if uniqueMember == userDN {
return true, nil
}
}
}
// 模式2:通过用户名搜索用户并检查memberOf属性 (Pattern 2: Search for user by username and check memberOf attribute)
// 使用配置的用户过滤器或默认的uid过滤器 (Use configured user filter or default uid filter)
userFilter := fmt.Sprintf(l.config.LdapUserFilter, username)
if l.config.LdapUserFilter == "" {
userFilter = fmt.Sprintf("(uid=%s)", username)
}
userSearchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, 0, false,
userFilter,
[]string{"memberOf", "distinguishedName"},
nil,
)
sr, err = conn.Search(userSearchRequest)
if err != nil {
log.Warn().Msgf("User search for memberOf failed: %v", err)
} else {
for _, entry := range sr.Entries {
memberOfValues := entry.GetAttributeValues("memberOf")
for _, memberOf := range memberOfValues {
if strings.Contains(strings.ToLower(memberOf), strings.ToLower("cn="+groupName)) {
return true, nil
}
}
}
}
return false, nil
}
// 查找用户的实际DN (Find the user's actual DN)
func (l *LDAPAuthenticator) findActualUserDN(conn *ldap.Conn, username string) (string, error) {
// 使用配置的用户过滤器搜索用户 (Search for user using configured user filter)
userFilter := fmt.Sprintf(l.config.LdapUserFilter, username)
if l.config.LdapUserFilter == "" {
userFilter = fmt.Sprintf("(uid=%s)", username)
}
userSearchRequest := ldap.NewSearchRequest(
l.config.LdapBaseDN,
ldap.ScopeWholeSubtree,
ldap.NeverDerefAliases,
0, 0, false,
userFilter,
[]string{"distinguishedName"},
nil,
)
sr, err := conn.Search(userSearchRequest)
if err != nil {
return "", err
}
if len(sr.Entries) == 0 {
return "", fmt.Errorf("user not found")
}
if len(sr.Entries) > 1 {
return "", fmt.Errorf("multiple users found")
}
return sr.Entries[0].DN, nil
}
// 执行用户认证,支持LDAP和传统密码认证 (Perform user authentication with LDAP and legacy password support)
func AuthenticateUser(cfg *Config, username, password, authMethod string) bool {
success, _ := AuthenticateUserWithError(cfg, username, password, authMethod)
return success
}
// 执行用户认证并返回错误类型,支持LDAP和传统密码认证 (Perform user authentication with error type, supporting LDAP and legacy password authentication)
func AuthenticateUserWithError(cfg *Config, username, password, authMethod string) (bool, string) {
// 处理LDAP认证 (Handle LDAP authentication)
if cfg.LdapEnabled && authMethod == "ldap" && username != "" {
ldapAuth := NewLDAPAuthenticator(cfg)
success, err := ldapAuth.Authenticate(username, password)
if err != nil {
log.Error().Msgf("LDAP authentication error: %v", err)
// 检查错误类型以区分认证和授权错误 (Check error type to distinguish between authentication and authorization errors)
if strings.Contains(err.Error(), "user not authorized") {
return false, "authorization"
}
return false, "authentication"
}
return success, ""
}
// 回退到原始密码认证以保持向后兼容 (Fallback to original password authentication for backward compatibility)
if authMethod == "legacy" || authMethod == "" {
if cfg.Password == password {
return true, ""
}
return false, "authentication"
}
return false, "authentication"
}
-308
View File
@@ -1,308 +0,0 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
import (
"context"
"encoding/json"
_ "net/http/pprof"
"os"
"rttys/db"
"runtime"
"runtime/debug"
xlog "rttys/log"
"github.com/rs/zerolog"
"github.com/rs/zerolog/log"
"github.com/urfave/cli/v3"
)
const RttysVersion = "5.2.0"
const KVMCloudVersion = "v1.8.0"
var (
GitCommit = ""
BuildTime = ""
)
func main() {
defaultLogPath := "/var/log/rttys.log"
if runtime.GOOS == "windows" {
defaultLogPath = "rttys.log"
}
cmd := &cli.Command{
Name: "rttys",
Usage: "The server side for rtty",
Version: RttysVersion,
Flags: []cli.Flag{
&cli.StringFlag{
Name: "log",
Value: defaultLogPath,
Usage: "log file path",
},
&cli.StringFlag{
Name: "log-level",
Value: "info",
Usage: "log level(debug, info, warn, error)",
},
&cli.StringFlag{
Name: "conf",
Aliases: []string{"c"},
Usage: "config file to load",
},
&cli.StringFlag{
Name: "addr-dev",
Value: ":5912",
Usage: "address to listen device",
},
&cli.StringFlag{
Name: "addr-user",
Value: ":5913",
Usage: "address to listen user",
},
&cli.StringFlag{
Name: "addr-http-proxy",
Usage: "address to listen for HTTP proxy (default auto)",
},
&cli.StringFlag{
Name: "http-proxy-redir-url",
Usage: "url to redirect for HTTP proxy",
},
&cli.StringFlag{
Name: "http-proxy-redir-domain",
Usage: "domain for HTTP proxy set cookie",
},
&cli.StringFlag{
Name: "token",
Aliases: []string{"t"},
Usage: "token to use",
},
&cli.StringFlag{
Name: "dev-hook-url",
Usage: "called when the device is connected",
},
&cli.StringFlag{
Name: "user-hook-url",
Usage: "called when user accesses /connect/:devid, /cmd/:devid, /web/, or /web2/ APIs",
},
&cli.BoolFlag{
Name: "local-auth",
Value: true,
Usage: "need auth for local",
},
&cli.StringFlag{
Name: "password",
Usage: "web management password",
},
&cli.BoolFlag{
Name: "allow-origins",
Usage: "allow all origins for cross-domain request",
},
&cli.BoolFlag{
Name: "ldap-enabled",
Usage: "enable LDAP authentication",
},
&cli.StringFlag{
Name: "ldap-server",
Usage: "LDAP server hostname or IP",
},
&cli.IntFlag{
Name: "ldap-port",
Value: 389,
Usage: "LDAP server port",
},
&cli.BoolFlag{
Name: "ldap-use-tls",
Usage: "use TLS/SSL for LDAP connection",
},
&cli.StringFlag{
Name: "ldap-bind-dn",
Usage: "LDAP bind DN for service account",
},
&cli.StringFlag{
Name: "ldap-bind-password",
Usage: "LDAP bind password for service account",
},
&cli.StringFlag{
Name: "ldap-base-dn",
Usage: "LDAP base DN for user searches",
},
&cli.StringFlag{
Name: "ldap-user-filter",
Value: "(uid=%s)",
Usage: "LDAP user filter",
},
&cli.StringFlag{
Name: "ldap-allowed-groups",
Usage: "comma-separated list of allowed LDAP groups",
},
&cli.StringFlag{
Name: "ldap-allowed-users",
Usage: "comma-separated list of allowed LDAP users",
},
&cli.StringFlag{
Name: "pprof",
Usage: "enable pprof and listen on specified address (e.g. localhost:6060)",
},
// ---- OIDC Authentication (generic OIDC provider) ----
&cli.BoolFlag{
Name: "oidc-enabled",
Usage: "enable OIDC authentication (OpenID Connect)",
},
&cli.StringFlag{
Name: "oidc-generic-client-id",
Usage: "OIDC client ID (issued by the identity provider)",
},
&cli.StringFlag{
Name: "oidc-generic-client-secret",
Usage: "OIDC client secret (read from OIDC_GENERIC_CLIENT_SECRET env by default)",
},
&cli.StringFlag{
Name: "oidc-generic-auth-url",
Usage: "OIDC authorization endpoint URL",
},
&cli.StringFlag{
Name: "oidc-generic-token-url",
Usage: "OIDC token endpoint URL",
},
&cli.StringFlag{
Name: "oidc-generic-redirect-url",
Usage: "OIDC redirect/callback URL (must match one registered in IdP)",
},
&cli.StringFlag{
Name: "oidc-generic-scopes",
Value: "openid profile email",
Usage: "space-separated list of OIDC scopes",
},
&cli.StringFlag{
Name: "oidc-generic-allowed-users",
Usage: "optional email whitelist for OIDC logins (exact emails or @domain, space/comma-separated)",
},
&cli.StringFlag{
Name: "oidc-generic-allowed-subs",
Usage: "optional subject (sub) whitelist for OIDC logins (space/comma-separated)",
},
&cli.StringFlag{
Name: "oidc-generic-allowed-usernames",
Usage: "optional username whitelist for OIDC logins (preferred_username/name, space/comma-separated)",
},
&cli.StringFlag{
Name: "oidc-generic-allowed-groups",
Usage: "optional groups whitelist for OIDC logins (space/comma-separated)",
},
&cli.BoolFlag{
Name: "verbose",
Aliases: []string{"V"},
Usage: "more detailed output",
},
},
Action: cmdAction,
}
err := cmd.Run(context.Background(), os.Args)
if err != nil {
log.Fatal().Msg(err.Error())
}
}
func cmdAction(c context.Context, cmd *cli.Command) error {
defer logPanic()
xlog.SetPath(cmd.String("log"))
switch cmd.String("log-level") {
case "debug":
zerolog.SetGlobalLevel(zerolog.DebugLevel)
case "warn":
zerolog.SetGlobalLevel(zerolog.WarnLevel)
case "error":
zerolog.SetGlobalLevel(zerolog.ErrorLevel)
default:
zerolog.SetGlobalLevel(zerolog.InfoLevel)
}
if cmd.Bool("verbose") {
xlog.Verbose()
}
log.Info().Msg("Go Version: " + runtime.Version())
log.Info().Msgf("Go OS/Arch: %s/%s", runtime.GOOS, runtime.GOARCH)
log.Info().Msg("Rttys Version: " + RttysVersion)
if GitCommit != "" {
log.Info().Msg("Git Commit: " + GitCommit)
}
if BuildTime != "" {
log.Info().Msg("Build Time: " + BuildTime)
}
if runtime.GOOS != "windows" {
go signalHandle()
}
cfg := Config{
AddrDev: ":5912",
AddrUser: ":5913",
LocalAuth: true,
}
err := cfg.Parse(cmd)
if err != nil {
return err
}
// ===== 打印完整配置(验证配置是否加载正确) =====
{
importJSON, _ := json.MarshalIndent(cfg, "", " ")
log.Info().Msg("==== Loaded Configuration ====")
log.Info().Msg(string(importJSON))
log.Info().Msg("==============================")
}
// Initialize the SQLite database connection
db.Init()
srv := &RttyServer{cfg: cfg}
return srv.Run()
}
func logPanic() {
if r := recover(); r != nil {
saveCrashLog(r, debug.Stack())
os.Exit(2)
}
}
func saveCrashLog(p any, stack []byte) {
log.Error().Msgf("%v", p)
log.Error().Msg(string(stack))
}
Executable
+15
View File
@@ -0,0 +1,15 @@
package model
// DeviceMeta represents device metadata stored in the devices table.
type DeviceMeta struct {
DeviceID string `gorm:"column:ddns"` // DeviceID maps to devices.ddns.
Mac string `gorm:"column:mac"` // Mac is the unique and immutable MAC address of the device.
IP string `gorm:"column:ip"` // IP is the current IP address of the device.
Description string `gorm:"column:description"` // Description is a human-readable description of the device.
Client string `gorm:"column:client"` // Client reported by device (e.g. "rtty-go").
}
// TableName sets the name of the table in the database that this struct binds to.
func (DeviceMeta) TableName() string {
return "devices"
}
-15
View File
@@ -1,15 +0,0 @@
# Authentication token for device connections
token: b1d4cdb2a3cd6a5634aa3599afcddcf5
# Web management password
password: 'b1d4cdb2a3cdr~t!tys#7ngz4hyc'
# Webrtc
webrtc-ip: 107.173.152.173
webrtc-port: 3478
webrtc-username: r92k8a3mz7t1xn6p0fj4qgldwcev5u9h
webrtc-password: uq4z7nwb1x30sy5jfgm69lpdrhta2evi
addr-dev: :5912
addr-user: :443
addr-http-proxy: :10443
-14
View File
@@ -1,14 +0,0 @@
[Unit]
Description=RTTYS - Remote Terminal Access Server
Documentation=https://github.com/zhaojh329/rttys
After=network.target
[Service]
Type=simple
ExecStart=/usr/bin/rttys -c /etc/rttys/rttys.conf
Restart=always
RestartSec=5
TimeoutStopSec=10
[Install]
WantedBy=multi-user.target
-200
View File
@@ -1,200 +0,0 @@
/*
* MIT License
*
* Copyright (c) 2019 Jianhui Zhao <zhaojh329@gmail.com>
*
* Permission is hereby granted, free of charge, to any person obtaining a copy
* of this software and associated documentation files (the "Software"), to deal
* in the Software without restriction, including without limitation the rights
* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
* copies of the Software, and to permit persons to whom the Software is
* furnished to do so, subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
* SOFTWARE.
*/
package main
import (
"context"
"flag"
"io"
"net/http"
"net/http/cookiejar"
"net/url"
"strings"
"sync"
"testing"
"time"
"github.com/gorilla/websocket"
"github.com/rs/zerolog/log"
)
func TestRttysStress(t *testing.T) {
duration := 10 * time.Minute
timeoutFlag := flag.Lookup("test.timeout")
if timeoutFlag != nil {
duration = timeoutFlag.Value.(flag.Getter).Get().(time.Duration)
}
cfg := Config{
AddrDev: ":5912",
AddrUser: ":5913",
}
srv := &RttyServer{cfg: cfg}
go func() {
err := srv.Run()
if err != nil {
log.Fatal().Msg(err.Error())
}
}()
ctx, cancel := context.WithTimeout(context.Background(), duration-time.Second*2)
defer cancel()
time.Sleep(time.Millisecond * 100)
log.Info().Msg("Waiting for devices to connect for testing...")
devices := &sync.Map{}
for {
select {
case <-ctx.Done():
log.Info().Msg("Test timeout, exiting...")
return
default:
time.Sleep(time.Second * 1)
srv.groups.Range(func(key, value any) bool {
group := key.(string)
g := value.(*DeviceGroup)
g.devices.Range(func(key, value any) bool {
dev := value.(*Device)
if _, loaded := devices.LoadOrStore(dev.id, group+dev.id); !loaded {
go runDeviceTest(ctx, devices, group, dev.id)
}
return true
})
return true
})
}
}
}
func runDeviceTest(ctx context.Context, devices *sync.Map, group, devID string) {
ctx, cancel := context.WithCancel(ctx)
defer func() {
time.Sleep(time.Second)
cancel()
devices.Delete(group + devID)
}()
go runHttpTest(ctx, group, devID)
wg := &sync.WaitGroup{}
for range 7 {
wg.Add(1)
go runWebSocketTest(ctx, group, devID, wg)
}
wg.Wait()
}
func runWebSocketTest(ctx context.Context, group, devID string, wg *sync.WaitGroup) {
conn, _, err := websocket.DefaultDialer.Dial("ws://127.0.0.1:5913/connect/"+devID+"?group="+group, nil)
if err != nil {
log.Fatal().Msg(err.Error())
}
defer conn.Close()
defer wg.Done()
go func() {
<-ctx.Done()
conn.Close()
}()
go func() {
msg := []byte{0}
msg = append(msg, []byte("ttttttttttttttttttttttttttttt\n")...)
msg = append(msg, []byte("ttttttttttttttttttttttttttttt\n")...)
for {
err = conn.WriteMessage(websocket.BinaryMessage, msg)
if err != nil {
return
}
time.Sleep(time.Millisecond * 20)
}
}()
for {
_, _, err := conn.ReadMessage()
if err != nil {
return
}
}
}
func runHttpTest(ctx context.Context, group, devID string) {
for {
select {
case <-ctx.Done():
return
default:
runHttpTestOnce(ctx, group, devID)
}
}
}
func runHttpTestOnce(ctx context.Context, group, devID string) {
addr := ""
if group == "" {
addr = "http://127.0.0.1:5913/web/"
} else {
addr = "http://127.0.0.1:5913/web2/" + group + "/"
}
addr += devID + "/http/" + encodeURIComponent("127.0.0.1:80/")
jar, _ := cookiejar.New(nil)
client := &http.Client{
Jar: jar,
}
request, _ := http.NewRequestWithContext(ctx, "GET", addr, nil)
for range 10 {
res, err := client.Do(request)
if err != nil {
log.Info().Msg(err.Error())
return
}
defer res.Body.Close()
io.ReadAll(res.Body)
time.Sleep(10 * time.Millisecond)
}
}
func encodeURIComponent(str string) string {
r := url.QueryEscape(str)
r = strings.ReplaceAll(r, "+", "%20")
return r
}
+411
View File
@@ -0,0 +1,411 @@
#!/bin/sh
# ============================================================================
# rtty-go client one-click installer for Linux
# Supports: x86_64, x86, arm64, armv7 (Ubuntu/Debian/CentOS/OpenWrt/Raspberry Pi)
# ============================================================================
set -e
# --------------- Configuration (can be pre-filled by server) ----------------
RTTY_HOST=""
RTTY_PORT="5912"
RTTY_TOKEN=""
RTTY_SSL="-s -x"
DOWNLOAD_BASE_URL="https://kvm-cloud.gl-inet.com/selfhost/clients"
# ----------------------------------------------------------------------------
INSTALL_DIR="/usr/local/bin"
BINARY_NAME="rtty-go"
CONFIG_DIR="/etc/rtty-go"
CONFIG_FILE="${CONFIG_DIR}/env"
# ========================== Helper Functions =================================
log_info() { printf "\033[32m[INFO]\033[0m %s\n" "$1"; }
log_warn() { printf "\033[33m[WARN]\033[0m %s\n" "$1"; }
log_error() { printf "\033[31m[ERROR]\033[0m %s\n" "$1"; }
command_exists() { command -v "$1" >/dev/null 2>&1; }
usage() {
cat <<EOF
Usage: $0 -h <host> -t <token> [-p <port>] [-u <download_base_url>]
Options:
-h Server host or IP address (required)
-t Authorization token (required)
-p Server port (default: 5912)
-u Download base URL for binaries
--uninstall Uninstall rtty-go client
Example:
$0 -h 107.173.152.173 -t lHEP7GyyGt4S18KlyikfpvzdZTVxnD8v
EOF
exit 1
}
# ========================== Parse Arguments ==================================
UNINSTALL=0
while [ $# -gt 0 ]; do
case "$1" in
-h) RTTY_HOST="$2"; shift 2 ;;
-p) RTTY_PORT="$2"; shift 2 ;;
-t) RTTY_TOKEN="$2"; shift 2 ;;
-u) DOWNLOAD_BASE_URL="$2"; shift 2 ;;
--uninstall) UNINSTALL=1; shift ;;
*) usage ;;
esac
done
# Strip leading colon from port (e.g. ":5912" -> "5912")
RTTY_PORT="${RTTY_PORT#:}"
# ========================== Detect Platform ==================================
detect_arch() {
ARCH=$(uname -m)
case "$ARCH" in
x86_64|amd64) echo "linux-amd64" ;;
i386|i486|i586|i686) echo "linux-386" ;;
aarch64|arm64) echo "linux-arm64" ;;
armv7*|armhf) echo "linux-armv7" ;;
armv6*) echo "linux-armv7" ;;
*) log_error "Unsupported architecture: $ARCH"; exit 1 ;;
esac
}
# ========================== Detect Init System ===============================
detect_init() {
if [ -f /etc/openwrt_release ]; then
echo "procd"
elif command_exists systemctl && systemctl --version >/dev/null 2>&1; then
echo "systemd"
elif command_exists rc-service; then
echo "openrc"
else
echo "sysvinit"
fi
}
# ========================== Get MAC Address ==================================
get_mac() {
MAC=""
# Method 1: OpenWrt GL.iNet specific
if [ -f /proc/gl-hw-info/device_mac ]; then
MAC=$(cat /proc/gl-hw-info/device_mac 2>/dev/null)
fi
# Method 2: ip command (most modern Linux)
if [ -z "$MAC" ] && command_exists ip; then
MAC=$(ip link show 2>/dev/null | awk '
/^[0-9]+:/ { iface=$2; gsub(/:$/, "", iface) }
/ether/ && iface !~ /^(lo|docker|br-|veth|virbr)/ { print $2; exit }
')
fi
# Method 3: /sys/class/net (fallback)
if [ -z "$MAC" ]; then
for iface_path in /sys/class/net/*; do
iface=$(basename "$iface_path")
case "$iface" in lo|docker*|br-*|veth*|virbr*) continue ;; esac
if [ -f "${iface_path}/address" ]; then
addr=$(cat "${iface_path}/address" 2>/dev/null)
if [ -n "$addr" ] && [ "$addr" != "00:00:00:00:00:00" ]; then
MAC="$addr"
break
fi
fi
done
fi
# Method 4: Generate random locally-administered MAC
if [ -z "$MAC" ] || [ "$MAC" = "00:00:00:00:00:00" ]; then
log_warn "Could not detect MAC address, generating random one"
# Use /proc/sys/kernel/random/uuid (available on all Linux, no external tools needed)
_uuid=$(cat /proc/sys/kernel/random/uuid | tr -d '-')
MAC=$(echo "$_uuid" | sed 's/\(..\)/\1:/g' | cut -c1-14)
MAC="02:${MAC}"
fi
echo "$MAC"
}
# ========================== Generate Device ID ===============================
gen_device_id() {
# OpenWrt GL.iNet specific
if [ -f /proc/gl-hw-info/device_ddns ]; then
cat /proc/gl-hw-info/device_ddns
return
fi
# Generate random ID: 8 hex chars (use kernel uuid, no external tools needed)
cat /proc/sys/kernel/random/uuid | tr -d '-' | cut -c1-8
}
# ========================== Download Binary ==================================
download_file() {
url="$1"
dest="$2"
if command_exists curl; then
curl -fsSL -o "$dest" "$url"
elif command_exists wget; then
wget -qO "$dest" "$url"
else
log_error "Neither curl nor wget found. Please install one."
exit 1
fi
}
# ========================== Uninstall ========================================
do_uninstall() {
log_info "Uninstalling rtty-go client..."
INIT_SYSTEM=$(detect_init)
case "$INIT_SYSTEM" in
systemd)
systemctl stop rtty-go 2>/dev/null || true
systemctl disable rtty-go 2>/dev/null || true
rm -f /etc/systemd/system/rtty-go.service
systemctl daemon-reload 2>/dev/null || true
;;
procd)
/etc/init.d/rtty-go stop 2>/dev/null || true
/etc/init.d/rtty-go disable 2>/dev/null || true
rm -f /etc/init.d/rtty-go
;;
openrc)
rc-service rtty-go stop 2>/dev/null || true
rc-update del rtty-go default 2>/dev/null || true
rm -f /etc/init.d/rtty-go
;;
*)
crontab -l 2>/dev/null | grep -v "rtty-go" | crontab - 2>/dev/null || true
pkill -f "${INSTALL_DIR}/${BINARY_NAME}" 2>/dev/null || true
;;
esac
rm -f "${INSTALL_DIR}/${BINARY_NAME}"
rm -rf "${CONFIG_DIR}"
log_info "Uninstall complete."
exit 0
}
# ========================== Service Setup ====================================
setup_systemd() {
log_info "Setting up systemd service..."
cat > /etc/systemd/system/rtty-go.service <<EOF
[Unit]
Description=rtty-go remote terminal client
After=network-online.target
Wants=network-online.target
[Service]
Type=simple
EnvironmentFile=${CONFIG_FILE}
ExecStart=${INSTALL_DIR}/${BINARY_NAME} ${RTTY_SSL} -a -I \${DEVICE_ID} -h \${RTTY_HOST} -p \${RTTY_PORT} -t \${RTTY_TOKEN} -d \${DEVICE_MAC}
Restart=always
RestartSec=10
[Install]
WantedBy=multi-user.target
EOF
systemctl daemon-reload
systemctl enable rtty-go
systemctl restart rtty-go
log_info "systemd service started."
}
setup_procd() {
log_info "Setting up OpenWrt procd service..."
cat > /etc/init.d/rtty-go <<'INITEOF'
#!/bin/sh /etc/rc.common
START=99
STOP=10
USE_PROCD=1
start_service() {
. /etc/rtty-go/env
procd_open_instance
procd_set_param command /usr/local/bin/rtty-go \
INITEOF
# Append the dynamic parts
cat >> /etc/init.d/rtty-go <<EOF
${RTTY_SSL} -a \\
EOF
cat >> /etc/init.d/rtty-go <<'INITEOF'
-I "$DEVICE_ID" \
-h "$RTTY_HOST" \
-p "$RTTY_PORT" \
-t "$RTTY_TOKEN" \
-d "$DEVICE_MAC"
procd_set_param respawn 3600 5 0
procd_close_instance
}
INITEOF
chmod +x /etc/init.d/rtty-go
/etc/init.d/rtty-go enable
/etc/init.d/rtty-go restart
log_info "OpenWrt procd service started."
}
setup_openrc() {
log_info "Setting up OpenRC service..."
cat > /etc/init.d/rtty-go <<EOF
#!/sbin/openrc-run
description="rtty-go remote terminal client"
command="${INSTALL_DIR}/${BINARY_NAME}"
command_args="${RTTY_SSL} -a -I \${DEVICE_ID} -h \${RTTY_HOST} -p \${RTTY_PORT} -t \${RTTY_TOKEN} -d \${DEVICE_MAC}"
command_background=true
pidfile="/run/rtty-go.pid"
depend() {
need net
after firewall
}
start_pre() {
. ${CONFIG_FILE}
}
EOF
chmod +x /etc/init.d/rtty-go
rc-update add rtty-go default
rc-service rtty-go restart
log_info "OpenRC service started."
}
setup_crontab() {
log_info "Setting up crontab auto-start (fallback)..."
# Create a wrapper script
cat > "${CONFIG_DIR}/start.sh" <<EOF
#!/bin/sh
. ${CONFIG_FILE}
if ! pgrep -f "${INSTALL_DIR}/${BINARY_NAME}" >/dev/null 2>&1; then
${INSTALL_DIR}/${BINARY_NAME} ${RTTY_SSL} -a -I "\${DEVICE_ID}" -h "\${RTTY_HOST}" -p "\${RTTY_PORT}" -t "\${RTTY_TOKEN}" -d "\${DEVICE_MAC}" &
fi
EOF
chmod +x "${CONFIG_DIR}/start.sh"
# Add to crontab: check every minute + run at reboot
(crontab -l 2>/dev/null | grep -v "rtty-go"; \
echo "@reboot sleep 10 && ${CONFIG_DIR}/start.sh"; \
echo "*/5 * * * * ${CONFIG_DIR}/start.sh") | crontab -
# Start now
"${CONFIG_DIR}/start.sh"
log_info "Crontab auto-start configured."
}
# ========================== Main =============================================
# Handle uninstall
[ "$UNINSTALL" -eq 1 ] && do_uninstall
# Validate required parameters
[ -z "$RTTY_HOST" ] && log_error "Server host is required (-h)" && usage
[ -z "$RTTY_TOKEN" ] && log_error "Token is required (-t)" && usage
# Check root
if [ "$(id -u)" -ne 0 ]; then
log_error "This script must be run as root (use sudo)"
exit 1
fi
PLATFORM=$(detect_arch)
INIT_SYSTEM=$(detect_init)
# Reuse existing device ID and MAC if config exists (preserve identity across reinstalls)
# Only extract DEVICE_ID and DEVICE_MAC, do NOT source the whole file
# (sourcing would overwrite RTTY_HOST/RTTY_PORT/RTTY_TOKEN from CLI args)
EXISTING_ID=""
EXISTING_MAC=""
if [ -f "$CONFIG_FILE" ]; then
EXISTING_ID=$(grep '^DEVICE_ID=' "$CONFIG_FILE" | head -1 | cut -d'=' -f2- | tr -d '"')
EXISTING_MAC=$(grep '^DEVICE_MAC=' "$CONFIG_FILE" | head -1 | cut -d'=' -f2- | tr -d '"')
fi
DEVICE_ID="${EXISTING_ID:-$(gen_device_id)}"
DEVICE_MAC="${EXISTING_MAC:-$(get_mac)}"
log_info "========================================"
log_info " rtty-go Client Installer for Linux"
log_info "========================================"
log_info "Platform: ${PLATFORM}"
log_info "Init system: ${INIT_SYSTEM}"
log_info "Device ID: ${DEVICE_ID}"
log_info "MAC address: ${DEVICE_MAC}"
log_info "Server: ${RTTY_HOST}:${RTTY_PORT}"
log_info "========================================"
# Determine download URL
if [ -z "$DOWNLOAD_BASE_URL" ]; then
DOWNLOAD_BASE_URL="https://kvm-cloud.gl-inet.com/selfhost/clients"
fi
FILE_URL="${DOWNLOAD_BASE_URL}/${BINARY_NAME}-${PLATFORM}"
# Stop existing service before replacing binary and config
case "$INIT_SYSTEM" in
systemd) systemctl stop rtty-go 2>/dev/null || true ;;
procd) /etc/init.d/rtty-go stop 2>/dev/null || true ;;
openrc) rc-service rtty-go stop 2>/dev/null || true ;;
*) pkill -f "${INSTALL_DIR}/${BINARY_NAME}" 2>/dev/null || true ;;
esac
# Download binary
log_info "Downloading ${BINARY_NAME}-${PLATFORM}..."
download_file "$FILE_URL" "/tmp/${BINARY_NAME}"
chmod +x "/tmp/${BINARY_NAME}"
# Install binary
mkdir -p "$INSTALL_DIR"
mv -f "/tmp/${BINARY_NAME}" "${INSTALL_DIR}/${BINARY_NAME}"
log_info "Installed to ${INSTALL_DIR}/${BINARY_NAME}"
# Save configuration
mkdir -p "$CONFIG_DIR"
cat > "$CONFIG_FILE" <<EOF
RTTY_HOST="${RTTY_HOST}"
RTTY_PORT="${RTTY_PORT}"
RTTY_TOKEN="${RTTY_TOKEN}"
DEVICE_ID="${DEVICE_ID}"
DEVICE_MAC="${DEVICE_MAC}"
EOF
chmod 600 "$CONFIG_FILE"
log_info "Configuration saved to ${CONFIG_FILE}"
# Setup auto-start service
case "$INIT_SYSTEM" in
systemd) setup_systemd ;;
procd) setup_procd ;;
openrc) setup_openrc ;;
*) setup_crontab ;;
esac
log_info "========================================"
log_info " Installation complete!"
log_info " Device ID: ${DEVICE_ID}"
log_info " MAC address: ${DEVICE_MAC}"
log_info "========================================"
+253
View File
@@ -0,0 +1,253 @@
#!/bin/bash
# ============================================================================
# rtty-go client one-click installer for macOS
# Supports: Intel (amd64) / Apple Silicon (arm64)
# ============================================================================
set -e
# --------------- Configuration (can be pre-filled by server) ----------------
RTTY_HOST=""
RTTY_PORT="5912"
RTTY_TOKEN=""
RTTY_SSL="-s -x"
DOWNLOAD_BASE_URL="https://kvm-cloud.gl-inet.com/selfhost/clients"
# ----------------------------------------------------------------------------
INSTALL_DIR="/usr/local/bin"
BINARY_NAME="rtty-go"
CONFIG_DIR="/etc/rtty-go"
CONFIG_FILE="${CONFIG_DIR}/env"
PLIST_NAME="com.glkvm.rtty-go"
PLIST_PATH="/Library/LaunchDaemons/${PLIST_NAME}.plist"
# ========================== Helper Functions =================================
log_info() { printf "\033[32m[INFO]\033[0m %s\n" "$1"; }
log_warn() { printf "\033[33m[WARN]\033[0m %s\n" "$1"; }
log_error() { printf "\033[31m[ERROR]\033[0m %s\n" "$1"; }
usage() {
cat <<EOF
Usage: $0 -h <host> -t <token> [-p <port>] [-u <download_base_url>]
Options:
-h Server host or IP address (required)
-t Authorization token (required)
-p Server port (default: 5912)
-u Download base URL for binaries
--uninstall Uninstall rtty-go client
Example:
sudo $0 -h 107.173.152.173 -t lHEP7GyyGt4S18KlyikfpvzdZTVxnD8v
EOF
exit 1
}
# ========================== Parse Arguments ==================================
UNINSTALL=0
while [[ $# -gt 0 ]]; do
case "$1" in
-h) RTTY_HOST="$2"; shift 2 ;;
-p) RTTY_PORT="$2"; shift 2 ;;
-t) RTTY_TOKEN="$2"; shift 2 ;;
-u) DOWNLOAD_BASE_URL="$2"; shift 2 ;;
--uninstall) UNINSTALL=1; shift ;;
*) usage ;;
esac
done
# Strip leading colon from port (e.g. ":5912" -> "5912")
RTTY_PORT="${RTTY_PORT#:}"
# ========================== Detect Architecture ==============================
detect_arch() {
ARCH=$(uname -m)
case "$ARCH" in
x86_64) echo "darwin-amd64" ;;
arm64) echo "darwin-arm64" ;;
*) log_error "Unsupported architecture: $ARCH"; exit 1 ;;
esac
}
# ========================== Get MAC Address ==================================
get_mac() {
MAC=""
# Method 1: en0 (usually the primary interface)
if command -v ifconfig &>/dev/null; then
MAC=$(ifconfig en0 2>/dev/null | awk '/ether/{print $2}')
fi
# Method 2: networksetup
if [[ -z "$MAC" ]] && command -v networksetup &>/dev/null; then
MAC=$(networksetup -getmacaddress Wi-Fi 2>/dev/null | awk '{print $3}')
if [[ "$MAC" == *"not"* ]] || [[ -z "$MAC" ]]; then
MAC=$(networksetup -getmacaddress Ethernet 2>/dev/null | awk '{print $3}')
fi
fi
# Method 3: Generate random locally-administered MAC
if [[ -z "$MAC" ]] || [[ "$MAC" == "00:00:00:00:00:00" ]]; then
log_warn "Could not detect MAC address, generating random one"
MAC=$(printf '02:%02x:%02x:%02x:%02x:%02x' $((RANDOM%256)) $((RANDOM%256)) $((RANDOM%256)) $((RANDOM%256)) $((RANDOM%256)))
fi
echo "$MAC"
}
# ========================== Generate Device ID ===============================
gen_device_id() {
uuidgen | tr -d '-' | tr 'A-F' 'a-f' | cut -c1-8
}
# ========================== Uninstall ========================================
do_uninstall() {
log_info "Uninstalling rtty-go client..."
# Stop and remove LaunchDaemon
if [[ -f "$PLIST_PATH" ]]; then
launchctl bootout system "$PLIST_PATH" 2>/dev/null || true
rm -f "$PLIST_PATH"
fi
rm -f "${INSTALL_DIR}/${BINARY_NAME}"
rm -rf "${CONFIG_DIR}"
log_info "Uninstall complete."
exit 0
}
# ========================== Main =============================================
# Handle uninstall
[[ "$UNINSTALL" -eq 1 ]] && do_uninstall
# Validate required parameters
[[ -z "$RTTY_HOST" ]] && log_error "Server host is required (-h)" && usage
[[ -z "$RTTY_TOKEN" ]] && log_error "Token is required (-t)" && usage
# Check root
if [[ $(id -u) -ne 0 ]]; then
log_error "This script must be run as root (use sudo)"
exit 1
fi
PLATFORM=$(detect_arch)
# Reuse existing device ID and MAC if config exists (preserve identity across reinstalls)
# Only extract DEVICE_ID and DEVICE_MAC, do NOT source the whole file
# (sourcing would overwrite RTTY_HOST/RTTY_PORT/RTTY_TOKEN from CLI args)
EXISTING_ID=""
EXISTING_MAC=""
if [[ -f "$CONFIG_FILE" ]]; then
EXISTING_ID=$(grep '^DEVICE_ID=' "$CONFIG_FILE" | head -1 | cut -d'=' -f2- | tr -d '"')
EXISTING_MAC=$(grep '^DEVICE_MAC=' "$CONFIG_FILE" | head -1 | cut -d'=' -f2- | tr -d '"')
fi
DEVICE_ID="${EXISTING_ID:-$(gen_device_id)}"
DEVICE_MAC="${EXISTING_MAC:-$(get_mac)}"
log_info "========================================"
log_info " rtty-go Client Installer for macOS"
log_info "========================================"
log_info "Platform: ${PLATFORM}"
log_info "Device ID: ${DEVICE_ID}"
log_info "MAC address: ${DEVICE_MAC}"
log_info "Server: ${RTTY_HOST}:${RTTY_PORT}"
log_info "========================================"
# Determine download URL
if [[ -z "$DOWNLOAD_BASE_URL" ]]; then
DOWNLOAD_BASE_URL="https://kvm-cloud.gl-inet.com/selfhost/clients"
fi
FILE_URL="${DOWNLOAD_BASE_URL}/${BINARY_NAME}-${PLATFORM}"
# Stop existing service before replacing binary and config
if [[ -f "$PLIST_PATH" ]]; then
launchctl bootout system "$PLIST_PATH" 2>/dev/null || true
fi
# Download binary
log_info "Downloading ${BINARY_NAME}-${PLATFORM}..."
curl -fsSL -o "/tmp/${BINARY_NAME}" "$FILE_URL"
chmod +x "/tmp/${BINARY_NAME}"
# Remove quarantine attribute (macOS Gatekeeper)
xattr -d com.apple.quarantine "/tmp/${BINARY_NAME}" 2>/dev/null || true
# Install binary
mkdir -p "$INSTALL_DIR"
mv -f "/tmp/${BINARY_NAME}" "${INSTALL_DIR}/${BINARY_NAME}"
log_info "Installed to ${INSTALL_DIR}/${BINARY_NAME}"
# Save configuration
mkdir -p "$CONFIG_DIR"
cat > "$CONFIG_FILE" <<EOF
RTTY_HOST="${RTTY_HOST}"
RTTY_PORT="${RTTY_PORT}"
RTTY_TOKEN="${RTTY_TOKEN}"
DEVICE_ID="${DEVICE_ID}"
DEVICE_MAC="${DEVICE_MAC}"
EOF
chmod 600 "$CONFIG_FILE"
log_info "Configuration saved to ${CONFIG_FILE}"
# Create LaunchDaemon for auto-start on boot
log_info "Setting up LaunchDaemon for auto-start..."
cat > "$PLIST_PATH" <<EOF
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">
<plist version="1.0">
<dict>
<key>Label</key>
<string>${PLIST_NAME}</string>
<key>ProgramArguments</key>
<array>
<string>${INSTALL_DIR}/${BINARY_NAME}</string>
<string>-s</string>
<string>-x</string>
<string>-a</string>
<string>-I</string>
<string>${DEVICE_ID}</string>
<string>-h</string>
<string>${RTTY_HOST}</string>
<string>-p</string>
<string>${RTTY_PORT}</string>
<string>-t</string>
<string>${RTTY_TOKEN}</string>
<string>-d</string>
<string>${DEVICE_MAC}</string>
</array>
<key>RunAtLoad</key>
<true/>
<key>KeepAlive</key>
<true/>
<key>StandardOutPath</key>
<string>/var/log/rtty-go.log</string>
<key>StandardErrorPath</key>
<string>/var/log/rtty-go.log</string>
</dict>
</plist>
EOF
chmod 644 "$PLIST_PATH"
launchctl bootstrap system "$PLIST_PATH"
log_info "========================================"
log_info " Installation complete!"
log_info " Device ID: ${DEVICE_ID}"
log_info " MAC address: ${DEVICE_MAC}"
log_info ""
log_info " Manage service:"
log_info " Stop: sudo launchctl bootout system ${PLIST_PATH}"
log_info " Start: sudo launchctl bootstrap system ${PLIST_PATH}"
log_info " Log: tail -f /var/log/rtty-go.log"
log_info " Remove: sudo $0 --uninstall"
log_info "========================================"

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