Files
2026-07-26 00:09:57 +08:00

413 lines
10 KiB
Go

package core
import (
"crypto/sha256"
"database/sql"
"errors"
"fmt"
mysqlDriver "github.com/go-sql-driver/mysql"
"io"
"log"
"time"
"strconv"
"strings"
// mysql sql驱动
_ "github.com/go-sql-driver/mysql"
)
// Mysql 结构体
type Mysql struct {
Enabled bool `json:"enabled"`
ServerAddr string `json:"server_addr"`
ServerPort int `json:"server_port"`
Database string `json:"database"`
Username string `json:"username"`
Password string `json:"password"`
Cafile string `json:"cafile"`
}
// User 用户表记录结构体
type User struct {
ID uint
Username string
Password string
EncryptPass string
Quota int64
Download uint64
Upload uint64
UseDays uint
ExpiryDate string
}
// PageQuery 分页查询的结构体
type PageQuery struct {
PageNum int
CurPage int
Total int
PageSize int
DataList []*User
}
// CreateTableSql 创表sql
var CreateTableSql = `
CREATE TABLE IF NOT EXISTS users (
id INT UNSIGNED NOT NULL AUTO_INCREMENT,
username VARCHAR(64) NOT NULL,
password CHAR(56) NOT NULL,
passwordShow VARCHAR(255) NOT NULL,
quota BIGINT NOT NULL DEFAULT 0,
download BIGINT UNSIGNED NOT NULL DEFAULT 0,
upload BIGINT UNSIGNED NOT NULL DEFAULT 0,
useDays int(10) DEFAULT 0,
expiryDate char(10) DEFAULT '',
PRIMARY KEY (id),
INDEX (password)
) DEFAULT CHARSET=utf8mb4;
`
// GetDB 获取mysql数据库连接
func (mysql *Mysql) GetDB() *sql.DB {
// 屏蔽mysql驱动包的日志输出
mysqlDriver.SetLogger(log.New(io.Discard, "", 0))
conn := fmt.Sprintf("%s:%s@tcp(%s:%d)/%s", mysql.Username, mysql.Password, mysql.ServerAddr, mysql.ServerPort, mysql.Database)
db, err := sql.Open("mysql", conn)
if err != nil {
fmt.Println(err.Error())
return nil
}
return db
}
// CreateTable 不存在trojan user表则自动创建
func (mysql *Mysql) CreateTable() {
db := mysql.GetDB()
defer db.Close()
if _, err := db.Exec(CreateTableSql); err != nil {
fmt.Println(err)
}
}
func queryUserList(db *sql.DB, sql string) ([]*User, error) {
var (
username string
encryptPass string
passShow string
download uint64
upload uint64
quota int64
id uint
useDays uint
expiryDate string
)
var userList []*User
rows, err := db.Query(sql)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
if err := rows.Scan(&id, &username, &encryptPass, &passShow, &quota, &download, &upload, &useDays, &expiryDate); err != nil {
return nil, err
}
userList = append(userList, &User{
ID: id,
Username: username,
Password: passShow,
EncryptPass: encryptPass,
Download: download,
Upload: upload,
Quota: quota,
UseDays: useDays,
ExpiryDate: expiryDate,
})
}
return userList, nil
}
func queryUser(db *sql.DB, sql string) (*User, error) {
var (
username string
encryptPass string
passShow string
download uint64
upload uint64
quota int64
id uint
useDays uint
expiryDate string
)
row := db.QueryRow(sql)
if err := row.Scan(&id, &username, &encryptPass, &passShow, &quota, &download, &upload, &useDays, &expiryDate); err != nil {
return nil, err
}
return &User{ID: id, Username: username, Password: passShow, EncryptPass: encryptPass, Download: download, Upload: upload, Quota: quota, UseDays: useDays, ExpiryDate: expiryDate}, nil
}
// CreateUser 创建Trojan用户
func (mysql *Mysql) CreateUser(username string, base64Pass string, originPass string) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
encryPass := sha256.Sum224([]byte(originPass))
if _, err := db.Exec(fmt.Sprintf("INSERT INTO users(username, password, passwordShow, quota) VALUES ('%s', '%x', '%s', -1);", username, encryPass, base64Pass)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// UpdateUser 更新Trojan用户名和密码
func (mysql *Mysql) UpdateUser(id uint, username string, base64Pass string, originPass string) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
encryPass := sha256.Sum224([]byte(originPass))
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET username='%s', password='%x', passwordShow='%s' WHERE id=%d;", username, encryPass, base64Pass, id)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// DeleteUser 删除用户
func (mysql *Mysql) DeleteUser(id uint) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
if userList, err := mysql.GetData(strconv.Itoa(int(id))); err != nil {
return err
} else if userList != nil && len(userList) == 0 {
return fmt.Errorf("不存在id为%d的用户", id)
}
if _, err := db.Exec(fmt.Sprintf("DELETE FROM users WHERE id=%d;", id)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// MonthlyResetData 设置了过期时间的用户,每月定时清空使用流量
func (mysql *Mysql) MonthlyResetData() error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
userList, err := queryUserList(db, "SELECT * FROM users WHERE useDays != 0 AND quota != 0")
if err != nil {
return err
}
for _, user := range userList {
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET download=0, upload=0 WHERE id=%d;", user.ID)); err != nil {
return err
}
}
return nil
}
// DailyCheckExpire 检查是否有过期,过期了设置流量上限为0
func (mysql *Mysql) DailyCheckExpire() (bool, error) {
needRestart := false
now := time.Now()
utc, err := time.LoadLocation("Asia/Shanghai")
if err != nil {
return false, err
}
addDay, _ := time.ParseDuration("-24h")
yesterdayStr := now.Add(addDay).In(utc).Format("2006-01-02")
yesterday, _ := time.Parse("2006-01-02", yesterdayStr)
db := mysql.GetDB()
if db == nil {
return false, errors.New("can't connect mysql")
}
defer db.Close()
userList, err := queryUserList(db, "SELECT * FROM users WHERE quota != 0")
if err != nil {
return false, err
}
for _, user := range userList {
if expireDate, err := time.Parse("2006-01-02", user.ExpiryDate); err == nil {
if yesterday.Sub(expireDate).Seconds() >= 0 {
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET quota=0 WHERE id=%d;", user.ID)); err != nil {
return false, err
}
if !needRestart {
needRestart = true
}
}
}
}
return needRestart, nil
}
// CancelExpire 取消过期时间
func (mysql *Mysql) CancelExpire(id uint) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET useDays=0, expiryDate='' WHERE id=%d;", id)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// SetExpire 设置过期时间
func (mysql *Mysql) SetExpire(id uint, useDays uint) error {
now := time.Now()
utc, err := time.LoadLocation("Asia/Shanghai")
if err != nil {
fmt.Println(err)
return err
}
addDay, _ := time.ParseDuration(strconv.Itoa(int(24*useDays)) + "h")
expiryDate := now.Add(addDay).In(utc).Format("2006-01-02")
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET useDays=%d, expiryDate='%s' WHERE id=%d;", useDays, expiryDate, id)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// SetQuota 限制流量
func (mysql *Mysql) SetQuota(id uint, quota int) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET quota=%d WHERE id=%d;", quota, id)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// CleanData 清空流量统计
func (mysql *Mysql) CleanData(id uint) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET download=0, upload=0 WHERE id=%d;", id)); err != nil {
fmt.Println(err)
return err
}
return nil
}
// CleanDataByName 清空指定用户名流量统计数据
func (mysql *Mysql) CleanDataByName(usernames []string) error {
db := mysql.GetDB()
if db == nil {
return errors.New("can't connect mysql")
}
defer db.Close()
runSql := "UPDATE users SET download=0, upload=0 WHERE BINARY username in ("
for i, name := range usernames {
runSql = runSql + "'" + name + "'"
if i == len(usernames)-1 {
runSql = runSql + ")"
} else {
runSql = runSql + ","
}
}
if _, err := db.Exec(runSql); err != nil {
fmt.Println(err)
return err
}
return nil
}
// GetUserByName 通过用户名来获取用户
func (mysql *Mysql) GetUserByName(name string) *User {
db := mysql.GetDB()
if db == nil {
return nil
}
defer db.Close()
user, err := queryUser(db, fmt.Sprintf("SELECT * FROM users WHERE BINARY username='%s'", name))
if err != nil {
return nil
}
return user
}
// GetUserByPass 通过密码来获取用户
func (mysql *Mysql) GetUserByPass(pass string) *User {
db := mysql.GetDB()
if db == nil {
return nil
}
defer db.Close()
user, err := queryUser(db, fmt.Sprintf("SELECT * FROM users WHERE BINARY passwordShow='%s'", pass))
if err != nil {
return nil
}
return user
}
// PageList 通过分页获取用户记录
func (mysql *Mysql) PageList(curPage int, pageSize int) (*PageQuery, error) {
var (
total int
)
db := mysql.GetDB()
if db == nil {
return nil, errors.New("连接mysql失败")
}
defer db.Close()
offset := (curPage - 1) * pageSize
querySQL := fmt.Sprintf("SELECT * FROM users LIMIT %d, %d", offset, pageSize)
userList, err := queryUserList(db, querySQL)
if err != nil {
fmt.Println(err)
return nil, err
}
db.QueryRow("SELECT COUNT(id) FROM users").Scan(&total)
return &PageQuery{
CurPage: curPage,
PageSize: pageSize,
Total: total,
DataList: userList,
PageNum: (total + pageSize - 1) / pageSize,
}, nil
}
// GetData 获取用户记录
func (mysql *Mysql) GetData(ids ...string) ([]*User, error) {
querySQL := "SELECT * FROM users"
db := mysql.GetDB()
if db == nil {
return nil, errors.New("连接mysql失败")
}
defer db.Close()
if len(ids) > 0 {
querySQL = querySQL + " WHERE id in (" + strings.Join(ids, ",") + ")"
}
userList, err := queryUserList(db, querySQL)
if err != nil {
fmt.Println(err)
return nil, err
}
return userList, nil
}