release revsuit vBeta0.1.0 (#15)

Co-authored-by: E99p1ant <[email protected]>
This commit is contained in:
Li4n0
2021-05-16 12:11:35 +08:00
committed by GitHub
co-authored by E99p1ant
parent b39ee101f9
commit 5b63e90390
98 changed files with 2357 additions and 623 deletions
+147 -114
View File
@@ -11,12 +11,15 @@ import (
"github.com/li4n0/revsuit/internal/recycler"
"github.com/li4n0/revsuit/internal/rule"
"github.com/patrickmn/go-cache"
"github.com/pkg/errors"
log "unknwon.dev/clog/v2"
)
type Server struct {
rules []*Rule
rulesLock sync.RWMutex
Config
rules []*Rule
rulesLock sync.RWMutex
livingLock sync.Mutex
}
var (
@@ -27,7 +30,7 @@ var (
func GetServer() *Server {
once.Do(func() {
server = &Server{rulesLock: sync.RWMutex{}}
server = &Server{rulesLock: sync.RWMutex{}, livingLock: sync.Mutex{}}
})
return server
}
@@ -42,141 +45,171 @@ func (s *Server) updateRules() error {
db := database.DB.Model(new(Rule))
defer s.rulesLock.Unlock()
s.rulesLock.Lock()
return db.Order("rank desc").Find(&s.rules).Error
return errors.Wrap(db.Order("rank desc").Find(&s.rules).Error, "DNS update rules error")
}
func (s *Server) Run() {
func newSet(_rule *Rule, name, value, ip string, _type newdns.Type) []newdns.Set {
set := []newdns.Set{
{
Name: name,
Type: _type,
Records: func() []newdns.Record {
switch _rule.Type {
case newdns.TXT:
return []newdns.Record{{Data: []string{value}}}
case newdns.CNAME, newdns.NS:
return []newdns.Record{{Address: value + "."}}
case newdns.REBINDING:
if err := s.updateRules(); err != nil {
log.Fatal(err.Error())
// Get rebinding ip list
values, ok := rebindingCache.Get(ip)
if !ok {
rebindingCache.Set(ip, strings.Split(value, ","), cache.DefaultExpiration)
values = strings.Split(value, ",")
}
//Choose and delete first ip
value := values.([]string)[0]
if len(values.([]string)) > 1 {
rebindingCache.Set(ip, values.([]string)[1:len(values.([]string))], cache.DefaultExpiration)
} else {
rebindingCache.Delete(ip)
}
log.Trace("DNS rebinding client[ip:%v] to %v", ip, value)
return []newdns.Record{{Address: value}}
default:
return []newdns.Record{{Address: value}}
}
}(),
TTL: _rule.TTL * time.Second,
},
}
return set
}
//create new dns zone with root domain
newZone := func(name string) *newdns.Zone {
defer func() {
if err := recover(); err != nil {
recycler.Recycle(err)
}
}()
domain := strings.TrimSuffix(name, ".")
frags := strings.Split(domain, ".")
zoneName := ""
if len(frags) >= 2 {
zoneName = strings.Join(frags[len(frags)-2:], ".") + "."
} else {
zoneName = name
// newZone creates new dns zone with root domain
func (s *Server) newZone(name string) *newdns.Zone {
defer func() {
if err := recover(); err != nil {
recycler.Recycle(err)
}
return &newdns.Zone{
Name: zoneName,
MasterNameServer: "ns1.hostmaster.com.",
AllNameServers: []string{
"ns1.hostmaster.com.",
"ns2.hostmaster.com.",
"ns3.hostmaster.com.",
},
Handler: func(lookedName, remoteAddr string) ([]newdns.Set, error) {
ip := strings.Split(remoteAddr, ":")[0]
}()
for _, _rule := range s.getRules() {
flag, flagGroup, vars := _rule.Match(domain)
if flag == "" {
continue
}
domain := strings.TrimSuffix(name, ".")
frags := strings.Split(domain, ".")
zoneName := name
if len(frags) >= 2 {
zoneName = strings.Join(frags[len(frags)-2:], ".") + "."
}
zone := &newdns.Zone{
Name: zoneName,
MasterNameServer: "ns1.hostmaster.com.",
AllNameServers: []string{
"ns1.hostmaster.com.",
"ns2.hostmaster.com.",
"ns3.hostmaster.com.",
},
Handler: func(lookedName, remoteAddr string) ([]newdns.Set, error) {
ip := strings.Split(remoteAddr, ":")[0]
r, err := newRecord(_rule, flag, domain, ip, qqwry.Area(ip))
if err != nil {
log.Warn("DNS record(rule_id:%s) created failed :%s", _rule.Name, err)
return nil, nil
}
log.Info("DNS record[id:%d rule:%s remote_ip:%s] has been created", r.ID, _rule.Name, ip)
for _, _rule := range s.getRules() {
flag, flagGroup, vars := _rule.Match(domain)
if flag == "" {
continue
}
//only send to client when this connection recorded first time.
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and domain like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("DNS record[id%d] has been put to client message queue", r.ID)
}
} else {
r, err := newRecord(_rule, flag, domain, ip, qqwry.Area(ip))
if err != nil {
log.Warn("DNS record(rule_id:%s) created failed :%s", _rule.Name, err)
return nil, nil
}
log.Info("DNS record[id:%d rule:%s remote_ip:%s] has been created", r.ID, _rule.Name, ip)
//only send to client when this connection recorded first time.
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and domain like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("DNS record[id%d] has been put to client message queue", r.ID)
log.Trace("DNS record[id:%d, flagGroup:%s] has been put to client message queue", r.ID, flagGroup)
}
}
//send notice
if _rule.Notice {
go func() {
r.Notice()
log.Trace("DNS record[id%d] notice has been sent", r.ID)
}()
}
if _rule.Value != "" {
value := rule.CompileTpl(_rule.Value, vars)
_type := _rule.Type
if _rule.Type == newdns.REBINDING {
_type = newdns.A
}
return []newdns.Set{
{
Name: name,
Type: _type,
Records: func() []newdns.Record {
switch _rule.Type {
case newdns.TXT:
return []newdns.Record{{Data: []string{value}}}
case newdns.CNAME, newdns.NS:
return []newdns.Record{{Address: value + "."}}
case newdns.REBINDING:
// Get rebinding ip list
values, ok := rebindingCache.Get(ip)
if !ok {
rebindingCache.Set(ip, strings.Split(value, ","), cache.DefaultExpiration)
values = strings.Split(value, ",")
}
//Choose and delete first ip
value := values.([]string)[0]
if len(values.([]string)) > 1 {
rebindingCache.Set(ip, values.([]string)[1:len(values.([]string))], cache.DefaultExpiration)
} else {
rebindingCache.Delete(ip)
}
log.Trace("DNS rebinding client(ip:%v) to %v", ip, value)
return []newdns.Record{{Address: value}}
default:
return []newdns.Record{{Address: value}}
}
}(),
TTL: _rule.TTL * time.Second,
},
}, nil
} else {
r.PushToClient()
log.Trace("DNS record[id:%d, flag:%s] has been put to client message queue", r.ID, flag)
}
}
return nil, nil
},
//send notice
if _rule.Notice {
go func() {
r.Notice()
log.Trace("DNS record[id:%d] notice has been sent", r.ID)
}()
}
if _rule.Value != "" {
value := rule.CompileTpl(_rule.Value, vars)
_type := _rule.Type
if _rule.Type == newdns.REBINDING {
_type = newdns.A
}
return newSet(_rule, name, value, ip, _type), nil
}
}
return nil, nil
},
}
return zone
}
func (s *Server) Stop() {
log.Info("DNS server is stopping...")
s.Enable = false
s.livingLock.Unlock()
}
func (s *Server) Run() {
s.Enable = true
s.livingLock.Lock()
defer func() {
if s.Enable {
log.Error("DNS Server exited unexpectedly")
}
s.Enable = false
s.livingLock.Unlock()
}()
if err := s.updateRules(); err != nil {
log.Error(err.Error())
return
}
// create server
server := newdns.NewServer(newdns.Config{
Handler: func(name string) (*newdns.Zone, error) {
return newZone(name), nil
return s.newZone(name), nil
},
})
// run server
log.Info("Starting DNS Server at :53")
err := server.Run(":53")
if err != nil {
log.Fatal(err.Error())
}
go func() {
s.livingLock.Lock()
if !s.Enable {
server.Close()
}
}()
err := server.Run(":53")
defer server.Close()
if err != nil {
log.Error(err.Error())
return
}
}
+13 -4
View File
@@ -47,11 +47,20 @@ func ListRecords(c *gin.Context) {
res []Record
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&dnsRecord); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
}
@@ -74,7 +83,7 @@ func ListRecords(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -84,10 +93,10 @@ func ListRecords(c *gin.Context) {
order = "desc"
}
if err := db.Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
+34 -25
View File
@@ -13,17 +13,17 @@ import (
)
type Rule struct {
rule.BaseRule
Type newdns.Type `gorm:"default:1" form:"type" json:"type"`
Value string `form:"value" json:"value"`
TTL time.Duration `gorm:"ttl;default:10" form:"ttl" json:"ttl"`
rule.BaseRule `yaml:""`
Type newdns.Type `gorm:"default:1" form:"type" json:"type"`
Value string `form:"value" json:"value"`
TTL time.Duration `gorm:"ttl;default:10" form:"ttl" json:"ttl"`
}
func (Rule) TableName() string {
return "dns_rules"
}
// New dns rule struct
// NewRule new dns rule struct
func NewRule(name, flagFormat, value string, pushToClient, notice bool, _type newdns.Type, ttl time.Duration) *Rule {
return &Rule{
BaseRule: rule.BaseRule{
@@ -38,7 +38,7 @@ func NewRule(name, flagFormat, value string, pushToClient, notice bool, _type ne
}
}
// Create or update the dns rule in database and ruleSet
// CreateOrUpdate creates or updates the dns rule in database and ruleSet
func (r *Rule) CreateOrUpdate() (err error) {
db := database.DB.Model(r)
err = db.Clauses(clause.OnConflict{
@@ -62,7 +62,7 @@ func (r *Rule) CreateOrUpdate() (err error) {
return err
}
// Delete the dns rule in database and ruleSet
// Delete deletes the dns rule in database and ruleSet
func (r *Rule) Delete() (err error) {
db := database.DB.Model(r)
err = db.Delete(r).Error
@@ -73,19 +73,28 @@ func (r *Rule) Delete() (err error) {
return err
}
// List all dns rules those satisfy the filter
// ListRules lists all dns rules those satisfy the filter
func ListRules(c *gin.Context) {
var (
dnsRule Rule
res []Rule
count int64
order = c.Query("order")
dnsRule Rule
res []Rule
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&dnsRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -101,7 +110,7 @@ func ListRules(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -111,10 +120,10 @@ func ListRules(c *gin.Context) {
order = "desc"
}
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -127,7 +136,7 @@ func ListRules(c *gin.Context) {
})
}
// Create or update dns rule from user submit
// UpsertRules creates or updates dns rule from user submit
func UpsertRules(c *gin.Context) {
var (
dnsRule Rule
@@ -137,7 +146,7 @@ func UpsertRules(c *gin.Context) {
if err := c.ShouldBind(&dnsRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -150,16 +159,16 @@ func UpsertRules(c *gin.Context) {
if err := dnsRule.CreateOrUpdate(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
if update {
log.Trace("DNS rule[id%d] has been updated", dnsRule.ID)
log.Trace("DNS rule[id:%d] has been updated", dnsRule.ID)
} else {
log.Trace("DNS rule[id%d] has been created", dnsRule.ID)
log.Trace("DNS rule[id:%d] has been created", dnsRule.ID)
}
c.JSON(200, gin.H{
@@ -169,14 +178,14 @@ func UpsertRules(c *gin.Context) {
})
}
// Delete dns rule from user submit
// DeleteRules Delete dns rule from user submit
func DeleteRules(c *gin.Context) {
var dnsRule Rule
if err := c.ShouldBind(&dnsRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -185,13 +194,13 @@ func DeleteRules(c *gin.Context) {
if err := dnsRule.Delete(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
log.Trace("DNS rule[id%d] has been deleted", dnsRule.ID)
log.Trace("DNS rule[id:%d] has been deleted", dnsRule.ID)
c.JSON(200, gin.H{
"status": "succeed",
+130 -78
View File
@@ -1,6 +1,7 @@
package ftp
import (
"bufio"
"bytes"
"fmt"
"go/types"
@@ -12,11 +13,10 @@ import (
"time"
"github.com/li4n0/revsuit/internal/database"
"github.com/li4n0/revsuit/internal/file"
"github.com/li4n0/revsuit/internal/qqwry"
"github.com/li4n0/revsuit/internal/recycler"
"github.com/li4n0/revsuit/internal/rule"
"github.com/patrickmn/go-cache"
"github.com/pkg/errors"
log "unknwon.dev/clog/v2"
)
@@ -24,6 +24,7 @@ type Server struct {
Config
rules []*Rule
rulesLock sync.RWMutex
livingLock sync.Mutex
dataChannel chan map[string]interface{}
}
@@ -64,7 +65,12 @@ func (s *Server) updateRules() error {
db := database.DB.Model(new(Rule))
defer s.rulesLock.Unlock()
s.rulesLock.Lock()
return db.Order("rank desc").Find(&s.rules).Error
return errors.Wrap(db.Order("rank desc").Find(&s.rules).Error, "FTP update rules error")
}
func getClientPasvConnAddress(ip, port string) string {
dataPort, _ := strconv.Atoi(port)
return fmt.Sprintf("%s:%d", ip, dataPort+1)
}
func (s *Server) authenticate(user, password string) (_rule *Rule, flag, flagGroup string, vars map[string]string) {
@@ -99,7 +105,28 @@ func (s *Server) getPasvAddressFromCache(ip, pasvAddressTpl string) (pasvAddress
return pasvAddress
}
const (
NeedAccount = "332 Need account for login.\r\n"
PasswordPlease = "331 password please - version check\r\n"
PasswordError = "331 please specify the password\r\n"
UserLogged = "230 User logged in\r\n"
NoSuchFile = "550 %s: No such file or directory.\r\n"
CommandNotFound = "500 '%s': command not understood.\r\n"
EnteringPassiveMode = "227 Entering Passive Mode (%s,%v,%d)\r\n"
OpeningBinaryMode = "150 Opening BINARY mode data connection for '%s' (%d bytes).\r\n"
OpeningBinaryModeUpload = "150 Opening BINARY mode data connection for '%s'.\r\n"
TransferComplete = "226 Transfer complete.\r\n"
Goodbye = "221 Goodbye.\r\n"
DirectoryChanged = "250 Directory successfully changed.\r\n"
CurrentDirectory = "257 \"%s\" is the current directory\r\n"
)
func (s *Server) handleConnection(conn net.Conn) {
defer func() {
if err := recover(); err != nil {
recycler.Recycle(err)
}
}()
log.Trace("New FTP connection from addr [%s]", conn.RemoteAddr())
defer func() {
_ = conn.Close()
@@ -117,9 +144,9 @@ func (s *Server) handleConnection(conn net.Conn) {
}
ip, port, _ := net.SplitHostPort(conn.RemoteAddr().String())
dataPort, _ := strconv.Atoi(port)
clientPasvConnAddress := fmt.Sprintf("%s:%d", ip, dataPort+1)
clientPasvConnAddress := getClientPasvConnAddress(ip, port)
buf := &bytes.Buffer{}
connBuf := bufio.NewWriter(conn)
var user, password, method, flag, flagGroup, pasvAddress, filename string
status := CRASHED
@@ -148,23 +175,23 @@ loop:
log.Trace("FTP connection[%s] exec command: %s", conn.RemoteAddr(), strings.TrimRight(buf.String(), "\r\n"))
if _rule == nil && cmd != "USER" && cmd != "PASS" {
_, _ = conn.Write([]byte("332 Need account for login.\r\n"))
_, _ = connBuf.WriteString(NeedAccount)
break loop
}
switch cmd {
case "USER":
user = args
_, _ = conn.Write([]byte("331 password please - version check\r\n"))
_, _ = connBuf.WriteString(PasswordPlease)
case "PASS":
password = args
if _rule, flag, flagGroup, vars = s.authenticate(user, password); _rule == nil {
_, _ = conn.Write([]byte("331 please specify the password\r\n"))
_, _ = connBuf.WriteString(PasswordError)
break loop
}
log.Trace("FTP connection[%s] matched rule[rule_name: %s, flag: %s]", conn.RemoteAddr(), _rule.Name, flag)
_, _ = conn.Write([]byte("230 User logged in\r\n"))
_, _ = connBuf.WriteString(UserLogged)
if pasvAddress = s.getPasvAddressFromCache(ip, _rule.PasvAddress); pasvAddress == "" {
pasvAddress = fmt.Sprintf("%s:%d", s.PasvIP, s.PasvPort)
@@ -174,13 +201,13 @@ loop:
case "SIZE":
path += strings.TrimLeft(args, "/")
if _rule == nil || isRedirect || len(_rule.Data) == 0 {
_, _ = conn.Write([]byte(fmt.Sprintf("550 %s: No such file or directory.\r\n", args)))
_, _ = connBuf.WriteString(fmt.Sprintf(NoSuchFile, args))
break
}
_, _ = conn.Write([]byte(fmt.Sprintf("213 %d\r\n", len(_rule.Data))))
_, _ = connBuf.WriteString(fmt.Sprintf("213 %d\r\n", len(_rule.Data)))
case "EPSV", "EPRT", "PORT":
// refuse to use EPSV/EPRT/PORT in order to make the client to use PASV mode.
_, _ = conn.Write([]byte(fmt.Sprintf("500 '%s': command not understood.\r\n", cmd)))
_, _ = connBuf.WriteString(fmt.Sprintf(CommandNotFound, cmd))
case "PASV":
//Just so that ide does not prompt that there may be a nil value
if _rule != nil {
@@ -188,17 +215,17 @@ loop:
pasvAddress := rule.CompileTpl(pasvAddress, vars)
pasvIP, pasvPort, err := net.SplitHostPort(pasvAddress)
if err != nil {
log.Warn("FTP failed to split rule[id%d] pasv_address(%s) :%s", _rule.ID, pasvAddress, err)
log.Warn("FTP failed to split rule[id:%d] pasv_address(%s) :%s", _rule.ID, pasvAddress, err)
break
}
port, err := strconv.Atoi(pasvPort)
if err != nil {
log.Warn("FTP failed to convert rule[id%d] pasv_port(%s) :%s", _rule.ID, pasvPort, err)
log.Warn("FTP failed to convert rule[id:%d] pasv_port(%s) :%s", _rule.ID, pasvPort, err)
break
}
ret := fmt.Sprintf("227 Entering Passive Mode (%s,%v,%d)\r\n", strings.ReplaceAll(pasvIP, ".", ","), float64(port/256), port%256)
_, _ = conn.Write([]byte(ret))
ret := fmt.Sprintf(EnteringPassiveMode, strings.ReplaceAll(pasvIP, ".", ","), float64(port/256), port%256)
_, _ = connBuf.WriteString(ret)
if isRedirect {
log.Trace("FTP connection[%s] will be redirect[pasv_address: %s]", conn.RemoteAddr(), pasvAddress)
}
@@ -210,16 +237,18 @@ loop:
method = DOWNLOAD
//send data to client
_, _ = conn.Write([]byte(fmt.Sprintf("150 Opening BINARY mode data connection for '%s' (%d bytes).\r\n", filename, len(_rule.Data))))
_, _ = connBuf.WriteString(fmt.Sprintf(OpeningBinaryMode, filename, len(_rule.Data)))
_ = connBuf.Flush()
s.dataChannel <- map[string]interface{}{clientPasvConnAddress: []byte(rule.CompileTpl(_rule.Data, vars))}
_, _ = conn.Write([]byte("226 Transfer complete.\r\n"))
_, _ = connBuf.WriteString(TransferComplete)
}
case "STOR":
filename = args
method = UPLOAD
_, _ = conn.Write([]byte(fmt.Sprintf("150 Opening BINARY mode data connection for '%s'.\r\n", filename)))
_, _ = connBuf.WriteString(fmt.Sprintf(OpeningBinaryModeUpload, filename))
_ = connBuf.Flush()
//only could read data send to local pasv server.
if !isRedirect {
dataChannel := make(chan []byte)
@@ -227,68 +256,36 @@ loop:
uploadData = <-dataChannel
log.Trace("FTP connection[%s] uploaded %d bytes", conn.RemoteAddr(), len(uploadData))
}
_, _ = conn.Write([]byte("226 Transfer complete.\r\n"))
_, _ = connBuf.WriteString(TransferComplete)
case "QUIT":
_, _ = conn.Write([]byte("221 Goodbye.\r\n"))
_, _ = connBuf.WriteString(Goodbye)
status = FINISHED
break loop
case "CWD":
_, _ = conn.Write([]byte("250 Directory successfully changed.\r\n"))
_, _ = connBuf.WriteString(DirectoryChanged)
path += strings.TrimRight(args, "\r\n") + "/"
case "PWD":
_, _ = conn.Write([]byte(fmt.Sprintf("257 \"%s\" is the current directory\r\n", path)))
_, _ = connBuf.WriteString(fmt.Sprintf(CurrentDirectory, path))
default:
_, _ = conn.Write([]byte("230 more data please!\r\n"))
}
_ = connBuf.Flush()
}
buf = &bytes.Buffer{}
}
if _rule != nil {
area := qqwry.Area(ip)
var r *Record
var err error
// create new record
ftpFile := &file.FTPFile{}
if len(uploadData) != 0 {
ftpFile = &file.FTPFile{
Name: filename,
Content: uploadData,
}
}
r, err = NewRecord(_rule, flag, user, password, method, path, ip, area, ftpFile, status)
if err != nil {
log.Warn("FTP record[rule_id:%d] created failed :%s", _rule.ID, err)
return
}
log.Info("FTP record[id:%d rule:%s remote_ip:%s] has been created", r.ID, _rule.Name, ip)
//only send to client when this connection recorded first time.
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and raw like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("FTP record[id%d] has been put to client message queue", r.ID)
}
}
r.PushToClient()
log.Trace("FTP record[id%d] has been put to client message queue", r.ID)
}
//send notice
if _rule.Notice {
go func() {
r.Notice()
log.Trace("FTP record[id%d] notice has been sent", r.ID)
}()
}
createRecord(_rule, flag, flagGroup, user, password, method, path, filename, ip, uploadData, status)
}
}
func (s *Server) handlePasvConnection(conn net.Conn, data map[string]interface{}) {
defer func() {
if err := recover(); err != nil {
recycler.Recycle(err)
}
}()
remoteAddress := conn.RemoteAddr().String()
switch v := data[remoteAddress].(type) {
case types.Nil:
@@ -312,42 +309,97 @@ func (s *Server) handlePasvConnection(conn net.Conn, data map[string]interface{}
_ = conn.Close()
}
func (s *Server) Run() {
if err := s.updateRules(); err != nil {
log.Fatal(err.Error())
// run pasv server
func (s *Server) runPasvServer() (net.Listener, error) {
pasvAddress := fmt.Sprintf("%s:%d", strings.Split(s.Addr, ":")[0], s.PasvPort)
log.Info("Start to listen FTP PASV port at %v, PasvIP is %v", pasvAddress, s.PasvIP)
listener, err := net.Listen("tcp", pasvAddress)
if err != nil {
return nil, errors.Wrap(err, "FTP failed to listen on pasv port")
}
// run pasv server
go func() {
pasvAddress := fmt.Sprintf("%s:%d", strings.Split(s.Addr, ":")[0], s.PasvPort)
log.Info("Start to listen FTP PASV port at %v", pasvAddress)
listener, err := net.Listen("tcp", pasvAddress)
if err != nil {
log.Fatal("FTP failed to listen on pasv port : %v", err)
}
for data := range s.dataChannel {
tcpConn, err := listener.Accept()
if err != nil {
log.Warn("FTP accept connection error: %v", err)
if !strings.Contains(err.Error(), net.ErrClosed.Error()) {
log.Warn("FTP accept connection error: %v", err)
} else {
break
}
continue
}
s.handlePasvConnection(tcpConn, data)
}
}()
return listener, nil
}
func (s *Server) Stop() {
log.Info("FTP Server is stopping...")
s.Enable = false
s.livingLock.Unlock()
}
func (s *Server) Restart() {
s.Stop()
time.Sleep(time.Second * 2)
go s.Run()
}
func (s *Server) Run() {
s.Enable = true
s.livingLock.Lock()
defer func() {
if s.Enable {
log.Error("FTP Server exited unexpectedly")
}
s.Enable = false
s.livingLock.Unlock()
}()
if err := s.updateRules(); err != nil {
log.Error(err.Error())
return
}
pasvListener, err := s.runPasvServer()
if err != nil {
log.Error(err.Error())
}
defer func() {
if pasvListener != nil {
_ = pasvListener.Close()
}
}()
// run ftp server
log.Info("Starting FTP Server at %v", s.Addr)
listener, err := net.Listen("tcp", s.Addr)
if err != nil {
log.Fatal(err.Error())
log.Error(errors.Wrap(err, "FTP failed to start").Error())
return
}
for {
go func() {
s.livingLock.Lock()
if !s.Enable {
_ = listener.Close()
}
}()
for s.Enable {
tcpConn, err := listener.Accept()
if err != nil {
log.Warn("FTP accept connection error: %v", err)
if !strings.Contains(err.Error(), net.ErrClosed.Error()) {
log.Warn("FTP accept connection error: %v", err)
} else {
break
}
continue
}
go s.handleConnection(tcpConn)
}
}
+60 -4
View File
@@ -8,7 +8,9 @@ import (
"github.com/li4n0/revsuit/internal/database"
"github.com/li4n0/revsuit/internal/file"
"github.com/li4n0/revsuit/internal/notice"
"github.com/li4n0/revsuit/internal/qqwry"
"github.com/li4n0/revsuit/internal/record"
log "unknwon.dev/clog/v2"
)
var _ record.Record = (*Record)(nil)
@@ -57,12 +59,21 @@ func ListRecords(c *gin.Context) {
res []Record
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&ftpRecord); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -98,7 +109,7 @@ func ListRecords(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -108,10 +119,10 @@ func ListRecords(c *gin.Context) {
order = "desc"
}
if err := db.Preload("File").Order("id " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Preload("File").Order("id " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -123,3 +134,48 @@ func ListRecords(c *gin.Context) {
"result": gin.H{"count": count, "data": res},
})
}
func createRecord(_rule *Rule, flag, flagGroup, user, password, method, path, filename, ip string, uploadData []byte, status Status) {
// create new record
area := qqwry.Area(ip)
var ftpFile *file.FTPFile
var r *Record
var err error
if len(uploadData) != 0 {
ftpFile = &file.FTPFile{
Name: filename,
Content: uploadData,
}
}
r, err = NewRecord(_rule, flag, user, password, method, path, ip, area, ftpFile, status)
if err != nil {
log.Warn("FTP record[rule_id:%d] created failed :%s", _rule.ID, err)
return
}
log.Info("FTP record[id:%d rule:%s remote_ip:%s] has been created", r.ID, _rule.Name, ip)
//only send to client when this connection recorded first time.
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and (user like ? or password like ?)", _rule.Name, "%"+flagGroup+"%", "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("FTP record[id:%d, flagGroup:%s] has been put to client message queue", r.ID, flagGroup)
}
} else {
r.PushToClient()
log.Trace("FTP record[id:%d, flag:%s] has been put to client message queue", r.ID, flag)
}
}
//send notice
if _rule.Notice {
go func() {
r.Notice()
log.Trace("FTP record[id:%d] notice has been sent", r.ID)
}()
}
}
+30 -21
View File
@@ -10,11 +10,11 @@ import (
log "unknwon.dev/clog/v2"
)
// FTP rule struct
// Rule FTP rule struct
type Rule struct {
rule.BaseRule
PasvAddress string `gorm:"pasv_address" json:"pasv_address" form:"pasv_address"`
Data []byte `json:"data" form:"data"`
rule.BaseRule `yaml:",inline"`
PasvAddress string `gorm:"pasv_address" json:"pasv_address" form:"pasv_address" yaml:"pasv_address"`
Data []byte `json:"data" form:"data"`
}
func (Rule) TableName() string {
@@ -71,16 +71,25 @@ func (r *Rule) Delete() (err error) {
// ListRules lists all ftp rules those satisfy the filter
func ListRules(c *gin.Context) {
var (
ftpRule Rule
res []Rule
count int64
order = c.Query("order")
ftpRule Rule
res []Rule
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&ftpRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -96,7 +105,7 @@ func ListRules(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -106,10 +115,10 @@ func ListRules(c *gin.Context) {
order = "desc"
}
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -122,7 +131,7 @@ func ListRules(c *gin.Context) {
})
}
// Create or update ftp rule from user submit
// UpsertRules creates or updates ftp rule from user submit
func UpsertRules(c *gin.Context) {
var (
ftpRule Rule
@@ -132,7 +141,7 @@ func UpsertRules(c *gin.Context) {
if err := c.ShouldBind(&ftpRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -145,16 +154,16 @@ func UpsertRules(c *gin.Context) {
if err := ftpRule.CreateOrUpdate(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
}
if update {
log.Trace("FTP rule[id%d] has been updated", ftpRule.ID)
log.Trace("FTP rule[id:%d] has been updated", ftpRule.ID)
} else {
log.Trace("FTP rule[id%d] has been created", ftpRule.ID)
log.Trace("FTP rule[id:%d] has been created", ftpRule.ID)
}
c.JSON(200, gin.H{
@@ -164,14 +173,14 @@ func UpsertRules(c *gin.Context) {
})
}
// Delete ftp rule from user submit
// DeleteRules deletes ftp rule from user submit
func DeleteRules(c *gin.Context) {
var ftpRule Rule
if err := c.ShouldBind(&ftpRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -180,13 +189,13 @@ func DeleteRules(c *gin.Context) {
if err := ftpRule.Delete(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
log.Trace("FTP rule[id%d] has been deleted", ftpRule.ID)
log.Trace("FTP rule[id:%d] has been deleted", ftpRule.ID)
c.JSON(200, gin.H{
"status": "succeed",
+40 -11
View File
@@ -3,7 +3,6 @@ package mysql
import (
"encoding/base64"
"fmt"
"os"
"regexp"
"strings"
"sync"
@@ -13,6 +12,7 @@ import (
"github.com/li4n0/revsuit/internal/file"
"github.com/li4n0/revsuit/internal/qqwry"
"github.com/li4n0/revsuit/pkg/mysql/vmysql"
"github.com/pkg/errors"
log "unknwon.dev/clog/v2"
"vitess.io/vitess/go/sqltypes"
)
@@ -25,8 +25,9 @@ var (
type Server struct {
Config
rules []*Rule
rulesLock sync.RWMutex
rules []*Rule
rulesLock sync.RWMutex
livingLock sync.Mutex
listener *vmysql.Listener
Handler vmysql.Handler
@@ -36,7 +37,7 @@ type Server struct {
func GetServer() *Server {
once.Do(func() {
server = &Server{rulesLock: sync.RWMutex{}}
server = &Server{rulesLock: sync.RWMutex{}, livingLock: sync.Mutex{}}
})
return server
}
@@ -51,7 +52,7 @@ func (s *Server) updateRules() error {
db := database.DB.Model(new(Rule))
defer s.rulesLock.Unlock()
s.rulesLock.Lock()
return db.Order("rank desc").Find(&s.rules).Error
return errors.Wrap(db.Order("rank desc").Find(&s.rules).Error, "MySQL update rules error")
}
// NewConnection is part of the mysql.Handler interface.
@@ -146,14 +147,14 @@ func (s *Server) ConnectionClosed(c *vmysql.Conn) {
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and domain like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
database.DB.Where("rule_name=? and (user like ? or schema like ?)", _rule.Name, "%"+flagGroup+"%", "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("MySQL record[id%d] has been put to client message queue", r.ID)
log.Trace("MySQL record[id:%d, flagGroup:%s] has been put to client message queue", r.ID, flagGroup)
}
} else {
r.PushToClient()
log.Trace("MySQL record[id%d] has been put to client message queue", r.ID)
log.Trace("MySQL record[id:%d, flag:%s] has been put to client message queue", r.ID, flag)
}
}
@@ -266,9 +267,31 @@ func (s *Server) WarningCount(c *vmysql.Conn) uint16 {
return 0
}
func (s *Server) Stop() {
log.Info("MySQL Server is stopping...")
s.Enable = false
s.livingLock.Unlock()
}
func (s *Server) Restart() {
s.Enable = false
time.Sleep(time.Second * 2)
go s.Run()
}
func (s *Server) Run() {
s.Enable = true
s.livingLock.Lock()
defer func() {
if s.Enable {
log.Error("MySQL Server exited unexpectedly")
}
s.Enable = false
s.livingLock.Unlock()
}()
if err := s.updateRules(); err != nil {
log.Fatal(err.Error())
log.Error(err.Error())
return
}
s.Handler = s
@@ -279,9 +302,15 @@ func (s *Server) Run() {
log.Info("Starting MySQL Server at %s", s.Addr)
s.listener, err = vmysql.NewListener("tcp", s.Addr, authServer, s, s.VersionString, 0, 0)
if err != nil {
log.Warn("New MySQL Server failed: %s", err)
os.Exit(-1)
log.Error("New MySQL Server failed: %s", err)
}
go func() {
if !s.Enable {
s.listener.Close()
}
}()
s.listener.Accept()
}
+13 -4
View File
@@ -55,12 +55,21 @@ func ListRecords(c *gin.Context) {
res []Record
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&mysqlRecord); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
}
@@ -93,7 +102,7 @@ func ListRecords(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -103,10 +112,10 @@ func ListRecords(c *gin.Context) {
order = "desc"
}
if err := db.Preload("Files").Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Preload("Files").Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
+27 -18
View File
@@ -11,9 +11,9 @@ import (
)
type Rule struct {
rule.BaseRule
rule.BaseRule `yaml:",inline"`
Files string `form:"files" json:"files"`
ExploitJdbcClient bool `gorm:"exploit_jdbc_client" form:"exploit_jdbc_client" json:"exploit_jdbc_client"`
ExploitJdbcClient bool `gorm:"exploit_jdbc_client" form:"exploit_jdbc_client" json:"exploit_jdbc_client" yaml:"exploit_jdbc_client"`
Payloads database.MapField `json:"payloads" form:"payloads"`
}
@@ -21,7 +21,7 @@ func (Rule) TableName() string {
return "mysql_rules"
}
// Create or update the mysql rule in database and ruleSet
// CreateOrUpdate creates or updates the mysql rule in database and ruleSet
func (r *Rule) CreateOrUpdate() (err error) {
db := database.DB.Model(r)
err = db.Clauses(clause.OnConflict{
@@ -45,7 +45,7 @@ func (r *Rule) CreateOrUpdate() (err error) {
return err
}
// Delete the mysql rule in database and ruleSet
// Delete deletes the mysql rule in database and ruleSet
func (r *Rule) Delete() (err error) {
db := database.DB.Model(r)
err = db.Delete(r).Error
@@ -56,19 +56,28 @@ func (r *Rule) Delete() (err error) {
return err
}
// List all mysql rules those satisfy the filter
// ListRules lists all mysql rules those satisfy the filter
func ListRules(c *gin.Context) {
var (
mysqlRule Rule
res []Rule
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&mysqlRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -84,7 +93,7 @@ func ListRules(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -94,10 +103,10 @@ func ListRules(c *gin.Context) {
order = "desc"
}
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -110,7 +119,7 @@ func ListRules(c *gin.Context) {
})
}
// Create or update mysql rule from user submit
// UpsertRules create or update mysql rule from user submit
func UpsertRules(c *gin.Context) {
var (
mysqlRule Rule
@@ -120,7 +129,7 @@ func UpsertRules(c *gin.Context) {
if err := c.ShouldBind(&mysqlRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -133,16 +142,16 @@ func UpsertRules(c *gin.Context) {
if err := mysqlRule.CreateOrUpdate(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
if update {
log.Trace("MySQL rule[id%d] has been updated", mysqlRule.ID)
log.Trace("MySQL rule[id:%d] has been updated", mysqlRule.ID)
} else {
log.Trace("MySQL rule[id%d] has been created", mysqlRule.ID)
log.Trace("MySQL rule[id:%d] has been created", mysqlRule.ID)
}
c.JSON(200, gin.H{
@@ -152,14 +161,14 @@ func UpsertRules(c *gin.Context) {
})
}
// Delete mysql rule from user submit
// DeleteRules Delete mysql rule from user submit
func DeleteRules(c *gin.Context) {
var mysqlRule Rule
if err := c.ShouldBind(&mysqlRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -168,13 +177,13 @@ func DeleteRules(c *gin.Context) {
if err := mysqlRule.Delete(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
log.Trace("MySQL rule[id%d] has been deleted", mysqlRule.ID)
log.Trace("MySQL rule[id:%d] has been deleted", mysqlRule.ID)
c.JSON(200, gin.H{
"status": "succeed",
+12 -12
View File
@@ -259,7 +259,7 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
// Catch panics, and close the connection in any case.
defer func() {
if x := recover(); x != nil {
log.Warn("mysql_server caught panic:\n%v\n%s", x, tb.Stack(4))
log.Trace("mysql_server caught panic:\n%v\n%s", x, tb.Stack(4))
}
// We call flush here in case there's a premature return after
// startWriterBuffering is called
@@ -275,7 +275,7 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
salt, err := c.writeHandshakeV10(l.ServerVersion, l.authServer, l.TLSConfig != nil)
if err != nil {
if err != io.EOF {
log.Warn("Cannot send HandshakeV10 packet to %s: %v", c, err)
log.Trace("Cannot send HandshakeV10 packet to %s: %v", c, err)
}
return
}
@@ -286,13 +286,13 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
if err != nil {
// Don't log EOF errors. They cause too much spam, same as main read loop.
if err != io.EOF {
log.Warn("Cannot read client handshake response from %s: %v", c, err)
log.Trace("Cannot read client handshake response from %s: %v", c, err)
}
return
}
user, authMethod, authResponse, err := l.parseClientHandshakePacket(c, true, response)
if err != nil {
log.Warn("Cannot parse client handshake response from %s: %v", c, err)
log.Trace("Cannot parse client handshake response from %s: %v", c, err)
return
}
@@ -305,14 +305,14 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
// SSL was enabled. We need to re-read the auth packet.
response, err = c.readEphemeralPacket()
if err != nil {
log.Warn("Cannot read post-SSL client handshake response from %s: %v", c, err)
log.Trace("Cannot read post-SSL client handshake response from %s: %v", c, err)
return
}
// Returns copies of the data, so we can recycle the buffer.
user, authMethod, authResponse, err = l.parseClientHandshakePacket(c, false, response)
if err != nil {
log.Warn("Cannot parse post-SSL client handshake response from %s: %v", c, err)
log.Trace("Cannot parse post-SSL client handshake response from %s: %v", c, err)
return
}
c.RecycleReadPacket()
@@ -366,13 +366,13 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
data := make([]byte, 21) //nolint:ineffassign,staticcheck // SA4006 This line is required because the binary protocol requires padding with 0
data = append(salt, byte(0x00))
if err := c.writeAuthSwitchRequest(MysqlNativePassword, data); err != nil {
log.Warn("Error writing auth switch packet for %s: %v", c, err)
log.Trace("Error writing auth switch packet for %s: %v", c, err)
return
}
response, err := c.readEphemeralPacket()
if err != nil {
log.Warn("Error reading auth switch response for %s: %v", c, err)
log.Trace("Error reading auth switch response for %s: %v", c, err)
return
}
c.RecycleReadPacket()
@@ -402,7 +402,7 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
data = authServerDialogSwitchData()
}
if err := c.writeAuthSwitchRequest(authServerMethod, data); err != nil {
log.Warn("Error writing auth switch packet for %s: %v", c, err)
log.Trace("Error writing auth switch packet for %s: %v", c, err)
return
}
@@ -424,7 +424,7 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
// Negotiation worked, send OK packet.
if err := c.writeOKPacket(0, 0, c.StatusFlags, 0); err != nil {
log.Warn("Cannot write OK packet to %s: %v", c, err)
log.Trace("Cannot write OK packet to %s: %v", c, err)
return
}
@@ -435,7 +435,7 @@ func (l *Listener) handle(conn net.Conn, connectionID uint32, acceptTime time.Ti
connectTime := time.Since(acceptTime)
if l.SlowConnectWarnThreshold != 0 && connectTime > l.SlowConnectWarnThreshold {
connSlow.Add(1)
log.Warn("Slow connection from %s: %v", c, connectTime)
log.Trace("Slow connection from %s: %v", c, connectTime)
}
for {
@@ -692,7 +692,7 @@ func (l *Listener) parseClientHandshakePacket(c *Conn, firstTime bool, data []by
// Decode connection attributes send by the client
if clientFlags&CapabilityClientConnAttr != 0 {
if connAttrs, _, err := parseConnAttrs(data, pos); err != nil {
log.Warn("Decode connection attributes send by the client: %v", err)
log.Trace("Decode connection attributes send by the client: %v", err)
} else {
c.ConnAttrs = connAttrs
}
+38 -9
View File
@@ -1,6 +1,7 @@
package rhttp
import (
"context"
"math/rand"
"net/http"
"net/http/httputil"
@@ -8,6 +9,7 @@ import (
"strconv"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/li4n0/revsuit/internal/database"
@@ -23,6 +25,8 @@ type Server struct {
ApiGroup *gin.RouterGroup
rules []*Rule
rulesLock sync.RWMutex
httpServer http.Server
}
const (
@@ -86,14 +90,38 @@ func (s *Server) updateRules() error {
return db.Order("rank desc").Find(&s.rules).Error
}
func (s *Server) startHttpServer() {
log.Info("Starting HTTP Server at %s, token:%s", s.Addr, s.Token)
s.httpServer = http.Server{
Addr: s.Addr,
Handler: s.Router,
}
err := s.httpServer.ListenAndServe()
if err != nil && err != http.ErrServerClosed {
log.Fatal(err.Error())
}
}
func (s *Server) stopHttpServer() {
log.Info("HTTP Server is stopping...")
err := s.httpServer.Shutdown(context.TODO())
if err != nil {
log.Fatal(err.Error())
}
}
func (s *Server) Restart() {
//only need to stop http server, because it will start in a loop
s.stopHttpServer()
}
func (s *Server) Run() {
if err := s.updateRules(); err != nil {
log.Warn(err.Error())
}
log.Info("Starting HTTP Server at %s, token:%s", s.Addr, s.Token)
err := s.Router.Run(s.Addr)
if err != nil {
log.Fatal(err.Error())
for {
s.startHttpServer()
time.Sleep(2 * time.Second)
}
}
@@ -167,21 +195,22 @@ func (s *Server) Receive(c *gin.Context) {
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and raw like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
database.DB.Where("rule_name=? and raw_request like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("HTTP record[id%d] has been put to client message queue", r.ID)
log.Trace("HTTP record[id:%d, flagGroup:%s] has been put to client message queue", r.ID, flagGroup)
}
} else {
r.PushToClient()
log.Trace("HTTP record[id:%d, flag:%s] has been put to client message queue", r.ID, r.Flag)
}
r.PushToClient()
log.Trace("HTTP record[id%d] has been put to client message queue", r.ID)
}
//send notice
if _rule.Notice {
go func() {
r.Notice()
log.Trace("HTTP record[id%d] notice has been sent", r.ID)
log.Trace("HTTP record[id:%d] notice has been sent", r.ID)
}()
}
+14 -4
View File
@@ -50,12 +50,22 @@ func ListRecords(c *gin.Context) {
res []Record
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
pageSize = n
if pageSize <= 0 || pageSize > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&httpRecord); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -82,7 +92,7 @@ func ListRecords(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -92,10 +102,10 @@ func ListRecords(c *gin.Context) {
order = "desc"
}
if err := db.Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
+31 -22
View File
@@ -10,19 +10,19 @@ import (
log "unknwon.dev/clog/v2"
)
// Http rule struct
// Rule Http rule struct
type Rule struct {
rule.BaseRule
ResponseStatusCode string `gorm:"index;default:200;not null" form:"response_status_code" json:"response_status_code"`
ResponseHeaders database.MapField `form:"response_headers" json:"response_headers"`
ResponseBody string `gorm:"default:Hello RevSuit!" form:"response_body" json:"response_body"`
rule.BaseRule `yaml:",inline"`
ResponseStatusCode string `gorm:"index;default:200;not null" form:"response_status_code" json:"response_status_code" yaml:"response_status_code"`
ResponseHeaders database.MapField `form:"response_headers" json:"response_headers" yaml:"response_headers"`
ResponseBody string `gorm:"default:Hello RevSuit!" form:"response_body" json:"response_body" yaml:"response_body"`
}
func (Rule) TableName() string {
return "http_rules"
}
// New http rule struct
// NewRule new http rule struct
func NewRule(name, flagFormat, responseBody string, pushToClient, notice bool, responseStatus string, responseHeaders database.MapField) *Rule {
return &Rule{
BaseRule: rule.BaseRule{
@@ -37,7 +37,7 @@ func NewRule(name, flagFormat, responseBody string, pushToClient, notice bool, r
}
}
// Create or update the http rule in database and ruleSet
// CreateOrUpdate creates or updates the http rule in database and ruleSet
func (r *Rule) CreateOrUpdate() (err error) {
db := database.DB.Model(r)
err = db.Clauses(clause.OnConflict{
@@ -62,7 +62,7 @@ func (r *Rule) CreateOrUpdate() (err error) {
return err
}
// Delete the http rule in database and ruleSet
// Delete deletes the http rule in database and ruleSet
func (r *Rule) Delete() (err error) {
db := database.DB.Model(r)
err = db.Delete(r).Error
@@ -74,19 +74,28 @@ func (r *Rule) Delete() (err error) {
return err
}
// List all http rules those satisfy the filter
// ListRules lists all http rules those satisfy the filter
func ListRules(c *gin.Context) {
var (
httpRule Rule
res []Rule
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&httpRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -102,7 +111,7 @@ func ListRules(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -112,10 +121,10 @@ func ListRules(c *gin.Context) {
order = "desc"
}
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -128,7 +137,7 @@ func ListRules(c *gin.Context) {
})
}
// Create or update http rule from user submit
// UpsertRules create or update http rule from user submit
func UpsertRules(c *gin.Context) {
var (
httpRule Rule
@@ -138,7 +147,7 @@ func UpsertRules(c *gin.Context) {
if err := c.ShouldBind(&httpRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -151,16 +160,16 @@ func UpsertRules(c *gin.Context) {
if err := httpRule.CreateOrUpdate(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
}
if update {
log.Trace("HTTP rule[id%d] has been updated", httpRule.ID)
log.Trace("HTTP rule[id:%d] has been updated", httpRule.ID)
} else {
log.Trace("HTTP rule[id%d] has been created", httpRule.ID)
log.Trace("HTTP rule[id:%d] has been created", httpRule.ID)
}
c.JSON(200, gin.H{
@@ -170,14 +179,14 @@ func UpsertRules(c *gin.Context) {
})
}
// Delete http rule from user submit
// DeleteRules deletes http rule from user submit
func DeleteRules(c *gin.Context) {
var httpRule Rule
if err := c.ShouldBind(&httpRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -186,13 +195,13 @@ func DeleteRules(c *gin.Context) {
if err := httpRule.Delete(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
log.Trace("HTTP rule[id%d] has been deleted", httpRule.ID)
log.Trace("HTTP rule[id:%d] has been deleted", httpRule.ID)
c.JSON(200, gin.H{
"status": "succeed",
+13 -4
View File
@@ -46,12 +46,21 @@ func ListRecords(c *gin.Context) {
res []Record
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&rmiRecord); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -75,7 +84,7 @@ func ListRecords(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -85,10 +94,10 @@ func ListRecords(c *gin.Context) {
order = "desc"
}
if err := db.Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
+51 -14
View File
@@ -12,13 +12,15 @@ import (
"github.com/li4n0/revsuit/internal/database"
"github.com/li4n0/revsuit/internal/qqwry"
"github.com/li4n0/revsuit/internal/recycler"
"github.com/pkg/errors"
log "unknwon.dev/clog/v2"
)
type Server struct {
Config
rules []*Rule
rulesLock sync.RWMutex
rules []*Rule
rulesLock sync.RWMutex
livingLock sync.Mutex
}
var (
@@ -28,7 +30,7 @@ var (
func GetServer() *Server {
once.Do(func() {
server = &Server{rulesLock: sync.RWMutex{}}
server = &Server{rulesLock: sync.RWMutex{}, livingLock: sync.Mutex{}}
})
return server
}
@@ -43,7 +45,7 @@ func (s *Server) updateRules() error {
db := database.DB.Model(new(Rule))
defer s.rulesLock.Unlock()
s.rulesLock.Lock()
return db.Order("rank desc").Find(&s.rules).Error
return errors.Wrap(db.Order("rank desc").Find(&s.rules).Error, "RMI update rules error")
}
func (s *Server) handleConnection(conn net.Conn) {
@@ -119,29 +121,53 @@ func (s *Server) handleConnection(conn net.Conn) {
if _rule.PushToClient {
if flagGroup != "" {
var count int64
database.DB.Where("rule_name=? and raw like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
database.DB.Where("rule_name=? and path like ?", _rule.Name, "%"+flagGroup+"%").Model(&Record{}).Count(&count)
if count <= 1 {
r.PushToClient()
log.Trace("RMI record[id%d] has been put to client message queue", r.ID)
log.Trace("RMI record[id:%d, flagGroup:%s] has been put to client message queue", r.ID, flagGroup)
}
} else {
r.PushToClient()
log.Trace("RMI record[id:%d, flag:%s] has been put to client message queue", r.ID, flag)
}
r.PushToClient()
log.Trace("RMI record[id%d] has been put to client message queue", r.ID)
}
//send notice
if _rule.Notice {
go func() {
r.Notice()
log.Trace("RMI record[id%d] notice has been sent", r.ID)
log.Trace("RMI record[id:%d] notice has been sent", r.ID)
}()
}
}
}
func (s *Server) Stop() {
log.Info("RMI Server is stopping...")
s.Enable = false
s.livingLock.Unlock()
}
func (s *Server) Restart() {
s.Enable = false
time.Sleep(time.Second * 2)
go s.Run()
}
func (s *Server) Run() {
s.Enable = true
s.livingLock.Lock()
defer func() {
if s.Enable {
log.Error("RMI Server exited unexpectedly")
}
s.Enable = false
s.livingLock.Unlock()
}()
if err := s.updateRules(); err != nil {
log.Fatal(err.Error())
log.Error(err.Error())
return
}
// run server
@@ -149,16 +175,27 @@ func (s *Server) Run() {
listener, err := net.Listen("tcp", s.Addr)
if err != nil {
log.Fatal(err.Error())
log.Error(err.Error())
return
}
for {
go func() {
s.livingLock.Lock()
if !s.Enable {
_ = listener.Close()
}
}()
for s.Enable {
tcpConn, err := listener.Accept()
if err != nil {
log.Warn("RMI accept connection error: %v", err)
if !strings.Contains(err.Error(), net.ErrClosed.Error()) {
log.Warn("RMI accept connection error: %v", err)
} else {
break
}
continue
}
go s.handleConnection(tcpConn)
}
}
+32 -23
View File
@@ -10,16 +10,16 @@ import (
log "unknwon.dev/clog/v2"
)
// RMI rule struct
// Rule RMI rule struct
type Rule struct {
rule.BaseRule
rule.BaseRule `yaml:",inline"`
}
func (Rule) TableName() string {
return "rmi_rules"
}
// New rmi rule struct
// NewRule new rmi rule struct
func NewRule(name, flagFormat string, pushToClient, notice bool) *Rule {
return &Rule{
BaseRule: rule.BaseRule{
@@ -31,7 +31,7 @@ func NewRule(name, flagFormat string, pushToClient, notice bool) *Rule {
}
}
// Create or update the rmi rule in database and ruleSet
// CreateOrUpdate creates or updates the rmi rule in database and ruleSet
func (r *Rule) CreateOrUpdate() (err error) {
db := database.DB.Model(r)
err = db.Clauses(clause.OnConflict{
@@ -53,7 +53,7 @@ func (r *Rule) CreateOrUpdate() (err error) {
return err
}
// Delete the rmi rule in database and ruleSet
// Delete deletes the rmi rule in database and ruleSet
func (r *Rule) Delete() (err error) {
db := database.DB.Model(r)
err = db.Delete(r).Error
@@ -65,19 +65,28 @@ func (r *Rule) Delete() (err error) {
return err
}
// List all rmi rules those satisfy the filter
// ListRules lists all rmi rules those satisfy the filter
func ListRules(c *gin.Context) {
var (
rmiRule Rule
res []Rule
count int64
order = c.Query("order")
rmiRule Rule
res []Rule
count int64
order = c.Query("order")
pageSize int
)
if c.Query("pageSize") == "" {
pageSize = 10
} else if n, err := strconv.Atoi(c.Query("pageSize")); err == nil {
if n <= 0 || n > 100 {
pageSize = 10
}
}
if err := c.ShouldBind(&rmiRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -93,7 +102,7 @@ func ListRules(c *gin.Context) {
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
@@ -103,10 +112,10 @@ func ListRules(c *gin.Context) {
order = "desc"
}
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * 10).Limit(10).Find(&res).Error; err != nil {
if err := db.Order("rank desc").Order("id" + " " + order).Count(&count).Offset((page - 1) * pageSize).Limit(pageSize).Find(&res).Error; err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -119,7 +128,7 @@ func ListRules(c *gin.Context) {
})
}
// Create or update rmi rule from user submit
// UpsertRules creates or updates rmi rule from user submit
func UpsertRules(c *gin.Context) {
var (
rmiRule Rule
@@ -129,7 +138,7 @@ func UpsertRules(c *gin.Context) {
if err := c.ShouldBind(&rmiRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -142,16 +151,16 @@ func UpsertRules(c *gin.Context) {
if err := rmiRule.CreateOrUpdate(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"result": nil,
})
return
}
if update {
log.Trace("RMI rule[id%d] has been updated", rmiRule.ID)
log.Trace("RMI rule[id:%d] has been updated", rmiRule.ID)
} else {
log.Trace("RMI rule[id%d] has been created", rmiRule.ID)
log.Trace("RMI rule[id:%d] has been created", rmiRule.ID)
}
c.JSON(200, gin.H{
@@ -161,14 +170,14 @@ func UpsertRules(c *gin.Context) {
})
}
// Delete rmi rule from user submit
// DeleteRules deletes rmi rule from user submit
func DeleteRules(c *gin.Context) {
var rmiRule Rule
if err := c.ShouldBind(&rmiRule); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
@@ -177,13 +186,13 @@ func DeleteRules(c *gin.Context) {
if err := rmiRule.Delete(); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err,
"error": err.Error(),
"data": nil,
})
return
}
log.Trace("RMI rule[id%d] has been deleted", rmiRule.ID)
log.Trace("RMI rule[id:%d] has been deleted", rmiRule.ID)
c.JSON(200, gin.H{
"status": "succeed",
+1 -1
View File
@@ -19,7 +19,7 @@ type Config struct {
Addr string
Token string
Database string
LogLevel string
LogLevel string `yaml:"log_level"`
Notice noticeConfig
rhttp.Config
DNS dns.Config
+10 -3
View File
@@ -27,9 +27,8 @@ func ping(c *gin.Context) {
}
func events(c *gin.Context) {
log.Info("Receive connection from %v", c.Request.RemoteAddr)
log.Info("Receive client connection from %v", c.Request.RemoteAddr)
c.Stream(func(w io.Writer) bool {
c.SSEvent("message", "connect succeed")
select {
case <-c.Writer.CloseNotify():
return false
@@ -38,7 +37,7 @@ func events(c *gin.Context) {
}
return true
})
log.Info(c.Request.RemoteAddr, "disconnect")
log.Info("Client %s disconnect", c.Request.RemoteAddr)
}
func recovery(c *gin.Context) {
@@ -85,3 +84,11 @@ func recovery(c *gin.Context) {
}()
c.Next()
}
func version(c *gin.Context) {
c.JSON(200, gin.H{
"status": "succeed",
"error": nil,
"result": VERSION,
})
}
+18
View File
@@ -42,6 +42,7 @@ func (revsuit *Revsuit) registerPlatformRouter() {
api.GET("/auth", auth)
api.GET("/events", events)
api.GET("/ping", ping)
api.GET("/version", version)
}
func (revsuit *Revsuit) registerHttpRouter() {
@@ -54,6 +55,23 @@ func (revsuit *Revsuit) registerHttpRouter() {
}
revsuit.http.Router.StaticFS("/revsuit/admin", http.FS(fe))
// init settings router group
settingsGroup := revsuit.http.ApiGroup.Group("setting")
settingsGroup.GET("/exportRules", exportRules)
settingsGroup.POST("/importRules", importRules)
settingsGroup.GET("/getHttpConfig", revsuit.getHttpConfig)
settingsGroup.POST("/updateHttpConfig", revsuit.updateHttpConfig)
settingsGroup.GET("/getDnsConfig", revsuit.getDnsConfig)
settingsGroup.POST("/updateDnsConfig", revsuit.updateDnsConfig)
settingsGroup.GET("/getRmiConfig", revsuit.getRmiConfig)
settingsGroup.POST("/updateRmiConfig", revsuit.updateRmiConfig)
settingsGroup.GET("/getMySQLConfig", revsuit.getMySQLConfig)
settingsGroup.POST("/updateMySQLConfig", revsuit.updateMySQLConfig)
settingsGroup.GET("/getFtpConfig", revsuit.getFtpConfig)
settingsGroup.POST("/updateFtpConfig", revsuit.updateFtpConfig)
settingsGroup.GET("/getNoticeConfig", revsuit.getNoticeConfig)
settingsGroup.POST("/updateNoticeConfig", revsuit.updateNoticeConfig)
// init record router group
recordGroup := revsuit.http.ApiGroup.Group("/record")
+32 -22
View File
@@ -14,7 +14,10 @@ import (
log "unknwon.dev/clog/v2"
)
const VERSION = "Beta0.1"
type Revsuit struct {
config *Config
logLevel log.Level
http *http.Server
@@ -25,6 +28,11 @@ type Revsuit struct {
}
func initDatabase(dsn string) {
_ = log.NewConsole(100,
log.ConsoleConfig{
Level: log.LevelInfo,
})
err := database.InitDB("sqlite", dsn)
if err != nil {
log.Fatal(err.Error())
@@ -104,6 +112,10 @@ func initLog(level string) (logLevel log.Level) {
gin.SetMode(gin.ReleaseMode)
database.DB.Logger.LogMode(logger.Error)
logLevel = log.LevelFatal
default:
gin.SetMode(gin.DebugMode)
database.DB.Logger.LogMode(logger.Info)
logLevel = log.LevelInfo
}
_ = log.NewConsole(100,
log.ConsoleConfig{
@@ -138,29 +150,27 @@ func initNotice(nc noticeConfig) {
func New(c *Config) *Revsuit {
logLevel := initLog(c.LogLevel)
initDatabase(c.Database)
logLevel := initLog(c.LogLevel)
initNotice(c.Notice)
s := &Revsuit{
config: c,
logLevel: logLevel,
http: http.GetServer(),
}
if c.DNS.Enable {
s.dns = dns.GetServer()
}
if c.MySQL.Enable {
s.mysql = mysql.GetServer()
s.mysql.Config = c.MySQL
}
if c.RMI.Enable {
s.rmi = rmi.GetServer()
s.rmi.Config = c.RMI
}
if c.FTP.Enable {
s.ftp = ftp.GetServer()
s.ftp.Config = c.FTP
}
s.dns = dns.GetServer()
s.dns.Config = c.DNS
s.mysql = mysql.GetServer()
s.mysql.Config = c.MySQL
s.rmi = rmi.GetServer()
s.rmi.Config = c.RMI
s.ftp = ftp.GetServer()
s.ftp.Config = c.FTP
if c.Addr != "" {
s.http.SetAddr(c.Addr)
@@ -178,16 +188,16 @@ func (revsuit *Revsuit) Run() {
defer log.Stop()
revsuit.registerRouter()
if revsuit.dns != nil {
if revsuit.dns != nil && revsuit.dns.Enable {
go revsuit.dns.Run()
}
if revsuit.mysql != nil {
go revsuit.mysql.Run()
}
if revsuit.rmi != nil {
if revsuit.rmi != nil && revsuit.rmi.Enable {
go revsuit.rmi.Run()
}
if revsuit.ftp != nil {
if revsuit.mysql != nil && revsuit.mysql.Enable {
go revsuit.mysql.Run()
}
if revsuit.ftp != nil && revsuit.ftp.Enable {
go revsuit.ftp.Run()
}
+367
View File
@@ -0,0 +1,367 @@
package server
import (
"fmt"
"io"
"time"
"github.com/gin-gonic/gin"
"github.com/li4n0/revsuit/internal/database"
"github.com/li4n0/revsuit/pkg/dns"
"github.com/li4n0/revsuit/pkg/ftp"
"github.com/li4n0/revsuit/pkg/mysql"
"github.com/li4n0/revsuit/pkg/rhttp"
"github.com/li4n0/revsuit/pkg/rmi"
"github.com/pkg/errors"
"gopkg.in/yaml.v3"
log "unknwon.dev/clog/v2"
)
type Rules struct {
Http []rhttp.Rule
Dns []dns.Rule
Mysql []mysql.Rule
Rmi []rmi.Rule
Ftp []ftp.Rule
}
func exportRules(c *gin.Context) {
var (
db = database.DB
rules Rules
)
db.Model(&rhttp.Rule{}).Find(&rules.Http)
db.Model(&dns.Rule{}).Find(&rules.Dns)
db.Model(&mysql.Rule{}).Find(&rules.Mysql)
db.Model(&rmi.Rule{}).Find(&rules.Rmi)
db.Model(&ftp.Rule{}).Find(&rules.Ftp)
out, err := yaml.Marshal(rules)
if err != nil {
log.Warn("export rules error: %s", err)
}
c.Header("Content-Disposition", fmt.Sprintf("attachment;filename=revsuit_rules_%s.yaml", time.Now().Format("20060102150405")))
c.String(200, string(out))
}
func importRules(c *gin.Context) {
var (
db = database.DB
rules Rules
count int
errs []string
)
f, err := c.FormFile("rules")
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"result": nil,
})
log.Trace("%v", err)
return
}
file, _ := f.Open()
content, err := io.ReadAll(file)
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"result": nil,
})
log.Trace("%v", err)
return
}
err = yaml.Unmarshal(content, &rules)
if err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"result": nil,
})
return
}
for _, rule := range rules.Http {
err := db.Model(&rhttp.Rule{}).Create(&rule).Error
if err != nil {
errs = append(errs, errors.Wrap(err, fmt.Sprintf("http rule[%s]", rule.Name)).Error())
continue
}
count += 1
}
for _, rule := range rules.Dns {
err := db.Model(&dns.Rule{}).Create(&rule).Error
if err != nil {
errs = append(errs, errors.Wrap(err, fmt.Sprintf("dns rule[%s]", rule.Name)).Error())
continue
}
count += 1
}
for _, rule := range rules.Mysql {
err := db.Model(&mysql.Rule{}).Create(&rule).Error
if err != nil {
errs = append(errs, errors.Wrap(err, fmt.Sprintf("mysql rule[%s]", rule.Name)).Error())
continue
}
count += 1
}
for _, rule := range rules.Rmi {
err := db.Model(&rmi.Rule{}).Create(&rule).Error
if err != nil {
errs = append(errs, errors.Wrap(err, fmt.Sprintf("rmi rule[%s]", rule.Name)).Error())
continue
}
count += 1
}
for _, rule := range rules.Ftp {
err := db.Model(&ftp.Rule{}).Create(&rule).Error
if err != nil {
errs = append(errs, errors.Wrap(err, fmt.Sprintf("ftp rule[%s]", rule.Name)).Error())
continue
}
count += 1
}
c.JSON(200, gin.H{
"status": "succeed",
"error": errs,
"result": fmt.Sprintf("%d rules were imported successfully, %d failed.", count, len(errs)),
})
}
func (revsuit *Revsuit) getHttpConfig(c *gin.Context) {
var res = make(map[string]string)
res["Addr"] = revsuit.config.Addr
res["Token"] = revsuit.config.Token
res["Database"] = revsuit.config.Database
res["LogLevel"] = revsuit.config.LogLevel
res["IpHeader"] = revsuit.config.IpHeader
c.JSON(200, res)
}
func (revsuit *Revsuit) updateHttpConfig(c *gin.Context) {
var form = make(map[string]string)
if err := c.ShouldBindJSON(&form); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"data": nil,
})
return
}
if form["LogLevel"] != revsuit.config.LogLevel {
revsuit.logLevel = initLog(form["LogLevel"])
revsuit.config.LogLevel = form["LogLevel"]
if revsuit.logLevel == log.LevelTrace {
revsuit.http.Router.Use(gin.Logger())
}
log.Info("Update platform config [log_level] to %s", form["LogLevel"])
}
if form["Token"] != revsuit.config.Token {
revsuit.config.Token = form["Token"]
revsuit.http.SetToken(form["Token"])
log.Info("Update platform config [token] to %s", form["Token"])
}
if form["Database"] != revsuit.config.Database {
revsuit.config.Database = form["Database"]
initDatabase(form["Database"])
log.Info("Update platform config [database] to %s", form["Database"])
}
if form["IpHeader"] != revsuit.config.IpHeader {
revsuit.config.IpHeader = form["IpHeader"]
revsuit.http.SetIpHeader(form["IpHeader"])
log.Info("Update http config [ip_header] to %s", form["IpHeader"])
}
c.JSON(200, gin.H{
"status": "succeed",
"error": nil,
"result": "update succeed",
})
}
func (revsuit *Revsuit) getFtpConfig(c *gin.Context) {
c.JSON(200, revsuit.ftp.Config)
}
func (revsuit *Revsuit) updateFtpConfig(c *gin.Context) {
var form = ftp.Config{}
if err := c.ShouldBindJSON(&form); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"data": nil,
})
return
}
if form.Addr != revsuit.ftp.Addr {
revsuit.ftp.Addr = form.Addr
log.Info("Update ftp config [addr] to %s", form.Addr)
}
if form.PasvIP != revsuit.ftp.PasvIP {
revsuit.ftp.PasvIP = form.PasvIP
log.Info("Update ftp config [pasv_ip] to %s", form.PasvIP)
}
if form.PasvPort != revsuit.ftp.PasvPort {
revsuit.ftp.PasvPort = form.PasvPort
log.Info("Update ftp config [pasv_port] to %d", form.PasvPort)
}
if form.Enable != revsuit.ftp.Enable {
log.Info("Update ftp config [enable] to %v", form.Enable)
if form.Enable {
go revsuit.ftp.Run()
} else {
revsuit.ftp.Stop()
}
return
}
if revsuit.ftp.Enable {
revsuit.ftp.Restart()
}
}
func (revsuit *Revsuit) getDnsConfig(c *gin.Context) {
c.JSON(200, revsuit.dns.Config)
}
func (revsuit *Revsuit) updateDnsConfig(c *gin.Context) {
var form = dns.Config{}
if err := c.ShouldBindJSON(&form); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"data": nil,
})
return
}
if form.Enable != revsuit.dns.Enable {
log.Info("Update dns config [enable] to %v", form.Enable)
if form.Enable {
go revsuit.dns.Run()
} else {
revsuit.dns.Stop()
}
return
}
}
func (revsuit *Revsuit) getMySQLConfig(c *gin.Context) {
c.JSON(200, revsuit.mysql.Config)
}
func (revsuit *Revsuit) updateMySQLConfig(c *gin.Context) {
var form = mysql.Config{}
if err := c.ShouldBindJSON(&form); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"data": nil,
})
return
}
if form.Addr != revsuit.mysql.Addr {
revsuit.mysql.Addr = form.Addr
log.Info("Update mysql config [addr] to %s", form.Addr)
}
if form.VersionString != revsuit.mysql.VersionString {
revsuit.mysql.VersionString = form.VersionString
log.Info("Update mysql config [version_string] to %s", form.VersionString)
}
if form.Enable != revsuit.mysql.Enable {
log.Info("Update mysql config [enable] to %v", form.Enable)
if form.Enable {
go revsuit.mysql.Run()
} else {
revsuit.mysql.Stop()
}
return
}
if revsuit.mysql.Enable {
revsuit.mysql.Restart()
}
}
func (revsuit *Revsuit) getRmiConfig(c *gin.Context) {
c.JSON(200, revsuit.rmi.Config)
}
func (revsuit *Revsuit) updateRmiConfig(c *gin.Context) {
var form = rmi.Config{}
if err := c.ShouldBindJSON(&form); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"data": nil,
})
return
}
if form.Addr != revsuit.rmi.Addr {
revsuit.rmi.Addr = form.Addr
log.Info("Update rmi config [addr] to %s", form.Addr)
}
if form.Enable != revsuit.rmi.Enable {
log.Info("Update rmi config [enable] to %v", form.Enable)
if form.Enable {
go revsuit.rmi.Run()
} else {
revsuit.rmi.Stop()
}
return
}
if revsuit.rmi.Enable {
revsuit.rmi.Restart()
}
}
func (revsuit *Revsuit) getNoticeConfig(c *gin.Context) {
c.JSON(200, revsuit.config.Notice)
}
func (revsuit *Revsuit) updateNoticeConfig(c *gin.Context) {
var form = noticeConfig{}
if err := c.ShouldBindJSON(&form); err != nil {
c.JSON(400, gin.H{
"status": "failed",
"error": err.Error(),
"data": nil,
})
return
}
log.Info("Update notice config")
revsuit.config.Notice = form
initNotice(form)
}