Compare commits
53 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e6da02f711 | |||
| 75a0416221 | |||
| dc15e534b8 | |||
| ea96caef90 | |||
| 13dca8b190 | |||
| 3810f9e72c | |||
| a8414dd025 | |||
| cea6b19f9a | |||
| ba07f6386b | |||
| 206b0fa68b | |||
| e3857e9f61 | |||
| f3c32b186c | |||
| 769b298707 | |||
| 3a503d0e77 | |||
| 929e9526d9 | |||
| 82a8037244 | |||
| 9c56f78493 | |||
| a23912fd89 | |||
| dff4d06107 | |||
| 359f29484a | |||
| f1fb5b52fb | |||
| c1b987e6b3 | |||
| 1320b58cbb | |||
| a1ebef3326 | |||
| fdba252fa8 | |||
| 9faeac6ff9 | |||
| 3e8d898a67 | |||
| b2a9c6c7dd | |||
| 9c4592e697 | |||
| f783fa14fe | |||
| 6a01778bec | |||
| 3b5360478d | |||
| 39dc768b9c | |||
| 606377fbf8 | |||
| b331a7c7f3 | |||
| 2b6a6eff77 | |||
| d93e5953ad | |||
| 0d130ebc29 | |||
| 2d2658383b | |||
| 7931fa0efd | |||
| ab77f53a77 | |||
| a5bce3b289 | |||
| 6c357f2842 | |||
| 8703f91ebf | |||
| 1c473458cf | |||
| c6c09dbae1 | |||
| 9258aa5f43 | |||
| 8994cccb25 | |||
| 60994b9513 | |||
| c99b96ca01 | |||
| 27792291ed | |||
| 4adb10f577 | |||
| cd231e996b |
@@ -12,3 +12,5 @@
|
||||
|
||||
# Project-local glide cache, RE: https://github.com/Masterminds/glide/issues/736
|
||||
.glide/
|
||||
|
||||
.vscode
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
{
|
||||
"cSpell.words": [
|
||||
"ddns",
|
||||
"glkvm",
|
||||
"repassword",
|
||||
"webrtc"
|
||||
]
|
||||
}
|
||||
@@ -1,4 +1,7 @@
|
||||
FROM alpine:latest
|
||||
WORKDIR /home
|
||||
COPY ./rttys /usr/bin/rttys
|
||||
|
||||
ARG TARGETARCH
|
||||
COPY ./dist/rttys-linux-${TARGETARCH} /usr/bin/rttys
|
||||
|
||||
ENTRYPOINT ["/usr/bin/rttys"]
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -1,45 +1,74 @@
|
||||
# Makefile
|
||||
|
||||
# Go binary name
|
||||
BINARY_NAME = rttys
|
||||
# ---------------- Project ----------------
|
||||
BINARY_NAME ?= rttys
|
||||
UI_DIR ?= ui
|
||||
GO_MAIN ?= ./cmd/glkvm-cloud
|
||||
|
||||
# Go build flags
|
||||
BUILD_FLAGS := -ldflags "-s -w"
|
||||
BUILD_FLAGS ?= -ldflags "-s -w"
|
||||
DIST_DIR ?= dist
|
||||
|
||||
# Go build command
|
||||
GO_BUILD_CMD = go build $(BUILD_FLAGS) -o $(BINARY_NAME)
|
||||
# Image name
|
||||
IMAGE_NAME ?= glkvm-cloud
|
||||
IMAGE_TAG ?= build
|
||||
|
||||
# Paths
|
||||
UI_DIR = ui
|
||||
CONF_FILE = ./rttys.conf
|
||||
GOARCH ?= $(shell go env GOARCH)
|
||||
|
||||
.PHONY: all ui build run build-run full-run
|
||||
# ---------------- Commands ----------------
|
||||
.PHONY: all ui debug-local \
|
||||
build-linux-amd64 build-linux-arm64 build-linux-all \
|
||||
docker-buildx docker-buildx-full
|
||||
|
||||
all: build-linux-amd64 build-linux-arm64
|
||||
|
||||
# Build frontend files only
|
||||
ui:
|
||||
cd $(UI_DIR) && npm install && npm run build
|
||||
|
||||
# Build Go binary only
|
||||
build:
|
||||
CGO_ENABLED=0 $(GO_BUILD_CMD)
|
||||
# ---------------- 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_MAIN)
|
||||
|
||||
# Run Go program only
|
||||
run:
|
||||
./$(BINARY_NAME) -c $(CONF_FILE)
|
||||
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 $(GO_MAIN)
|
||||
|
||||
# Build frontend and Go binary
|
||||
build-all: ui build
|
||||
# ---------------- Docker Buildx ----------------
|
||||
# Multi-arch build
|
||||
# Usage:
|
||||
# make docker-buildx GOARCH=amd64 IMAGE_TAG=build-amd64
|
||||
# make docker-buildx GOARCH=arm64 IMAGE_TAG=build-arm64
|
||||
REGISTRY ?=
|
||||
|
||||
# Build Go binary and run
|
||||
build-run: build run
|
||||
# If REGISTRY is set, tag becomes: REGISTRY/IMAGE_NAME:IMAGE_TAG
|
||||
ifdef REGISTRY
|
||||
IMAGE_REF := $(REGISTRY)/$(IMAGE_NAME):$(IMAGE_TAG)
|
||||
else
|
||||
IMAGE_REF := $(IMAGE_NAME):$(IMAGE_TAG)
|
||||
endif
|
||||
|
||||
# Build frontend, build Go binary, and run
|
||||
full-run: ui build run
|
||||
docker-buildx:
|
||||
@docker buildx version >/dev/null 2>&1 || (echo "docker buildx not available" && exit 1)
|
||||
@echo "==> buildx (load local image): $(IMAGE_REF) [linux/$(GOARCH)]"
|
||||
docker buildx build \
|
||||
--platform linux/$(GOARCH) \
|
||||
-t $(IMAGE_REF) \
|
||||
--load .
|
||||
|
||||
# Build Docker image without updating ui
|
||||
docker-build: build
|
||||
docker build -t glkvm-cloud:build .
|
||||
docker-buildx-full: ui
|
||||
@$(MAKE) docker-buildx
|
||||
|
||||
# Full Build Docker image
|
||||
docker-fullbuild: ui build
|
||||
docker build -t glkvm-cloud:build .
|
||||
|
||||
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"
|
||||
|
||||
@@ -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,6 +17,9 @@ 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 & 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
|
||||
|
||||
The following mainstream operating systems have been tested and verified
|
||||
@@ -61,9 +64,11 @@ If your server provider uses a **cloud security group** (e.g., AWS, Aliyun, etc.
|
||||
|
||||
We provide **two** ways to install GLKVM Cloud:
|
||||
|
||||
#### A) One-line installer (recommended)
|
||||
#### A) One-line installer (recommended, x86_64/amd64)
|
||||
|
||||
> **Note:** The one-line installer is **Docker-based**. It automates Docker/Compose setup, pulls images, renders configs from templates, and starts services for you.
|
||||
>
|
||||
> **Platform:** currently supports **x86_64 (amd64)** only.
|
||||
|
||||
Run **as root**:
|
||||
|
||||
@@ -74,24 +79,27 @@ Run **as root**:
|
||||
#### B) Docker manual install
|
||||
|
||||
> Full reference: see [`docker-compose/README.md`](https://github.com/gl-inet/glkvm-cloud/blob/main/docker-compose/README.md)
|
||||
>
|
||||
> **Platform:** supports both **x86_64 (amd64)** and **arm64 (AArch64)**.
|
||||
|
||||
### 🌐 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>
|
||||
```
|
||||
|
||||

|
||||
@@ -121,6 +129,10 @@ The default login password for the Web UI will be displayed in the installation
|
||||
|
||||

|
||||
|
||||
### Web Proxy
|
||||
|
||||

|
||||
|
||||
|
||||
|
||||
## Use your own SSL Certificate (Optional)
|
||||
@@ -227,4 +239,4 @@ Once everything is configured, you can access the platform via your domain:
|
||||
|
||||
```
|
||||
https://www.your-domain.com
|
||||
```
|
||||
```
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
|
||||
#### 主要功能与特性
|
||||
|
||||
* **设备管理** - 实时查看设备在线状态
|
||||
* **用户组和设备支持** - 支持用户组管理特定的设备组设备,实现不同用户管理不同设备
|
||||
* **脚本部署** - 通过脚本快速添加设备
|
||||
* **远程 SSH** - Web SSH 远程连接
|
||||
* **远程控制** - Web远程桌面控制
|
||||
@@ -18,6 +18,9 @@
|
||||
* **轻量设计** - 专为小型企业和个人优化
|
||||
* **企业级认证** - 同时支持 **LDAP** 和 **OIDC** 登录方式,适用于企业用户。
|
||||
|
||||
- **部署与平台兼容性** - 同时支持 **内网部署** 和 **公网部署**,并兼容 **x86_64** 与 **arm64** 平台
|
||||
- **HTTP/HTTPS Web代理功能支持** - 支持 OpenWrt、ImmortalWrt、树莓派、Linux VPS、macOS、Windows 等主机接入自部署 GLKVM Cloud 进行统一管理,并可作为 HTTP/HTTPS Web 代理节点实现内网穿透访问
|
||||
|
||||
## 自部署指南
|
||||
|
||||
以下主流操作系统已通过测试验证:
|
||||
@@ -61,9 +64,11 @@
|
||||
|
||||
我们提供 **两种** 安装 GLKVM Cloud 的方式:
|
||||
|
||||
#### A) 一键安装脚本(推荐)
|
||||
#### A) 一键安装脚本(推荐,仅支持 x86_64 / amd64)
|
||||
|
||||
> **注意:** 一键安装脚本基于 **Docker**。它会自动完成 Docker / Docker Compose 的安装、拉取镜像、根据模板渲染配置文件,并启动所有服务。
|
||||
>
|
||||
> **平台支持:** 当前仅支持 **x86_64(amd64)** 平台。
|
||||
|
||||
使用 **root 权限** 运行以下命令安装 GLKVM 轻量云:
|
||||
|
||||
@@ -74,25 +79,28 @@
|
||||
#### B) 使用 Docker 手动安装
|
||||
|
||||
> 完整参考文档请查看:[`docker-compose/README-CN.md`](https://github.com/gl-inet/glkvm-cloud/blob/main/docker-compose/README-CN.md)
|
||||
>
|
||||
> 平台支持: 同时支持 x86_64(amd64) 与 arm64(AArch64) 平台。
|
||||
|
||||
|
||||
### 🌐 平台访问
|
||||
|
||||
安装完成后,你可以通过以下方式访问平台:
|
||||
安装完成后,安装脚本会在控制台输出平台访问地址和管理员登录信息。你可以通过以下方式访问平台:
|
||||
|
||||
```
|
||||
https://<你的服务器公网IP>
|
||||
```
|
||||
|
||||
⚠️ **提示**:通过 IP 访问时,浏览器会提示 **证书不受信任**。
|
||||
如果想消除该提示,建议配置 **自定义域名 + 有效 SSL 证书**。
|
||||
如需消除该提示,建议配置 **自定义域名 + 有效 SSL 证书**。
|
||||
|
||||
### 🔑 Web UI 登录密码
|
||||
### 🔑 Web UI 登录信息
|
||||
|
||||
Web UI 的默认登录密码会在安装脚本运行结束时显示:
|
||||
安装脚本运行结束后,安装控制台会显示 Web UI 管理员用户名和密码(示例):
|
||||
|
||||
```
|
||||
🔐 请在安装控制台查看 Web 登录密码
|
||||
```text
|
||||
👤 管理员用户名:admin
|
||||
🔑 管理员密码:<自动生成密码>
|
||||
```
|
||||
|
||||

|
||||
@@ -121,6 +129,10 @@ Web UI 的默认登录密码会在安装脚本运行结束时显示:
|
||||
|
||||

|
||||
|
||||
#### Web代理功能
|
||||
|
||||

|
||||
|
||||
## 使用自有 SSL 证书(可选)
|
||||
|
||||
⚠️ **可选配置**:
|
||||
|
||||
@@ -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 (
|
||||
"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) {
|
||||
c.Next()
|
||||
log.Debug().Msgf("%s - \"%s %s %s %d\"", c.ClientIP(),
|
||||
c.Request.Method, c.Request.URL.Path, c.Request.Proto, c.Writer.Status())
|
||||
})
|
||||
|
||||
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,
|
||||
}
|
||||
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
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
# Images
|
||||
GLKVM_IMAGE=glzhitong/glkvm-cloud:latest-arm64
|
||||
COTURN_IMAGE=coturn/coturn:edge-alpine-arm64v8
|
||||
|
||||
# Enable reverse proxy mode (e.g. Nginx in front of GLKVM Cloud).
|
||||
# When enabled, TLS is handled by the proxy and GLKVM Cloud runs in plain HTTP.
|
||||
#
|
||||
# Note:
|
||||
# In reverse-proxy mode, remote device access depends on the correct forwarded headers
|
||||
# from the front-end proxy. If these headers are missing or incorrect, GLKVM Cloud may
|
||||
# generate redirect URLs with the internal port (e.g. :10443).
|
||||
#
|
||||
# Please make sure your Nginx config includes:
|
||||
# proxy_set_header Host $host;
|
||||
# proxy_set_header X-Forwarded-Host $host;
|
||||
# proxy_set_header X-Forwarded-Proto $scheme;
|
||||
# proxy_set_header X-Forwarded-Port $server_port;
|
||||
# proxy_set_header X-Real-IP $remote_addr;
|
||||
# proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
#
|
||||
# Reference (verified working example):
|
||||
# https://github.com/gl-inet/glkvm-cloud/blob/main/docker-compose/nginx-reverse-proxy-example.conf
|
||||
REVERSE_PROXY_ENABLED=false
|
||||
|
||||
# =====================================================
|
||||
# Device Remote Access Domain (Reverse Proxy Mode Only)
|
||||
# =====================================================
|
||||
# This option is used to generate the Remote Control URL for devices when
|
||||
# running behind a reverse proxy.
|
||||
#
|
||||
# Effective ONLY when:
|
||||
# REVERSE_PROXY_ENABLED=true
|
||||
#
|
||||
# When set, GLKVM Cloud will generate device access addresses as:
|
||||
# https://<deviceId>.<DEVICE_ENDPOINT_HOST>/... (scheme is taken from X-Forwarded-Proto)
|
||||
#
|
||||
# Examples:
|
||||
# DEVICE_ENDPOINT_HOST=kvm.example.com
|
||||
# DEVICE_ENDPOINT_HOST=kvm.example.com:443
|
||||
#
|
||||
# Notes:
|
||||
# - Do NOT include scheme (http:// or https://)
|
||||
# - Do NOT include path (/xxx)
|
||||
#
|
||||
# Leave empty to derive the host/port from X-Forwarded-* headers (auto-detect).
|
||||
DEVICE_ENDPOINT_HOST=
|
||||
|
||||
|
||||
# =====================================================
|
||||
# Platform Access Domain Restriction
|
||||
# =====================================================
|
||||
# Restrict the domain used to access the GLKVM Cloud platform.
|
||||
#
|
||||
# When set, only requests with a matching domain are allowed to access
|
||||
# the Web UI and API. Requests using other domains will be rejected
|
||||
# as invalid access.
|
||||
#
|
||||
# Examples:
|
||||
# WEB_UI_HOST=www.example.com
|
||||
#
|
||||
# Notes:
|
||||
# - Do NOT include scheme (http:// or https://)
|
||||
# - Do NOT include path (/xxx)
|
||||
# - Leave empty to disable domain restriction (allow access via any domain)
|
||||
WEB_UI_HOST=
|
||||
|
||||
# GLKVM access IP seen by devices/users.
|
||||
# Leave empty to auto-detect at container start.
|
||||
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
|
||||
|
||||
# TURN
|
||||
TURN_PORT=3478
|
||||
TURN_USER=glkvmcloudwebrtcuser
|
||||
TURN_PASS=AnotherS3cret
|
||||
|
||||
# LDAP Authentication (Optional)
|
||||
LDAP_ENABLED=false
|
||||
LDAP_SERVER=your-ldap-server.com
|
||||
LDAP_PORT=389
|
||||
LDAP_USE_TLS=false
|
||||
LDAP_BIND_DN=cn=service-account,ou=users,dc=company,dc=com
|
||||
LDAP_BIND_PASSWORD=service-password
|
||||
LDAP_BASE_DN=ou=users,dc=company,dc=com
|
||||
|
||||
# User filter examples for different LDAP implementations:
|
||||
# Active Directory: (&(objectClass=person)(sAMAccountName=%s))
|
||||
# OpenLDAP: (&(objectClass=inetOrgPerson)(uid=%s))
|
||||
# FreeIPA: (&(objectClass=person)(uid=%s))
|
||||
# Generic LDAP: (uid=%s)
|
||||
LDAP_USER_FILTER=(uid=%s)
|
||||
|
||||
LDAP_ALLOWED_GROUPS=admins,operators
|
||||
LDAP_ALLOWED_USERS=user1,user2
|
||||
|
||||
# OIDC Authentication (Optional, generic OIDC provider)
|
||||
OIDC_ENABLED=false
|
||||
OIDC_ISSUER=
|
||||
OIDC_CLIENT_ID=
|
||||
OIDC_CLIENT_SECRET=
|
||||
OIDC_AUTH_URL=
|
||||
OIDC_TOKEN_URL=
|
||||
|
||||
# Redirect URL registered in your OIDC provider.
|
||||
# The path part (/auth/oidc/callback) is fixed by GLKVM Cloud and must not be changed.
|
||||
# Example:
|
||||
# OIDC_REDIRECT_URL=https://your-domain.example.com/auth/oidc/callback
|
||||
OIDC_REDIRECT_URL=
|
||||
|
||||
OIDC_SCOPES="openid profile email"
|
||||
|
||||
# Email-based whitelist (exact email or domain like @example.com)
|
||||
OIDC_ALLOWED_USERS=
|
||||
# Subject (sub) whitelist (stable user IDs)
|
||||
OIDC_ALLOWED_SUBS=
|
||||
# Username whitelist (preferred_username or name)
|
||||
OIDC_ALLOWED_USERNAMES=
|
||||
# Groups whitelist (e.g. admin, devops)
|
||||
OIDC_ALLOWED_GROUPS=
|
||||
@@ -4,15 +4,75 @@ COTURN_IMAGE=coturn/coturn:edge-alpine
|
||||
|
||||
# Enable reverse proxy mode (e.g. Nginx in front of GLKVM Cloud).
|
||||
# When enabled, TLS is handled by the proxy and GLKVM Cloud runs in plain HTTP.
|
||||
#
|
||||
# Note:
|
||||
# In reverse-proxy mode, remote device access depends on the correct forwarded headers
|
||||
# from the front-end proxy. If these headers are missing or incorrect, GLKVM Cloud may
|
||||
# generate redirect URLs with the internal port (e.g. :10443).
|
||||
#
|
||||
# Please make sure your Nginx config includes:
|
||||
# proxy_set_header Host $host;
|
||||
# proxy_set_header X-Forwarded-Host $host;
|
||||
# proxy_set_header X-Forwarded-Proto $scheme;
|
||||
# proxy_set_header X-Forwarded-Port $server_port;
|
||||
# proxy_set_header X-Real-IP $remote_addr;
|
||||
# proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
#
|
||||
# Reference (verified working example):
|
||||
# https://github.com/gl-inet/glkvm-cloud/blob/main/docker-compose/nginx-reverse-proxy-example.conf
|
||||
REVERSE_PROXY_ENABLED=false
|
||||
|
||||
# GLKVM access IP seen by devices/users.
|
||||
# =====================================================
|
||||
# Device Remote Access Domain (Reverse Proxy Mode Only)
|
||||
# =====================================================
|
||||
# This option is used to generate the Remote Control URL for devices when
|
||||
# running behind a reverse proxy.
|
||||
#
|
||||
# Effective ONLY when:
|
||||
# REVERSE_PROXY_ENABLED=true
|
||||
#
|
||||
# When set, GLKVM Cloud will generate device access addresses as:
|
||||
# https://<deviceId>.<DEVICE_ENDPOINT_HOST>/... (scheme is taken from X-Forwarded-Proto)
|
||||
#
|
||||
# Examples:
|
||||
# DEVICE_ENDPOINT_HOST=kvm.example.com
|
||||
# DEVICE_ENDPOINT_HOST=kvm.example.com:443
|
||||
#
|
||||
# Notes:
|
||||
# - Do NOT include scheme (http:// or https://)
|
||||
# - Do NOT include path (/xxx)
|
||||
#
|
||||
# Leave empty to derive the host/port from X-Forwarded-* headers (auto-detect).
|
||||
DEVICE_ENDPOINT_HOST=
|
||||
|
||||
# =====================================================
|
||||
# Platform Access Domain Restriction
|
||||
# =====================================================
|
||||
# Restrict the domain used to access the GLKVM Cloud platform.
|
||||
#
|
||||
# When set, only requests with a matching domain are allowed to access
|
||||
# the Web UI and API. Requests using other domains will be rejected
|
||||
# as invalid access.
|
||||
#
|
||||
# Examples:
|
||||
# WEB_UI_HOST=www.example.com
|
||||
#
|
||||
# Notes:
|
||||
# - Do NOT include scheme (http:// or https://)
|
||||
# - Do NOT include path (/xxx)
|
||||
# - Leave empty to disable domain restriction (allow access via any domain)
|
||||
WEB_UI_HOST=
|
||||
|
||||
GLKVM access IP seen by devices/users.
|
||||
# Leave empty to auto-detect at container start.
|
||||
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
|
||||
|
||||
@@ -7,9 +7,19 @@
|
||||
```bash
|
||||
git clone https://github.com/gl-inet/glkvm-cloud.git
|
||||
cd glkvm-cloud/docker-compose/
|
||||
cp .env.example .env
|
||||
```
|
||||
|
||||
* **x86_64(amd64)平台**:
|
||||
```bash
|
||||
cp .env.example .env
|
||||
```
|
||||
* **arm64(AArch64)平台**:
|
||||
|
||||
```bash
|
||||
cp .env.arm64.example .env
|
||||
```
|
||||
|
||||
|
||||
### 2. **配置环境变量**
|
||||
|
||||
编辑 `.env` 文件,并根据需求更新关键参数:
|
||||
@@ -53,27 +63,24 @@ cp .env.example .env
|
||||
- `OIDC_ALLOWED_USERNAMES`:允许的用户名列表(可选)
|
||||
- `OIDC_ALLOWED_GROUPS`:允许的用户组列表(可选)
|
||||
|
||||
#### **反向代理模式(可选)**
|
||||
#### 反向代理模式(可选)
|
||||
|
||||
```env
|
||||
# 启用反向代理模式(例如在 GLKVM Cloud 前使用 Nginx)
|
||||
# 启用后,TLS 由反向代理终止,GLKVM Cloud 内部使用明文 HTTP
|
||||
REVERSE_PROXY_ENABLED=false
|
||||
```
|
||||
|
||||
当 `REVERSE_PROXY_ENABLED` 设置为 `true` 时,GLKVM Cloud 将运行在 **反向代理(如 Nginx)之后**:
|
||||
启用后(`REVERSE_PROXY_ENABLED=true`):
|
||||
|
||||
- HTTPS 证书由反向代理管理(而不是由 GLKVM Cloud 本身管理)
|
||||
- GLKVM Cloud 内部以明文 HTTP 方式监听
|
||||
- 同一个 HTTPS 端口可同时用于:
|
||||
- 访问 GLKVM Cloud Web 管理界面
|
||||
- 访问远程 KVM 设备
|
||||
- GLKVM Cloud 运行在反向代理(如 Nginx)之后
|
||||
- TLS 由反向代理终止,GLKVM Cloud 内部使用 HTTP
|
||||
- Web UI 与设备远程访问可共用同一个 HTTPS 端口(通常为 443)
|
||||
|
||||
例如,在正确配置 Nginx 的情况下:
|
||||
|
||||
##### 必需的反向代理请求头
|
||||
|
||||
反向代理必须转发以下请求头,否则可能生成包含内部端口(如 `:10443`)的访问地址:
|
||||
|
||||
```nginx
|
||||
# 转发原始的主机名、协议、端口以及客户端 IP
|
||||
# 在反向代理模式下,这些 Header 是必须的
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Forwarded-Host $host;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
@@ -82,14 +89,34 @@ proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
```
|
||||
|
||||
你可以通过以下地址访问:
|
||||
##### 设备远程访问域名(可选)
|
||||
|
||||
```text
|
||||
https://www.example.com → GLKVM Cloud 管理界面
|
||||
https://<device_id>.example.com → 远程设备访问
|
||||
```env
|
||||
DEVICE_ENDPOINT_HOST=
|
||||
```
|
||||
|
||||
- **仅在** `REVERSE_PROXY_ENABLED=true` 时生效
|
||||
- 用于指定设备远程访问使用的域名
|
||||
- 生成的设备访问地址格式为:
|
||||
|
||||
```text
|
||||
https://<deviceId>.<DEVICE_ENDPOINT_HOST>/
|
||||
```
|
||||
|
||||
**说明:**
|
||||
|
||||
- 不需要包含 `http(s)://` 或路径
|
||||
- 可与 Web UI 域名不同
|
||||
- 留空时,将从 `X-Forwarded-*` 请求头自动推导
|
||||
|
||||
**示例:**
|
||||
|
||||
```text
|
||||
https://www.example.com → Web UI
|
||||
https://<deviceId>.kvm.example.com → 设备远程访问
|
||||
DEVICE_ENDPOINT_HOST=kvm.example.com
|
||||
```
|
||||
|
||||
这两个地址可以共用 **同一个 HTTPS 端口(443)**,由反向代理根据访问的域名进行路由区分。
|
||||
|
||||
⚠️ **注意:所有配置均需在 `.env` 中完成,不需要修改 `docker-compose.yml`、模板或脚本。**
|
||||
|
||||
|
||||
@@ -7,8 +7,16 @@
|
||||
```bash
|
||||
git clone https://github.com/gl-inet/glkvm-cloud.git
|
||||
cd glkvm-cloud/docker-compose/
|
||||
cp .env.example .env
|
||||
```
|
||||
* For **x86_64 (amd64)**:
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
```
|
||||
* For **arm64 (AArch64)**:
|
||||
```bash
|
||||
cp .env.arm64.example .env
|
||||
```
|
||||
|
||||
2. **Configure environment variables**
|
||||
|
||||
@@ -54,43 +62,66 @@
|
||||
- `OIDC_ALLOWED_USERNAMES`: comma-separated list of allowed usernames (`preferred_username` or `name`) (optional)
|
||||
- `OIDC_ALLOWED_GROUPS`: comma-separated list of allowed OIDC groups (optional)
|
||||
|
||||
**Reverse Proxy Mode (Optional)**
|
||||
|
||||
```env
|
||||
# Enable reverse proxy mode (e.g. Nginx in front of GLKVM Cloud).
|
||||
# When enabled, TLS is terminated by the reverse proxy and GLKVM Cloud runs in plain HTTP.
|
||||
REVERSE_PROXY_ENABLED=false
|
||||
```
|
||||
|
||||
When `REVERSE_PROXY_ENABLED` is set to `true`, GLKVM Cloud is designed to run **behind a reverse proxy** such as Nginx:
|
||||
|
||||
- HTTPS certificates are managed by the reverse proxy (not by GLKVM Cloud itself)
|
||||
- GLKVM Cloud listens on plain HTTP internally
|
||||
- The same HTTPS port can be used for both:
|
||||
- Accessing the GLKVM Cloud web UI
|
||||
- Accessing remote KVM devices
|
||||
|
||||
For example, with proper Nginx configuration,
|
||||
|
||||
```nginx
|
||||
# Forward original host, scheme, port and client IP
|
||||
# These headers are required when running GLKVM Cloud behind a reverse proxy.
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Forwarded-Host $host;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_set_header X-Forwarded-Port $server_port;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
```
|
||||
|
||||
you can use:
|
||||
|
||||
```text
|
||||
https://www.example.com → GLKVM Cloud web interface
|
||||
https://<device_id>.example.com → Remote device access
|
||||
```
|
||||
|
||||
Both addresses can share the **same HTTPS port (443)**, while routing is handled by the reverse proxy based on the domain name.
|
||||
|
||||
#### Reverse Proxy Mode (Optional)
|
||||
|
||||
```env
|
||||
REVERSE_PROXY_ENABLED=false
|
||||
```
|
||||
|
||||
When enabled (`REVERSE_PROXY_ENABLED=true`):
|
||||
|
||||
- GLKVM Cloud runs behind a reverse proxy (e.g. Nginx)
|
||||
- TLS is terminated at the reverse proxy; GLKVM Cloud uses plain HTTP internally
|
||||
- The Web UI and remote device access can share the same HTTPS port (usually 443)
|
||||
|
||||
|
||||
##### Required Reverse Proxy Headers
|
||||
|
||||
The reverse proxy **must** forward the following headers; otherwise, GLKVM Cloud may generate URLs containing internal ports (e.g. `:10443`):
|
||||
|
||||
```nginx
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Forwarded-Host $host;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_set_header X-Forwarded-Port $server_port;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
```
|
||||
|
||||
|
||||
##### Device Remote Access Domain (Optional)
|
||||
|
||||
```env
|
||||
DEVICE_ENDPOINT_HOST=
|
||||
```
|
||||
|
||||
- **Effective only when** `REVERSE_PROXY_ENABLED=true`
|
||||
- Used to specify the domain for device remote access
|
||||
- Device access URLs are generated as:
|
||||
|
||||
```text
|
||||
https://<deviceId>.<DEVICE_ENDPOINT_HOST>/
|
||||
```
|
||||
|
||||
**Notes:**
|
||||
|
||||
- Do not include the scheme (`http://` or `https://`)
|
||||
- Do not include any path
|
||||
- The domain may differ from the Web UI domain
|
||||
- If left empty, the host/port will be derived from `X-Forwarded-*` headers
|
||||
|
||||
**Example:**
|
||||
|
||||
```text
|
||||
https://www.example.com → Web UI
|
||||
https://<deviceId>.kvm.example.com → Device remote access
|
||||
DEVICE_ENDPOINT_HOST=kvm.example.com
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
⚠️ **Note:** All configuration should be done in the `.env` file.
|
||||
You don’t need to modify `docker-compose.yml`, templates, or scripts directly.
|
||||
|
||||
@@ -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;
|
||||
@@ -12,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
|
||||
@@ -52,6 +53,11 @@ services:
|
||||
|
||||
# ---- Reverse Proxy ----
|
||||
REVERSE_PROXY_ENABLED: ${REVERSE_PROXY_ENABLED:-false}
|
||||
|
||||
# ---- Device Endpoint Host ----
|
||||
DEVICE_ENDPOINT_HOST: ${DEVICE_ENDPOINT_HOST:-}
|
||||
# ---- Web UI Host ----
|
||||
WEB_UI_HOST: ${WEB_UI_HOST:-}
|
||||
volumes:
|
||||
- ./templates/rttys.conf.template:/tpl/rttys.conf.tmpl:ro
|
||||
- ./scripts/docker-entrypoint.sh:/docker-entrypoint.sh:ro
|
||||
|
||||
@@ -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}}
|
||||
|
||||
@@ -1,726 +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
|
||||
}
|
||||
// 获取 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
|
||||
}
|
||||
|
||||
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 {
|
||||
// 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"
|
||||
}
|
||||
|
||||
// 1) external port: prefer the one user actually accessed (Host or forwarded headers)
|
||||
port := ""
|
||||
|
||||
// Prefer port from Host
|
||||
if _, p, err := net.SplitHostPort(c.Request.Host); err == nil && p != "" {
|
||||
port = p
|
||||
}
|
||||
|
||||
// Fallback to forwarded headers
|
||||
if 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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2) If still empty, fallback to cfg.AddrHttpProxy (which is a PORT, not a path)
|
||||
if port == "" && strings.TrimSpace(cfg.AddrHttpProxy) != "" {
|
||||
portTmp := strings.TrimSpace(cfg.AddrHttpProxy)
|
||||
// Common cases: ":10443", "0.0.0.0:10443", "[::]:10443"
|
||||
if _, p, err := net.SplitHostPort(portTmp); err == nil {
|
||||
port = p
|
||||
}
|
||||
}
|
||||
|
||||
// 3) Build host: in proxy mode redirect domain to be redirHost
|
||||
hostPort := redirHost
|
||||
if port != "" {
|
||||
// avoid adding default ports
|
||||
if (scheme == "https" && port != "443") || (scheme == "http" && port != "80") {
|
||||
hostPort = net.JoinHostPort(redirHost, port)
|
||||
}
|
||||
}
|
||||
|
||||
// 4) Path: use the current request path
|
||||
redirectPath := c.Request.URL.Path
|
||||
if redirectPath == "" {
|
||||
redirectPath = "/"
|
||||
}
|
||||
|
||||
u := &url.URL{
|
||||
Scheme: scheme,
|
||||
Host: hostPort,
|
||||
Path: redirectPath,
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("sid", sid)
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
location = u.String()
|
||||
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))
|
||||
}
|
||||
|
||||
// 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"
|
||||
// - "www.l1.l2.example.com" -> "devid.l1.l2.example.com"
|
||||
// - Two-level domain "example.com" -> "devid.example.com"
|
||||
// - Single label / abnormal cases -> "devid." + hostname (fallback)
|
||||
//
|
||||
// The input hostname must be a pure hostname without port.
|
||||
func buildRedirectHost(hostname, devid string) string {
|
||||
// Allow FQDN with trailing dot like "example.com."
|
||||
hostname = strings.TrimSuffix(hostname, ".")
|
||||
|
||||
// Split into labels
|
||||
labels := strings.Split(hostname, ".")
|
||||
// Remove empty labels (in case of consecutive dots)
|
||||
compact := make([]string, 0, len(labels))
|
||||
for _, l := range labels {
|
||||
if l != "" {
|
||||
compact = append(compact, l)
|
||||
}
|
||||
}
|
||||
labels = compact
|
||||
|
||||
switch len(labels) {
|
||||
case 0:
|
||||
return devid // extreme case: just return devid
|
||||
case 1:
|
||||
// Single label (e.g., "localhost") — keep original as suffix
|
||||
return devid + "." + labels[0]
|
||||
default:
|
||||
// >=2: drop the leftmost label
|
||||
suffix := strings.Join(labels[1:], ".")
|
||||
return devid + "." + suffix
|
||||
}
|
||||
}
|
||||
|
Before Width: | Height: | Size: 55 KiB After Width: | Height: | Size: 52 KiB |
|
Before Width: | Height: | Size: 38 KiB After Width: | Height: | Size: 44 KiB |
|
After Width: | Height: | Size: 58 KiB |
|
Before Width: | Height: | Size: 61 KiB After Width: | Height: | Size: 14 KiB |
|
Before Width: | Height: | Size: 1.5 MiB |
|
Before Width: | Height: | Size: 418 KiB |
|
Before Width: | Height: | Size: 400 KiB |
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
package devicegroup
|
||||
|
||||
type DeviceGroup struct {
|
||||
ID int64
|
||||
Name string
|
||||
Description string
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
package devicegroup
|
||||
|
||||
import "context"
|
||||
|
||||
type Repository interface {
|
||||
ListDeviceGroupsVisibleToUser(ctx context.Context, userID int64, isAdmin bool) ([]DeviceGroup, error)
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
package group
|
||||
|
||||
type UserGroup struct {
|
||||
ID int64
|
||||
Name string
|
||||
Description string
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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{ /* ... */ }
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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{}
|
||||
@@ -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()},
|
||||
}
|
||||
}
|
||||
@@ -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{}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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{}
|
||||
@@ -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{}
|
||||
@@ -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{}))
|
||||
}
|
||||
@@ -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{}))
|
||||
}
|
||||
@@ -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{}{}))
|
||||
}
|
||||
@@ -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,
|
||||
}))
|
||||
}
|
||||
@@ -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}))
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
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
|
||||
}
|
||||
|
||||
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{}))
|
||||
}
|
||||
@@ -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{}))
|
||||
}
|
||||
@@ -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}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type HostInfo struct {
|
||||
Host string // pure host without port
|
||||
Port string // external port if known
|
||||
Scheme string // http/https
|
||||
RawHost string // req.Host (may include port)
|
||||
XFHost string // X-Forwarded-Host (raw)
|
||||
XFProto string // X-Forwarded-Proto (raw)
|
||||
XFPort string // X-Forwarded-Port (raw)
|
||||
}
|
||||
|
||||
func GetHostInfoFromRequest(req *http.Request) HostInfo {
|
||||
hi := HostInfo{
|
||||
RawHost: req.Host,
|
||||
XFHost: req.Header.Get("X-Forwarded-Host"),
|
||||
XFProto: req.Header.Get("X-Forwarded-Proto"),
|
||||
XFPort: req.Header.Get("X-Forwarded-Port"),
|
||||
}
|
||||
|
||||
// host: prefer X-Forwarded-Host
|
||||
host := strings.TrimSpace(hi.XFHost)
|
||||
if host != "" {
|
||||
host = strings.TrimSpace(strings.Split(host, ",")[0])
|
||||
} else {
|
||||
host = strings.TrimSpace(req.Host)
|
||||
}
|
||||
|
||||
// split port if host contains it
|
||||
if h, p, err := net.SplitHostPort(host); err == nil {
|
||||
hi.Host = h
|
||||
hi.Port = p
|
||||
} else {
|
||||
hi.Host = strings.TrimSuffix(host, ".")
|
||||
}
|
||||
|
||||
// scheme
|
||||
proto := strings.TrimSpace(hi.XFProto)
|
||||
if proto != "" {
|
||||
proto = strings.ToLower(strings.TrimSpace(strings.Split(proto, ",")[0]))
|
||||
hi.Scheme = proto
|
||||
} else if req.TLS != nil {
|
||||
hi.Scheme = "https"
|
||||
} else {
|
||||
hi.Scheme = "http"
|
||||
}
|
||||
|
||||
// forwarded port overrides
|
||||
fp := strings.TrimSpace(hi.XFPort)
|
||||
if fp != "" {
|
||||
hi.Port = strings.TrimSpace(strings.Split(fp, ",")[0])
|
||||
}
|
||||
|
||||
return hi
|
||||
}
|
||||
|
||||
// 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.
|
||||
// Allow:
|
||||
// - exact match: base
|
||||
// - subdomain: *.base
|
||||
func DomainAllowed(host, base string) bool {
|
||||
host = strings.ToLower(strings.TrimSuffix(strings.TrimSpace(host), "."))
|
||||
base = strings.ToLower(strings.TrimSuffix(strings.TrimSpace(base), "."))
|
||||
|
||||
if host == "" || base == "" {
|
||||
return false
|
||||
}
|
||||
if host == base {
|
||||
return true
|
||||
}
|
||||
return strings.HasSuffix(host, "."+base)
|
||||
}
|
||||
|
||||
// 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"
|
||||
// - "www.l1.l2.example.com" -> "devid.l1.l2.example.com"
|
||||
// - Two-level domain "example.com" -> "devid.example.com"
|
||||
// - Single label / abnormal cases -> "devid." + hostname (fallback)
|
||||
//
|
||||
// The input hostname must be a pure hostname without port.
|
||||
func BuildRedirectHost(hostname, devid string) string {
|
||||
// Allow FQDN with trailing dot like "example.com."
|
||||
hostname = strings.TrimSuffix(hostname, ".")
|
||||
|
||||
// Split into labels
|
||||
labels := strings.Split(hostname, ".")
|
||||
// Remove empty labels (in case of consecutive dots)
|
||||
compact := make([]string, 0, len(labels))
|
||||
for _, l := range labels {
|
||||
if l != "" {
|
||||
compact = append(compact, l)
|
||||
}
|
||||
}
|
||||
labels = compact
|
||||
|
||||
switch len(labels) {
|
||||
case 0:
|
||||
return devid // extreme case: just return devid
|
||||
case 1:
|
||||
// Single label (e.g., "localhost") keep original as suffix
|
||||
return devid + "." + labels[0]
|
||||
default:
|
||||
// >=2: drop the leftmost label
|
||||
suffix := strings.Join(labels[1:], ".")
|
||||
return devid + "." + suffix
|
||||
}
|
||||
}
|
||||
|
||||
func JoinHostPortIfNeeded(host, scheme, port string) string {
|
||||
if port == "" {
|
||||
return host
|
||||
}
|
||||
// avoid adding default ports
|
||||
if (scheme == "https" && port == "443") || (scheme == "http" && port == "80") {
|
||||
return host
|
||||
}
|
||||
return net.JoinHostPort(host, port)
|
||||
}
|
||||
|
||||
func BuildRedirectLocation(scheme, hostPort, path, sid string) string {
|
||||
if path == "" {
|
||||
path = "/"
|
||||
}
|
||||
u := &url.URL{
|
||||
Scheme: scheme,
|
||||
Host: hostPort,
|
||||
Path: path,
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("rttysid", sid)
|
||||
u.RawQuery = q.Encode()
|
||||
return u.String()
|
||||
}
|
||||
|
||||
// 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) {
|
||||
// 1) Reverse-proxy headers
|
||||
xfh := strings.TrimSpace(req.Header.Get("X-Forwarded-Host"))
|
||||
xfp := strings.TrimSpace(req.Header.Get("X-Forwarded-Proto"))
|
||||
xfport := strings.TrimSpace(req.Header.Get("X-Forwarded-Port"))
|
||||
|
||||
// X-Forwarded-Host may contain a comma-separated list. Take the first one.
|
||||
if xfh != "" {
|
||||
if i := strings.IndexByte(xfh, ','); i >= 0 {
|
||||
xfh = strings.TrimSpace(xfh[:i])
|
||||
}
|
||||
host = xfh
|
||||
}
|
||||
|
||||
// 2) Fallback to Host header
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(req.Host)
|
||||
}
|
||||
|
||||
// Split host:port if present
|
||||
if h, p, err := net.SplitHostPort(host); err == nil {
|
||||
host = h
|
||||
port = p
|
||||
} else {
|
||||
// no explicit port in Host header
|
||||
port = ""
|
||||
}
|
||||
|
||||
// scheme/proto
|
||||
if xfp != "" {
|
||||
if i := strings.IndexByte(xfp, ','); i >= 0 {
|
||||
xfp = strings.TrimSpace(xfp[:i])
|
||||
}
|
||||
proto = xfp
|
||||
} else if req.TLS != nil {
|
||||
proto = "https"
|
||||
} else {
|
||||
proto = "http"
|
||||
}
|
||||
|
||||
// forwarded port overrides parsed port if present
|
||||
if xfport != "" {
|
||||
if i := strings.IndexByte(xfport, ','); i >= 0 {
|
||||
xfport = strings.TrimSpace(xfport[:i])
|
||||
}
|
||||
port = xfport
|
||||
}
|
||||
|
||||
return host, port, proto
|
||||
}
|
||||
|
||||
// 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) {
|
||||
host = strings.TrimSpace(host)
|
||||
if host == "" {
|
||||
return "", false
|
||||
}
|
||||
|
||||
// remove trailing dot
|
||||
host = strings.TrimSuffix(host, ".")
|
||||
|
||||
// If host is IP, skip
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return "", false
|
||||
}
|
||||
|
||||
labels := strings.Split(host, ".")
|
||||
for _, l := range labels {
|
||||
if l != "" {
|
||||
return l, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
@@ -0,0 +1,405 @@
|
||||
/*
|
||||
* 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.
|
||||
return db.WithContext(ctx).Exec(
|
||||
`INSERT INTO users (username, description, password_hash, role, status, is_system)
|
||||
VALUES (?, 'Admin', ?, 'admin', 'active', 1)
|
||||
ON CONFLICT(username) DO UPDATE SET
|
||||
password_hash=excluded.password_hash,
|
||||
role='admin',
|
||||
status='active',
|
||||
is_system=1`,
|
||||
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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
package server
|
||||
|
||||
const RttysVersion = "5.2.0"
|
||||
const KVMCloudVersion = "v2.0.0"
|
||||
|
||||
var (
|
||||
GitCommit = ""
|
||||
BuildTime = ""
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -1,307 +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"
|
||||
|
||||
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))
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||