413 lines
10 KiB
Go
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, "a, &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, "a, &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
|
|
}
|