mirror of
https://github.com/p4gefau1t/trojan-go.git
synced 2024-04-21 12:21:34 +00:00
add server-side ip limit api
This commit is contained in:
+740
-393
File diff suppressed because it is too large
Load Diff
+5
-3
@@ -1,6 +1,6 @@
|
||||
syntax = "proto3";
|
||||
package trojan.api;
|
||||
option go_package = "api";
|
||||
option go_package = ".;api";
|
||||
|
||||
message Traffic {
|
||||
uint64 upload_traffic = 1;
|
||||
@@ -39,17 +39,19 @@ message ListUserResponse {
|
||||
Traffic traffic_total = 3;
|
||||
Speed speed_current = 4;
|
||||
Speed speed_limit = 5;
|
||||
int32 ip_limit = 6;
|
||||
}
|
||||
|
||||
message SetUserRequest {
|
||||
User user = 1;
|
||||
Speed speed_limit = 2;
|
||||
enum Operation {
|
||||
Add = 0;
|
||||
Delete = 1;
|
||||
Modify = 2;
|
||||
}
|
||||
Operation operation = 3;
|
||||
Operation operation = 2;
|
||||
Speed speed_limit = 3;
|
||||
int32 ip_limit = 4;
|
||||
}
|
||||
|
||||
message SetUserResponse {
|
||||
|
||||
+3
-3
@@ -33,12 +33,12 @@ func (s *ClientAPI) GetTraffic(ctx context.Context, req *GetTrafficRequest) (*Ge
|
||||
if req.User.Hash == "" {
|
||||
req.User.Hash = common.SHA224String(req.User.Password)
|
||||
}
|
||||
valid, meter := s.auth.AuthUser(req.User.Hash)
|
||||
valid, user := s.auth.AuthUser(req.User.Hash)
|
||||
if !valid {
|
||||
return nil, common.NewError("User " + req.User.Hash + " not found")
|
||||
}
|
||||
sent, recv := meter.Get()
|
||||
sentSpeed, recvSpeed := meter.GetSpeed()
|
||||
sent, recv := user.GetTraffic()
|
||||
sentSpeed, recvSpeed := user.GetSpeed()
|
||||
resp := &GetTrafficResponse{
|
||||
Success: true,
|
||||
TrafficTotal: &Traffic{
|
||||
|
||||
+2
-2
@@ -22,11 +22,11 @@ func TestClientAPI(t *testing.T) {
|
||||
},
|
||||
}, auth)
|
||||
common.Must(auth.AddUser("hash1234"))
|
||||
valid, meter := auth.AuthUser("hash1234")
|
||||
valid, user := auth.AuthUser("hash1234")
|
||||
if !valid {
|
||||
t.Fail()
|
||||
}
|
||||
meter.Count(1234, 5678)
|
||||
user.AddTraffic(1234, 5678)
|
||||
time.Sleep(time.Second)
|
||||
conn, err := grpc.Dial("127.0.0.1:10000", grpc.WithInsecure())
|
||||
common.Must(err)
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
package api
|
||||
|
||||
import "github.com/p4gefau1t/trojan-go/common"
|
||||
|
||||
// TODO implement api service client
|
||||
|
||||
type apiOption struct {
|
||||
common.OptionHandler
|
||||
}
|
||||
|
||||
func (apiOption) Name() string {
|
||||
return "api"
|
||||
}
|
||||
|
||||
func (o *apiOption) Handle() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (o *apiOption) Priority() int {
|
||||
return 50
|
||||
}
|
||||
+16
-9
@@ -36,7 +36,7 @@ func (s *ServerAPI) GetTraffic(stream TrojanServerService_GetTrafficServer) erro
|
||||
if req.User.Hash == "" {
|
||||
req.User.Hash = common.SHA224String(req.User.Password)
|
||||
}
|
||||
valid, meter := s.auth.AuthUser(req.User.Hash)
|
||||
valid, user := s.auth.AuthUser(req.User.Hash)
|
||||
if !valid {
|
||||
stream.Send(&GetTrafficResponse{
|
||||
Success: false,
|
||||
@@ -44,9 +44,9 @@ func (s *ServerAPI) GetTraffic(stream TrojanServerService_GetTrafficServer) erro
|
||||
})
|
||||
continue
|
||||
}
|
||||
downloadTraffic, uploadTraffic := meter.Get()
|
||||
downloadSpeed, uploadSpeed := meter.GetSpeed()
|
||||
downloadSpeedLimit, uploadSpeedLimit := meter.GetSpeedLimit()
|
||||
downloadTraffic, uploadTraffic := user.GetTraffic()
|
||||
downloadSpeed, uploadSpeed := user.GetSpeed()
|
||||
downloadSpeedLimit, uploadSpeedLimit := user.GetSpeedLimit()
|
||||
err = stream.Send(&GetTrafficResponse{
|
||||
Success: true,
|
||||
TrafficTotal: &Traffic{
|
||||
@@ -88,20 +88,25 @@ func (s *ServerAPI) SetUsers(stream TrojanServerService_SetUsersServer) error {
|
||||
case SetUserRequest_Add:
|
||||
err = s.auth.AddUser(req.User.Hash)
|
||||
if req.SpeedLimit != nil {
|
||||
valid, meter := s.auth.AuthUser(req.User.Hash)
|
||||
valid, user := s.auth.AuthUser(req.User.Hash)
|
||||
if !valid {
|
||||
return common.NewError("Failed to add new user")
|
||||
}
|
||||
meter.LimitSpeed(int(req.SpeedLimit.DownloadSpeed), int(req.SpeedLimit.UploadSpeed))
|
||||
user.SetSpeedLimit(int(req.SpeedLimit.DownloadSpeed), int(req.SpeedLimit.UploadSpeed))
|
||||
}
|
||||
case SetUserRequest_Delete:
|
||||
err = s.auth.DelUser(req.User.Hash)
|
||||
case SetUserRequest_Modify:
|
||||
valid, meter := s.auth.AuthUser(req.User.Hash)
|
||||
valid, user := s.auth.AuthUser(req.User.Hash)
|
||||
if !valid {
|
||||
err = common.NewError("Invalid user " + req.User.Hash)
|
||||
} else {
|
||||
meter.LimitSpeed(int(req.SpeedLimit.DownloadSpeed), int(req.SpeedLimit.UploadSpeed))
|
||||
if req.SpeedLimit.DownloadSpeed > 0 || req.SpeedLimit.UploadSpeed > 0 {
|
||||
user.SetSpeedLimit(int(req.SpeedLimit.DownloadSpeed), int(req.SpeedLimit.UploadSpeed))
|
||||
}
|
||||
if req.IpLimit > 0 {
|
||||
user.SetIPLimit(int(req.IpLimit))
|
||||
}
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
@@ -121,9 +126,10 @@ func (s *ServerAPI) ListUsers(req *ListUserRequest, stream TrojanServerService_L
|
||||
log.Debug("API: ListUsers")
|
||||
users := s.auth.ListUsers()
|
||||
for _, meter := range users {
|
||||
downloadTraffic, uploadTraffic := meter.Get()
|
||||
downloadTraffic, uploadTraffic := meter.GetTraffic()
|
||||
downloadSpeed, uploadSpeed := meter.GetSpeed()
|
||||
downloadSpeedLimit, uploadSpeedLimit := meter.GetSpeedLimit()
|
||||
ipLimit := meter.GetIPLimit()
|
||||
online := false
|
||||
if downloadSpeed > 0 || uploadSpeed > 0 {
|
||||
online = true
|
||||
@@ -145,6 +151,7 @@ func (s *ServerAPI) ListUsers(req *ListUserRequest, stream TrojanServerService_L
|
||||
DownloadSpeed: uint64(downloadSpeedLimit),
|
||||
UploadSpeed: uint64(uploadSpeedLimit),
|
||||
},
|
||||
IpLimit: int32(ipLimit),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
+5
-5
@@ -23,7 +23,7 @@ func TestServerAPI(t *testing.T) {
|
||||
},
|
||||
}, auth)
|
||||
common.Must(auth.AddUser("hash1234"))
|
||||
_, meter := auth.AuthUser("hash1234")
|
||||
_, user := auth.AuthUser("hash1234")
|
||||
conn, err := grpc.Dial("127.0.0.1:10000", grpc.WithInsecure())
|
||||
server := NewTrojanServerServiceClient(conn)
|
||||
stream1, err := server.ListUsers(ctx, &ListUserRequest{})
|
||||
@@ -41,7 +41,7 @@ func TestServerAPI(t *testing.T) {
|
||||
fmt.Println(resp.SpeedLimit)
|
||||
}
|
||||
stream1.CloseSend()
|
||||
meter.Count(1234, 5678)
|
||||
user.AddTraffic(1234, 5678)
|
||||
time.Sleep(time.Millisecond * 1000)
|
||||
stream2, err := server.GetTraffic(ctx)
|
||||
common.Must(err)
|
||||
@@ -84,7 +84,7 @@ func TestServerAPI(t *testing.T) {
|
||||
if err != nil || !resp3.Success {
|
||||
t.Fail()
|
||||
}
|
||||
valid, meter = auth.AuthUser("newhash")
|
||||
valid, user = auth.AuthUser("newhash")
|
||||
if !valid {
|
||||
t.Fail()
|
||||
}
|
||||
@@ -100,12 +100,12 @@ func TestServerAPI(t *testing.T) {
|
||||
})
|
||||
go func() {
|
||||
for {
|
||||
meter.Count(200, 0)
|
||||
user.AddTraffic(200, 0)
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
for {
|
||||
meter.Count(0, 300)
|
||||
user.AddTraffic(0, 300)
|
||||
}
|
||||
}()
|
||||
time.Sleep(time.Second * 3)
|
||||
|
||||
Reference in New Issue
Block a user