diff --git a/api/server.go b/api/server.go index beb1b76..c8798e2 100644 --- a/api/server.go +++ b/api/server.go @@ -36,12 +36,13 @@ func (s *ServerAPI) GetTraffic(stream TrojanServerService_GetTrafficServer) erro if !valid { stream.Send(&GetTrafficResponse{ Success: false, - Info: "invalid user", + Info: "invalid user " + req.User.Hash, }) continue } downloadTraffic, uploadTraffic := meter.Get() downloadSpeed, uploadSpeed := meter.GetSpeed() + downloadSpeedLimit, uploadSpeedLimit := meter.GetSpeedLimit() err = stream.Send(&GetTrafficResponse{ Success: true, TrafficTotal: &Traffic{ @@ -52,6 +53,10 @@ func (s *ServerAPI) GetTraffic(stream TrojanServerService_GetTrafficServer) erro DownloadSpeed: downloadSpeed, UploadSpeed: uploadSpeed, }, + SpeedLimit: &Speed{ + DownloadSpeed: uint64(downloadSpeedLimit), + UploadSpeed: uint64(uploadSpeedLimit), + }, }) if err != nil { return err @@ -80,7 +85,12 @@ func (s *ServerAPI) SetUsers(stream TrojanServerService_SetUsersServer) error { case SetUserRequest_Delete: err = s.auth.DelUser(req.User.Hash) case SetUserRequest_Modify: - err = common.NewError("not support yet") + valid, meter := 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 err != nil { stream.Send(&SetUserResponse{ @@ -100,6 +110,7 @@ func (s *ServerAPI) ListUsers(req *ListUserRequest, stream TrojanServerService_L for _, meter := range users { downloadTraffic, uploadTraffic := meter.Get() downloadSpeed, uploadSpeed := meter.GetSpeed() + downloadSpeedLimit, uploadSpeedLimit := meter.GetSpeedLimit() err := stream.Send(&ListUserResponse{ User: &User{ Hash: meter.Hash(), @@ -112,6 +123,10 @@ func (s *ServerAPI) ListUsers(req *ListUserRequest, stream TrojanServerService_L DownloadSpeed: downloadSpeed, UploadSpeed: uploadSpeed, }, + SpeedLimit: &Speed{ + DownloadSpeed: uint64(downloadSpeedLimit), + UploadSpeed: uint64(uploadSpeedLimit), + }, }) if err != nil { return err diff --git a/api/server_test.go b/api/server_test.go index 7f01ad8..c7e8d75 100644 --- a/api/server_test.go +++ b/api/server_test.go @@ -36,11 +36,12 @@ func TestServerAPI(t *testing.T) { if resp.User.Hash != "hash1234" { t.Fail() } + fmt.Println(resp.SpeedCurrent) + fmt.Println(resp.SpeedLimit) } stream1.CloseSend() - meter.Count(1234, 5678) - time.Sleep(time.Millisecond * 400) + time.Sleep(time.Millisecond * 1000) stream2, err := server.GetTraffic(ctx) common.Must(err) stream2.Send(&GetTrafficRequest{ @@ -56,7 +57,6 @@ func TestServerAPI(t *testing.T) { if resp2.SpeedCurrent.DownloadSpeed != 1234 || resp2.TrafficTotal.UploadTraffic != 5678 { t.Fail() } - stream2.CloseSend() stream3, err := server.SetUsers(ctx) stream3.Send(&SetUserRequest{ @@ -83,9 +83,42 @@ func TestServerAPI(t *testing.T) { if err != nil || !resp3.Success { t.Fail() } - valid, _ = auth.AuthUser("newhash") + valid, meter = auth.AuthUser("newhash") if !valid { t.Fail() } + stream3.Send(&SetUserRequest{ + User: &User{ + Hash: "newhash", + }, + Operation: SetUserRequest_Modify, + SpeedLimit: &Speed{ + DownloadSpeed: 5000, + UploadSpeed: 3000, + }, + }) + go func() { + for { + meter.Count(200, 0) + } + }() + go func() { + for { + meter.Count(0, 300) + } + }() + time.Sleep(time.Second * 3) + for i := 0; i < 3; i++ { + stream2.Send(&GetTrafficRequest{ + User: &User{ + Hash: "newhash", + }, + }) + resp2, err = stream2.Recv() + fmt.Println(resp2.SpeedCurrent) + fmt.Println(resp2.SpeedLimit) + time.Sleep(time.Second) + } + stream2.CloseSend() cancel() } diff --git a/stat/memory/memory.go b/stat/memory/memory.go index ec9070d..1539c1b 100644 --- a/stat/memory/memory.go +++ b/stat/memory/memory.go @@ -45,11 +45,11 @@ func (m *MemoryTrafficMeter) Count(sent, recv int) { atomic.AddUint64(&m.recv, uint64(recv)) } -func (m *MemoryTrafficMeter) LimitSpeed(sent, recv int) { - if sent == 0 { +func (m *MemoryTrafficMeter) LimitSpeed(send, recv int) { + if send == 0 { m.sendLimiter = nil } else { - m.sendLimiter = rate.NewLimiter(rate.Limit(sent), sent*2) + m.sendLimiter = rate.NewLimiter(rate.Limit(send), send*2) } if recv == 0 { m.recvLimiter = nil @@ -58,6 +58,18 @@ func (m *MemoryTrafficMeter) LimitSpeed(sent, recv int) { } } +func (m *MemoryTrafficMeter) GetSpeedLimit() (send, recv int) { + sendLimit := 0 + recvLimit := 0 + if m.sendLimiter != nil { + sendLimit = int(m.sendLimiter.Limit()) + } + if m.recvLimiter != nil { + recvLimit = int(m.recvLimiter.Limit()) + } + return sendLimit, recvLimit +} + func (m *MemoryTrafficMeter) Hash() string { return m.hash } diff --git a/stat/stat.go b/stat/stat.go index d536bce..edec78b 100644 --- a/stat/stat.go +++ b/stat/stat.go @@ -11,12 +11,13 @@ import ( type TrafficMeter interface { io.Closer Hash() string - Count(sent int, recv int) - Get() (sent uint64, recv uint64) + Count(sent, recv int) + Get() (sent, recv uint64) Reset() - GetAndReset() (sent uint64, recv uint64) - GetSpeed() (sent uint64, recv uint64) - LimitSpeed(sent int, recv int) + GetAndReset() (sent, recv uint64) + GetSpeed() (sent, recv uint64) + LimitSpeed(send, recv int) + GetSpeedLimit() (send, recv int) } type Authenticator interface {