diff --git a/go.mod b/go.mod index 899e4cf4..fc4f646f 100644 --- a/go.mod +++ b/go.mod @@ -23,7 +23,7 @@ require ( github.com/hibiken/asynq v0.25.1 github.com/maypok86/otter/v2 v2.3.0 github.com/peterbourgon/diskv/v3 v3.0.1 - github.com/pressly/goose/v3 v3.15.1 + github.com/pressly/goose/v3 v3.24.0 github.com/redis/go-redis/extra/redisotel/v9 v9.16.0 github.com/redis/go-redis/v9 v9.16.0 github.com/robfig/cron/v3 v3.0.1 @@ -133,8 +133,10 @@ require ( github.com/leodido/go-urn v1.4.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-sqlite3 v1.14.22 // indirect + github.com/mfridman/interpolate v0.0.2 // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect github.com/modern-go/reflect2 v1.0.2 // indirect + github.com/ncruces/go-strftime v0.1.9 // indirect github.com/paulmach/orb v0.12.0 // indirect github.com/pelletier/go-toml/v2 v2.2.4 // indirect github.com/pierrec/lz4/v4 v4.1.22 // indirect @@ -145,6 +147,7 @@ require ( github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/sagikazarmark/locafero v0.12.0 // indirect github.com/segmentio/asm v1.2.1 // indirect + github.com/sethvargo/go-retry v0.3.0 // indirect github.com/spf13/afero v1.15.0 // indirect github.com/spf13/cast v1.10.0 // indirect github.com/spf13/pflag v1.0.10 // indirect @@ -173,8 +176,8 @@ require ( google.golang.org/protobuf v1.36.10 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect gorm.io/driver/mysql v1.6.0 // indirect - modernc.org/libc v1.24.1 // indirect + modernc.org/libc v1.55.3 // indirect modernc.org/mathutil v1.6.0 // indirect - modernc.org/memory v1.7.2 // indirect - modernc.org/sqlite v1.26.0 // indirect + modernc.org/memory v1.8.0 // indirect + modernc.org/sqlite v1.34.1 // indirect ) diff --git a/go.sum b/go.sum index 1da7d2d3..55501324 100644 --- a/go.sum +++ b/go.sum @@ -168,8 +168,8 @@ github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/gofuzz v1.2.0 h1:xRy4A+RhZaiKjJ1bPfwQ8sedCA+YS2YcCHW6ec7JMi0= github.com/google/gofuzz v1.2.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= -github.com/google/pprof v0.0.0-20221118152302-e6195bd50e26 h1:Xim43kblpZXfIBQsbuBVKCudVG457BR2GZFIz3uw3hQ= -github.com/google/pprof v0.0.0-20221118152302-e6195bd50e26/go.mod h1:dDKJzRmX4S37WGHujM7tX//fmj1uioxKzKxz3lo4HJo= +github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd h1:gbpYu9NMq8jhDVbvlGkMFWCjLFlqqEZjEmObmhUy6Vo= +github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd/go.mod h1:kf6iHlnVGwgKolg33glAes7Yg/8iWP8ukqeldJSO7jw= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/context v1.1.2 h1:WRkNAv2uoa03QNIc1A6u4O7DAGMUVoopZhkiXWA2V1o= @@ -184,6 +184,8 @@ github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5uk github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3/go.mod h1:ndYquD05frm2vACXE1nsccT4oJzjhw2arTS2cpUD1PI= github.com/hashicorp/go-version v1.7.0 h1:5tqGy27NaOTB8yJKUZELlFAS/LTKJkrmONwQKeRZfjY= github.com/hashicorp/go-version v1.7.0/go.mod h1:fltr4n8CU8Ke44wwGCBoEymUuxUHl09ZGVZPK5anwXA= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/hibiken/asynq v0.25.1 h1:phj028N0nm15n8O2ims+IvJ2gz4k2auvermngh9JhTw= github.com/hibiken/asynq v0.25.1/go.mod h1:pazWNOLBu0FEynQRBvHA26qdIKRSmfdIfUm4HdsLmXg= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= @@ -202,8 +204,6 @@ github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= -github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51 h1:Z9n2FFNUXsshfwJMBgNA0RU6/i7WVaAegv3PtuIHPMs= -github.com/kballard/go-shellquote v0.0.0-20180428030007-95032a82bc51/go.mod h1:CzGEWj7cYgsdH8dAjBGEr58BoE7ScuLd+fwFZ44+/x8= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= github.com/klauspost/compress v1.13.6/go.mod h1:/3/Vjq9QcHkK5uEr5lBEmyoZ1iFhe47etQ6QUkpK6sk= @@ -226,12 +226,16 @@ github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= github.com/maypok86/otter/v2 v2.3.0 h1:8H8AVVFUSzJwIegKwv1uF5aGitTY+AIrtktg7OcLs8w= github.com/maypok86/otter/v2 v2.3.0/go.mod h1:XgIdlpmL6jYz882/CAx1E4C1ukfgDKSaw4mWq59+7l8= +github.com/mfridman/interpolate v0.0.2 h1:pnuTK7MQIxxFz1Gr+rjSIx9u7qVjf5VOoM/u6BbAxPY= +github.com/mfridman/interpolate v0.0.2/go.mod h1:p+7uk6oE07mpE/Ik1b8EckO0O4ZXiGAfshKBWLUM9Xg= github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= github.com/montanaflynn/stats v0.0.0-20171201202039-1bf9dbcd8cbe/go.mod h1:wL8QJuTMNUDYhXwkmfOly8iTdp5TEcJFWZD2D7SIkUc= +github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= +github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/paulmach/orb v0.12.0 h1:z+zOwjmG3MyEEqzv92UN49Lg1JFYx0L9GpGKNVDKk1s= github.com/paulmach/orb v0.12.0/go.mod h1:5mULz1xQfs3bmQm63QEJA6lNGujuRafwA5S/EnuLaLU= github.com/paulmach/protoscan v0.2.1/go.mod h1:SpcSwydNLrxUGSDvXvO0P7g7AuhJ7lcKfDlhJCDw2gY= @@ -244,8 +248,8 @@ github.com/pierrec/lz4/v4 v4.1.22/go.mod h1:gZWDp/Ze/IJXGXf23ltt2EXimqmTUXEy0GFu github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/pressly/goose/v3 v3.15.1 h1:dKaJ1SdLvS/+HtS8PzFT0KBEtICC1jewLXM+b3emlv8= -github.com/pressly/goose/v3 v3.15.1/go.mod h1:0E3Yg/+EwYzO6Rz2P98MlClFgIcoujbVRs575yi3iIM= +github.com/pressly/goose/v3 v3.24.0 h1:sFbNms7Bd++2VMq6HSgDHDLWa7kHz1qXzPb3ZIU72VU= +github.com/pressly/goose/v3 v3.24.0/go.mod h1:rEWreU9uVtt0DHCyLzF9gRcWiiTF/V+528DV+4DORug= github.com/quic-go/qpack v0.5.1 h1:giqksBPnT/HDtZ6VhtFKgoLOWmlyo9Ei6u9PqzIMbhI= github.com/quic-go/qpack v0.5.1/go.mod h1:+PC4XFrEskIVkcLzpEkbLqq1uCoxPhQuvK5rH1ZgaEg= github.com/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk= @@ -267,6 +271,8 @@ github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88ee github.com/sagikazarmark/locafero v0.12.0/go.mod h1:sZh36u/YSZ918v0Io+U9ogLYQJ9tLLBmM4eneO6WwsI= github.com/segmentio/asm v1.2.1 h1:DTNbBqs57ioxAD4PrArqftgypG4/qNpXoJx8TVXxPR0= github.com/segmentio/asm v1.2.1/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= +github.com/sethvargo/go-retry v0.3.0 h1:EEt31A35QhrcRZtrYFDTBg91cqZVnFL2navjDrah2SE= +github.com/sethvargo/go-retry v0.3.0/go.mod h1:mNX17F0C/HguQMyMyJxcnU471gOZGxCLyYaFyAZraas= github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k= github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME= github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I= @@ -463,22 +469,28 @@ gorm.io/plugin/dbresolver v1.6.2 h1:F4b85TenghUeITqe3+epPSUtHH7RIk3fXr5l83DF8Pc= gorm.io/plugin/dbresolver v1.6.2/go.mod h1:tctw63jdrOezFR9HmrKnPkmig3m5Edem9fdxk9bQSzM= gorm.io/plugin/opentelemetry v0.1.14 h1:xivP39t/0JgcceDl+BLwVAJHihjFEUj0ZocMSBwZ7ZY= gorm.io/plugin/opentelemetry v0.1.14/go.mod h1:ZAp4v5vU1CCcK9Oo8/va5rl6NStrzpSU+a70evd+W/g= -lukechampine.com/uint128 v1.3.0 h1:cDdUVfRwDUDovz610ABgFD17nXD4/uDgVHl2sC3+sbo= -lukechampine.com/uint128 v1.3.0/go.mod h1:c4eWIwlEGaxC/+H1VguhU4PHXNWDCDMUlWdIWl2j1gk= -modernc.org/cc/v3 v3.41.0 h1:QoR1Sn3YWlmA1T4vLaKZfawdVtSiGx8H+cEojbC7v1Q= -modernc.org/cc/v3 v3.41.0/go.mod h1:Ni4zjJYJ04CDOhG7dn640WGfwBzfE0ecX8TyMB0Fv0Y= -modernc.org/ccgo/v3 v3.16.15 h1:KbDR3ZAVU+wiLyMESPtbtE/Add4elztFyfsWoNTgxS0= -modernc.org/ccgo/v3 v3.16.15/go.mod h1:yT7B+/E2m43tmMOT51GMoM98/MtHIcQQSleGnddkUNI= -modernc.org/libc v1.24.1 h1:uvJSeCKL/AgzBo2yYIPPTy82v21KgGnizcGYfBHaNuM= -modernc.org/libc v1.24.1/go.mod h1:FmfO1RLrU3MHJfyi9eYYmZBfi/R+tqZ6+hQ3yQQUkak= +modernc.org/cc/v4 v4.21.4 h1:3Be/Rdo1fpr8GrQ7IVw9OHtplU4gWbb+wNgeoBMmGLQ= +modernc.org/cc/v4 v4.21.4/go.mod h1:HM7VJTZbUCR3rV8EYBi9wxnJ0ZBRiGE5OeGXNA0IsLQ= +modernc.org/ccgo/v4 v4.19.2 h1:lwQZgvboKD0jBwdaeVCTouxhxAyN6iawF3STraAal8Y= +modernc.org/ccgo/v4 v4.19.2/go.mod h1:ysS3mxiMV38XGRTTcgo0DQTeTmAO4oCmJl1nX9VFI3s= +modernc.org/fileutil v1.3.0 h1:gQ5SIzK3H9kdfai/5x41oQiKValumqNTDXMvKo62HvE= +modernc.org/fileutil v1.3.0/go.mod h1:XatxS8fZi3pS8/hKG2GH/ArUogfxjpEKs3Ku3aK4JyQ= +modernc.org/gc/v2 v2.4.1 h1:9cNzOqPyMJBvrUipmynX0ZohMhcxPtMccYgGOJdOiBw= +modernc.org/gc/v2 v2.4.1/go.mod h1:wzN5dK1AzVGoH6XOzc3YZ+ey/jPgYHLuVckd62P0GYU= +modernc.org/gc/v3 v3.0.0-20240107210532-573471604cb6 h1:5D53IMaUuA5InSeMu9eJtlQXS2NxAhyWQvkKEgXZhHI= +modernc.org/gc/v3 v3.0.0-20240107210532-573471604cb6/go.mod h1:Qz0X07sNOR1jWYCrJMEnbW/X55x206Q7Vt4mz6/wHp4= +modernc.org/libc v1.55.3 h1:AzcW1mhlPNrRtjS5sS+eW2ISCgSOLLNyFzRh/V3Qj/U= +modernc.org/libc v1.55.3/go.mod h1:qFXepLhz+JjFThQ4kzwzOjA/y/artDeg+pcYnY+Q83w= modernc.org/mathutil v1.6.0 h1:fRe9+AmYlaej+64JsEEhoWuAYBkOtQiMEU7n/XgfYi4= modernc.org/mathutil v1.6.0/go.mod h1:Ui5Q9q1TR2gFm0AQRqQUaBWFLAhQpCwNcuhBOSedWPo= -modernc.org/memory v1.7.2 h1:Klh90S215mmH8c9gO98QxQFsY+W451E8AnzjoE2ee1E= -modernc.org/memory v1.7.2/go.mod h1:NO4NVCQy0N7ln+T9ngWqOQfi7ley4vpwvARR+Hjw95E= +modernc.org/memory v1.8.0 h1:IqGTL6eFMaDZZhEWwcREgeMXYwmW83LYW8cROZYkg+E= +modernc.org/memory v1.8.0/go.mod h1:XPZ936zp5OMKGWPqbD3JShgd/ZoQ7899TUuQqxY+peU= modernc.org/opt v0.1.3 h1:3XOZf2yznlhC+ibLltsDGzABUGVx8J6pnFMS3E4dcq4= modernc.org/opt v0.1.3/go.mod h1:WdSiB5evDcignE70guQKxYUl14mgWtbClRi5wmkkTX0= -modernc.org/sqlite v1.26.0 h1:SocQdLRSYlA8W99V8YH0NES75thx19d9sB/aFc4R8Lw= -modernc.org/sqlite v1.26.0/go.mod h1:FL3pVXie73rg3Rii6V/u5BoHlSoyeZeIgKZEgHARyCU= +modernc.org/sortutil v1.2.0 h1:jQiD3PfS2REGJNzNCMMaLSp/wdMNieTbKX920Cqdgqc= +modernc.org/sortutil v1.2.0/go.mod h1:TKU2s7kJMf1AE84OoiGppNHJwvB753OYfNl2WRb++Ss= +modernc.org/sqlite v1.34.1 h1:u3Yi6M0N8t9yKRDwhXcyp1eS5/ErhPTBggxWFuR6Hfk= +modernc.org/sqlite v1.34.1/go.mod h1:pXV2xHxhzXZsgT/RtTFAPY6JJDEvOTcTdwADQCCWD4k= modernc.org/strutil v1.2.0 h1:agBi9dp1I+eOnxXeiZawM8F4LawKv4NzGWSaLfyeNZA= modernc.org/strutil v1.2.0/go.mod h1:/mdcBmfOibveCTBxUl5B5l6W+TTH1FXPLHZE6bTosX0= modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= diff --git a/internal/apps/admin/push/custom_events/admin_login_test.go b/internal/apps/admin/push/custom_events/admin_login_test.go index b9c54f81..bc31b35e 100644 --- a/internal/apps/admin/push/custom_events/admin_login_test.go +++ b/internal/apps/admin/push/custom_events/admin_login_test.go @@ -13,6 +13,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/admin/push" "github.com/Rain-kl/Wavelet/internal/listener" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/internal/testhelper" "github.com/hibiken/asynq" @@ -88,7 +89,7 @@ func enableAdminLoginEvent(t *testing.T, dbConn *gorm.DB, channelName string, ta event.Enabled = true event.Channels = []string{channelName} event.Targets = targets - require.NoError(t, dbConn.Save(&event).Error) + require.NoError(t, repository.SavePushEvent(context.Background(), &event)) } func waitForAsyncTrigger(t *testing.T) { @@ -165,11 +166,11 @@ func TestAdminLoginPushIntegration(t *testing.T) { var event model.PushEvent require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error) event.Enabled = false - require.NoError(t, dbConn.Save(&event).Error) + require.NoError(t, repository.SavePushEvent(context.Background(), &event)) listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1") waitForAsyncTrigger(t) assert.Equal(t, int64(0), countPushTasks(t, dbConn)) }) -} \ No newline at end of file +} diff --git a/internal/apps/admin/push/push_test.go b/internal/apps/admin/push/push_test.go index f7aa94fd..a0ff2d5e 100644 --- a/internal/apps/admin/push/push_test.go +++ b/internal/apps/admin/push/push_test.go @@ -3,7 +3,8 @@ package push -import ("bytes" +import ( + "bytes" "context" "encoding/json" "net/http" @@ -13,7 +14,10 @@ import ("bytes" "testing" "time" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/internal/testhelper" pkgpush "github.com/Rain-kl/Wavelet/pkg/push" @@ -23,8 +27,7 @@ import ("bytes" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" - - "github.com/Rain-kl/Wavelet/internal/common/response") +) var adminLoginEvent = EventMetadata{ Key: "admin_login", @@ -187,7 +190,7 @@ func TestEventTrigger(t *testing.T) { event.Enabled = true event.Channels = []string{"mock_channel"} event.Targets = []string{"admin_user"} - err = dbConn.Save(&event).Error + err = repository.SavePushEvent(context.Background(), &event) require.NoError(t, err) // Trigger @@ -244,8 +247,8 @@ func TestEventTrigger(t *testing.T) { event.Enabled = true event.Channels = []string{"mock_channel"} - event.Targets = []string{"user.username"} // 动态目标 - err = dbConn.Save(&event).Error + event.Targets = []string{"user.username"} + err = repository.SavePushEvent(context.Background(), &event) require.NoError(t, err) // Trigger with empty body (simulates cron scheduler triggering) @@ -369,7 +372,7 @@ func TestPushRouters(t *testing.T) { // 2. 为该事件关联渠道后,再切换开启,应当成功 event.Channels = []string{"email"} - dbConn.Save(&event) + _ = repository.SavePushEvent(context.Background(), &event) req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil) w2 := httptest.NewRecorder() diff --git a/internal/apps/admin/user/logics.go b/internal/apps/admin/user/logics.go index fb263923..d5704dc5 100644 --- a/internal/apps/admin/user/logics.go +++ b/internal/apps/admin/user/logics.go @@ -9,6 +9,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" @@ -35,7 +37,22 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error { if !active && flags.IsAdmin { return errors.New(cannotDisable) } - return repository.UpdateUserActive(ctx, id, active) + + var tokens []model.AccessToken + if !active { + _ = db.DB(ctx).Where("user_id = ?", id).Find(&tokens).Error + } + + err = repository.UpdateUserActive(ctx, id, active) + if err == nil { + oauth.InvalidateCachedUser(ctx, id) + if !active { + for _, token := range tokens { + oauth.InvalidateCachedToken(ctx, token.TokenHash) + } + } + } + return err } func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { @@ -49,7 +66,18 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { if flags.IsAdmin { return errors.New(cannotDelete) } - return repository.DeleteUserWithRelations(ctx, targetID) + + var tokens []model.AccessToken + _ = db.DB(ctx).Where("user_id = ?", targetID).Find(&tokens).Error + + err = repository.DeleteUserWithRelations(ctx, targetID) + if err == nil { + oauth.InvalidateCachedUser(ctx, targetID) + for _, token := range tokens { + oauth.InvalidateCachedToken(ctx, token.TokenHash) + } + } + return err } func createUser(ctx context.Context, req createUserRequest) (model.User, error) { @@ -103,4 +131,4 @@ func createUser(ctx context.Context, req createUserRequest) (model.User, error) return model.User{}, err } return newUser, nil -} \ No newline at end of file +} diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go index 36acff97..b53c67f3 100644 --- a/internal/apps/admin/user/routers.go +++ b/internal/apps/admin/user/routers.go @@ -12,6 +12,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -103,7 +104,8 @@ func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbidde return true } } - response.AbortInternal(c, msg) + logger.ErrorF(c.Request.Context(), "Admin user error: %v", err) + response.AbortInternal(c, "内部服务器错误") return true } @@ -129,7 +131,8 @@ func ListUsers(c *gin.Context) { total, modelUsers, err := listUsers(c.Request.Context(), req) if err != nil { - response.AbortInternal(c, err.Error()) + logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err) + response.AbortInternal(c, "获取用户列表失败") return } @@ -285,4 +288,4 @@ func CreateUser(c *gin.Context) { } c.JSON(http.StatusOK, response.OK(toUser(newUser))) -} \ No newline at end of file +} diff --git a/internal/apps/cap/routers.go b/internal/apps/cap/routers.go index 1cc2eef9..3965e69a 100644 --- a/internal/apps/cap/routers.go +++ b/internal/apps/cap/routers.go @@ -6,9 +6,15 @@ package cap import ( "net/http" - "github.com/gin-gonic/gin" + "github.com/Rain-kl/Wavelet/internal/common/response" + pkgcap "github.com/Rain-kl/Wavelet/pkg/cap" + "github.com/Rain-kl/Wavelet/pkg/logger" + "github.com/gin-gonic/gin" ) +// ChallengeResponse is a local type alias for the pkg/cap.ChallengeResponse struct +type ChallengeResponse = pkgcap.ChallengeResponse + type challengeRequest struct { Scope string `json:"scope" form:"scope"` } @@ -26,8 +32,8 @@ type redeemRequest struct { // @Accept json // @Produce json // @Param request body challengeRequest false "可选范围限制参数" -// @Success 200 {object} cap.ChallengeResponse "成功返回 PoW 难题" -// @Failure 500 {object} RedeemResponse "内部服务错误" +// @Success 200 {object} response.Any{data=cap.ChallengeResponse} "成功返回 PoW 难题" +// @Failure 500 {object} response.Any "内部服务错误" // @Router /api/cap/challenge [post] func Challenge(c *gin.Context) { var req challengeRequest @@ -40,14 +46,12 @@ func Challenge(c *gin.Context) { mgr := GetDefaultManager() resp, err := mgr.Generate(c.Request.Context(), req.Scope) if err != nil { - c.JSON(http.StatusInternalServerError, RedeemResponse{ - Success: false, - Error: err.Error(), - }) + logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err) + response.AbortInternal(c, "生成验证难题失败,请稍后再试") return } - c.JSON(http.StatusOK, resp) + c.JSON(http.StatusOK, response.OK(resp)) } // Redeem 提交 PoW 解答并兑换一次性凭证 Token @@ -57,17 +61,14 @@ func Challenge(c *gin.Context) { // @Accept json // @Produce json // @Param request body redeemRequest true "难题 Token 与解答 solutions 数组" -// @Success 200 {object} RedeemResponse "核销成功,返回 X-Cap-Token" -// @Failure 400 {object} RedeemResponse "参数错误或核销失败" -// @Failure 500 {object} RedeemResponse "内部服务错误" +// @Success 200 {object} response.Any{data=cap.RedeemResponse} "核销成功,返回 X-Cap-Token" +// @Failure 400 {object} response.Any "参数错误或核销失败" +// @Failure 500 {object} response.Any "内部服务错误" // @Router /api/cap/redeem [post] func Redeem(c *gin.Context) { var req redeemRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, RedeemResponse{ - Success: false, - Error: "无效的参数", - }) + response.AbortBadRequest(c, "无效的参数") return } @@ -78,17 +79,15 @@ func Redeem(c *gin.Context) { mgr := GetDefaultManager() resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope) if err != nil { - c.JSON(http.StatusInternalServerError, RedeemResponse{ - Success: false, - Error: err.Error(), - }) + logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err) + response.AbortInternal(c, "校验验证解答失败,请稍后再试") return } if !resp.Success { - c.JSON(http.StatusBadRequest, resp) + response.AbortBadRequest(c, resp.Error) return } - c.JSON(http.StatusOK, resp) + c.JSON(http.StatusOK, response.OK(resp)) } diff --git a/internal/apps/cap/routers_test.go b/internal/apps/cap/routers_test.go index a99da38c..6db78767 100644 --- a/internal/apps/cap/routers_test.go +++ b/internal/apps/cap/routers_test.go @@ -46,10 +46,14 @@ func TestCapEndpointsAndMiddleware(t *testing.T) { t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String()) } - var challengeResp pkgcap.ChallengeResponse - if err := json.Unmarshal(w.Body.Bytes(), &challengeResp); err != nil { + var envelope struct { + ErrorMsg string `json:"error_msg"` + Data pkgcap.ChallengeResponse `json:"data"` + } + if err := json.Unmarshal(w.Body.Bytes(), &envelope); err != nil { t.Fatalf("failed to unmarshal challenge response: %v", err) } + challengeResp := envelope.Data if challengeResp.Token == "" { t.Fatalf("expected token in challenge response") @@ -99,10 +103,14 @@ func TestCapEndpointsAndMiddleware(t *testing.T) { t.Fatalf("expected 200 OK for redeem, got %d. Body: %s", w.Code, w.Body.String()) } - var redeemResp RedeemResponse - if err := json.Unmarshal(w.Body.Bytes(), &redeemResp); err != nil { + var redeemEnvelope struct { + ErrorMsg string `json:"error_msg"` + Data RedeemResponse `json:"data"` + } + if err := json.Unmarshal(w.Body.Bytes(), &redeemEnvelope); err != nil { t.Fatalf("failed to unmarshal redeem response: %v", err) } + redeemResp := redeemEnvelope.Data if !redeemResp.Success || redeemResp.Token == "" { t.Fatalf("redeem failed or returned empty token: %+v", redeemResp) diff --git a/internal/apps/oauth/cache.go b/internal/apps/oauth/cache.go new file mode 100644 index 00000000..86b3c326 --- /dev/null +++ b/internal/apps/oauth/cache.go @@ -0,0 +1,152 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "fmt" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +type cacheEntry struct { + value any + expiredAt time.Time +} + +type memoryCache struct { + sync.RWMutex + items map[string]cacheEntry +} + +var localCache = &memoryCache{ + items: make(map[string]cacheEntry), +} + +func (c *memoryCache) Set(key string, val any, ttl time.Duration) { + c.Lock() + defer c.Unlock() + c.items[key] = cacheEntry{ + value: val, + expiredAt: time.Now().Add(ttl), + } +} + +func (c *memoryCache) Get(key string) (any, bool) { + c.RLock() + item, ok := c.items[key] + if !ok { + c.RUnlock() + return nil, false + } + if time.Now().After(item.expiredAt) { + c.RUnlock() + c.Lock() + if item, ok = c.items[key]; ok && time.Now().After(item.expiredAt) { + delete(c.items, key) + } + c.Unlock() + return nil, false + } + c.RUnlock() + return item.value, true +} + +func (c *memoryCache) Delete(key string) { + c.Lock() + defer c.Unlock() + delete(c.items, key) +} + +const ( + tokenCacheTTL = 5 * time.Minute + userCacheTTL = 5 * time.Minute +) + +func tokenCacheKey(tokenHash string) string { + return "oauth:token:" + tokenHash +} + +func userCacheKey(userID uint64) string { + return fmt.Sprintf("oauth:user:%d", userID) +} + +// GetCachedToken 获取缓存的 AccessToken +func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) { + key := tokenCacheKey(tokenHash) + if val, ok := localCache.Get(key); ok { + if token, ok := val.(*model.AccessToken); ok { + return token, nil + } + } + + if db.Redis != nil { + var token model.AccessToken + if err := db.GetJSON(ctx, key, &token); err == nil { + // Write back to local cache + localCache.Set(key, &token, tokenCacheTTL) + return &token, nil + } + } + return nil, fmt.Errorf("cache miss") +} + +// SetCachedToken 设置 AccessToken 缓存 +func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) { + key := tokenCacheKey(tokenHash) + localCache.Set(key, token, tokenCacheTTL) + if db.Redis != nil { + _ = db.SetJSON(ctx, key, token, tokenCacheTTL) + } +} + +// InvalidateCachedToken 吊销/删除 token 缓存 +func InvalidateCachedToken(ctx context.Context, tokenHash string) { + key := tokenCacheKey(tokenHash) + localCache.Delete(key) + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() + } +} + +// GetCachedUser 获取缓存的 User +func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) { + key := userCacheKey(userID) + if val, ok := localCache.Get(key); ok { + if u, ok := val.(*model.User); ok { + return u, nil + } + } + + if db.Redis != nil { + var u model.User + if err := db.GetJSON(ctx, key, &u); err == nil { + // Write back to local cache + localCache.Set(key, &u, userCacheTTL) + return &u, nil + } + } + return nil, fmt.Errorf("cache miss") +} + +// SetCachedUser 设置 User 缓存 +func SetCachedUser(ctx context.Context, userID uint64, u *model.User) { + key := userCacheKey(userID) + localCache.Set(key, u, userCacheTTL) + if db.Redis != nil { + _ = db.SetJSON(ctx, key, u, userCacheTTL) + } +} + +// InvalidateCachedUser 吊销/失效 User 缓存 +func InvalidateCachedUser(ctx context.Context, userID uint64) { + key := userCacheKey(userID) + localCache.Delete(key) + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() + } +} diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 6ad871fa..6e341543 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -31,15 +31,26 @@ type loginRequiredAuditLog struct { func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) { tokenHash := model.HashToken(tokenStr) - var tokenRecord model.AccessToken - if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err != nil { - return nil, nil, err + tokenRecord, err := GetCachedToken(ctx, tokenHash) + if err != nil { + var dbToken model.AccessToken + if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&dbToken).Error; err != nil { + return nil, nil, err + } + tokenRecord = &dbToken + SetCachedToken(ctx, tokenHash, tokenRecord) } - var user model.User - if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err != nil { - return nil, nil, err + + user, err := GetCachedUser(ctx, tokenRecord.UserID) + if err != nil || !user.IsActive { + var dbUser model.User + if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil { + return nil, nil, err + } + user = &dbUser + SetCachedUser(ctx, tokenRecord.UserID, user) } - return &user, &tokenRecord, nil + return user, tokenRecord, nil } // GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error @@ -74,11 +85,16 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { return nil, errors.New("unauthorized") } - var user model.User - // load user from db to make sure is active - tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&user) - if tx.Error != nil { - return nil, tx.Error + user, err := GetCachedUser(ctx, userID) + if err != nil || !user.IsActive { + var dbUser model.User + // load user from db to make sure is active + tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&dbUser) + if tx.Error != nil { + return nil, tx.Error + } + user = &dbUser + SetCachedUser(ctx, userID, user) } // 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致 @@ -99,7 +115,7 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { return nil, errors.New("system user is not allowed to login") } - return &user, nil + return user, nil } // LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go index 41fc6ccc..8d2e13f6 100644 --- a/internal/apps/oauth/routers.go +++ b/internal/apps/oauth/routers.go @@ -4,15 +4,16 @@ package oauth -import ("net/http" +import ( + "net/http" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) // BasicUserInfo 用户基本信息结构体 type BasicUserInfo struct { @@ -94,6 +95,13 @@ func Logout(c *gin.Context) { username := session.Get(UserNameKey) if userID != nil { logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + if id, ok := userID.(uint64); ok { + InvalidateCachedUser(c.Request.Context(), id) + } else if idFloat, ok := userID.(float64); ok { + InvalidateCachedUser(c.Request.Context(), uint64(idFloat)) + } else if idInt, ok := userID.(int); ok && idInt >= 0 { + InvalidateCachedUser(c.Request.Context(), uint64(idInt)) + } } session.Options(GetSessionOptions(-1)) session.Clear() diff --git a/internal/apps/upload/handler/routers.go b/internal/apps/upload/handler/routers.go index 6940ebfb..d85998f8 100644 --- a/internal/apps/upload/handler/routers.go +++ b/internal/apps/upload/handler/routers.go @@ -7,6 +7,7 @@ package handler import ( "archive/zip" + "bufio" "bytes" "crypto/sha256" "encoding/hex" @@ -247,8 +248,12 @@ func BatchDownloadFiles(c *gin.Context) { c.Header("Content-Type", "application/zip") c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"") - zipWriter := zip.NewWriter(c.Writer) - defer func() { _ = zipWriter.Close() }() + bufferedWriter := bufio.NewWriter(c.Writer) + zipWriter := zip.NewWriter(bufferedWriter) + defer func() { + _ = zipWriter.Close() + _ = bufferedWriter.Flush() + }() usedNames := make(map[string]int) @@ -326,4 +331,3 @@ func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64) } return mimeType } - diff --git a/internal/apps/user/access_tokens.go b/internal/apps/user/access_tokens.go index 235aabae..647cfbca 100644 --- a/internal/apps/user/access_tokens.go +++ b/internal/apps/user/access_tokens.go @@ -11,7 +11,6 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-gonic/gin" @@ -43,8 +42,8 @@ func ListAccessTokens(c *gin.Context) { currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() - var tokens []model.AccessToken - if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil { + tokens, err := listAccessTokensLogic(ctx, currUser.ID) + if err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -91,8 +90,8 @@ func CreateAccessToken(c *gin.Context) { maxLimit = val } - var count int64 - if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil { + count, err := countAccessTokensLogic(ctx, currUser.ID) + if err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -120,7 +119,7 @@ func CreateAccessToken(c *gin.Context) { IsAdmin: req.IsAdmin, } - if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil { + if err := createAccessTokenLogic(ctx, &tokenRecord); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -152,14 +151,8 @@ func DeleteAccessToken(c *gin.Context) { return } - tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{}) - if tx.Error != nil { - response.AbortBadRequest(c, tx.Error.Error()) - return - } - - if tx.RowsAffected == 0 { - response.AbortBadRequest(c, errTokenNotFoundOrForbidden) + if err := deleteAccessTokenLogic(ctx, id, currUser.ID); err != nil { + response.AbortBadRequest(c, err.Error()) return } @@ -187,32 +180,14 @@ func RotateAccessToken(c *gin.Context) { return } - var tokenRecord model.AccessToken - if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil { - response.AbortBadRequest(c, errTokenNotFoundOrForbidden) - return - } - - // 生成新的 Token - newTokenStr, err := model.GenerateTokenString() + newTokenStr, tokenRecord, err := rotateAccessTokenLogic(ctx, id, currUser.ID) if err != nil { - response.AbortBadRequest(c, errGenerateTokenFailed) - return - } - - newTokenHash := model.HashToken(newTokenStr) - newMaskedToken := model.MaskTokenString(newTokenStr) - - tokenRecord.TokenHash = newTokenHash - tokenRecord.MaskedToken = newMaskedToken - - if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil { response.AbortBadRequest(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(tokenResponse{ Token: newTokenStr, - Record: tokenRecord, + Record: *tokenRecord, })) } diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 1fc9db2f..5d6e7b88 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -12,6 +12,7 @@ import ( "math/big" "strings" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" @@ -285,3 +286,124 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn } return &dbUser, nil } + +func getUserByUsernameOrEmail(ctx context.Context, input string) (*model.User, error) { + var user model.User + if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil { + return nil, err + } + return &user, nil +} + +func updateLastLogin(ctx context.Context, user *model.User) error { + return db.DB(ctx).Model(user).Update("last_login_at", user.LastLoginAt).Error +} + +func registerUserLogic(ctx context.Context, u *model.User) error { + if err := u.RegisterUser(ctx, db.DB(ctx)); err != nil { + if strings.Contains(err.Error(), "duplicate key") || strings.Contains(err.Error(), "UNIQUE") { + return errors.New("用户名或邮箱已被占用") + } + return errors.New("注册失败,请稍后再试") + } + return nil +} + +func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass string) error { + var dbUser model.User + if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil { + return errors.New(errUserNotFound) + } + + if !dbUser.CheckPassword(oldPass) { + return errors.New(errOldPasswordIncorrect) + } + + if err := dbUser.SetEncryptedPassword(newPass); err != nil { + return errors.New(errPasswordEncryptFailed) + } + + if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil { + return errors.New("更新密码失败,请稍后再试") + } + + // 吊销该用户所有的 Access Token + var tokens []model.AccessToken + if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Find(&tokens).Error; err == nil { + for _, token := range tokens { + oauth.InvalidateCachedToken(ctx, token.TokenHash) + } + } + if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { + return errors.New("吊销 Access Token 失败,请稍后再试") + } + + oauth.InvalidateCachedUser(ctx, dbUser.ID) + return nil +} + +func listAccessTokensLogic(ctx context.Context, userID uint64) ([]model.AccessToken, error) { + var tokens []model.AccessToken + if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil { + return nil, errors.New("获取令牌列表失败,请稍后再试") + } + return tokens, nil +} + +func countAccessTokensLogic(ctx context.Context, userID uint64) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil { + return 0, errors.New("查询令牌数量失败,请稍后再试") + } + return count, nil +} + +func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) error { + if err := db.DB(ctx).Create(record).Error; err != nil { + return errors.New("创建令牌失败,请稍后再试") + } + return nil +} + +func deleteAccessTokenLogic(ctx context.Context, id, userID uint64) error { + var tokenRecord model.AccessToken + if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil { + return errors.New(errTokenNotFoundOrForbidden) + } + oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) + + tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{}) + if tx.Error != nil { + return errors.New("删除令牌失败,请稍后再试") + } + if tx.RowsAffected == 0 { + return errors.New(errTokenNotFoundOrForbidden) + } + return nil +} + +func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *model.AccessToken, error) { + var tokenRecord model.AccessToken + if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil { + return "", nil, errors.New(errTokenNotFoundOrForbidden) + } + + oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) + + newTokenStr, err := model.GenerateTokenString() + if err != nil { + return "", nil, errors.New(errGenerateTokenFailed) + } + + newTokenHash := model.HashToken(newTokenStr) + newMaskedToken := model.MaskTokenString(newTokenStr) + + tokenRecord.TokenHash = newTokenHash + tokenRecord.MaskedToken = newMaskedToken + + if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil { + return "", nil, errors.New("轮换令牌失败,请稍后再试") + } + + return newTokenStr, &tokenRecord, nil +} diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index b404281a..b3903e79 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -13,7 +13,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/common" "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/config" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/listener" "github.com/Rain-kl/Wavelet/internal/model" @@ -117,8 +116,8 @@ func Login(c *gin.Context) { return } - var user model.User - if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil { + user, err := getUserByUsernameOrEmail(ctx, req.Username) + if err != nil { logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP()) response.AbortBadRequest(c, errUsernameOrPasswordWrong) return @@ -139,7 +138,7 @@ func Login(c *gin.Context) { } if isEmailLoginVerificationEnabled(ctx) { - result, err := processLoginEmailVerification(ctx, req.Code, &user) + result, err := processLoginEmailVerification(ctx, req.Code, user) if err != nil { response.AbortBadRequest(c, err.Error()) return @@ -160,20 +159,20 @@ func Login(c *gin.Context) { } user.LastLoginAt = time.Now() - if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil { - response.AbortBadRequest(c, err.Error()) + if err := updateLastLogin(ctx, user); err != nil { + response.AbortBadRequest(c, "更新登录时间失败,请稍后再试") return } - if err := setLoginSession(ctx, c, &user); err != nil { + if err := setLoginSession(ctx, c, user); err != nil { response.AbortBadRequest(c, errSaveSessionFailed) return } logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP()) + listener.EmitAdminLoggedIn(ctx, user, c.ClientIP()) - c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&user, needChangePassword))) + c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(user, needChangePassword))) } // Register 用户注册 @@ -247,7 +246,7 @@ func Register(c *gin.Context) { return } - if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil { + if err := registerUserLogic(ctx, &user); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -275,6 +274,13 @@ func Logout(c *gin.Context) { username := session.Get(oauth.UserNameKey) if userID != nil { logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + if id, ok := userID.(uint64); ok { + oauth.InvalidateCachedUser(c.Request.Context(), id) + } else if idFloat, ok := userID.(float64); ok { + oauth.InvalidateCachedUser(c.Request.Context(), uint64(idFloat)) + } else if idInt, ok := userID.(int); ok && idInt >= 0 { + oauth.InvalidateCachedUser(c.Request.Context(), uint64(idInt)) + } } session.Options(oauth.GetSessionOptions(-1)) session.Clear() @@ -327,35 +333,11 @@ func ChangePassword(c *gin.Context) { } ctx := c.Request.Context() - var dbUser model.User - if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { - response.AbortBadRequest(c, errUserNotFound) - return - } - - // 校验旧密码 - if !dbUser.CheckPassword(req.OldPassword) { - response.AbortBadRequest(c, errOldPasswordIncorrect) - return - } - - // 加密并更新为新密码 - if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil { - response.AbortBadRequest(c, errPasswordEncryptFailed) - return - } - - if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil { + if err := changePasswordLogic(ctx, userObj.ID, req.OldPassword, req.NewPassword); err != nil { response.AbortBadRequest(c, err.Error()) return } - // 吊销该用户所有的 Access Token - if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { - response.AbortBadRequest(c, "吊销 Access Token 失败: "+err.Error()) - return - } - // 销毁当前活跃会话以强制重新登录 session := sessions.Default(c) session.Clear() @@ -431,6 +413,7 @@ func UpdateProfile(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } + oauth.InvalidateCachedUser(ctx, userObj.ID) session := sessions.Default(c) needChange := session.Get("need_change_password") == true diff --git a/internal/db/migrator/clickhouse.go b/internal/db/migrator/clickhouse.go index d301ba50..dc08de00 100644 --- a/internal/db/migrator/clickhouse.go +++ b/internal/db/migrator/clickhouse.go @@ -15,6 +15,7 @@ import ( "github.com/ClickHouse/clickhouse-go/v2" "github.com/Rain-kl/Wavelet/internal/config" "github.com/pressly/goose/v3" + "github.com/pressly/goose/v3/database" ) const ( @@ -63,11 +64,17 @@ func MigrateClickHouse() { log.Fatalf("[ClickHouse] get sub fs failed: %v\n", err) } + store, err := database.NewStore(database.DialectClickHouse, clickhouseGooseVersionTable) + if err != nil { + closeClickHouseDB(sqlDB) + log.Fatalf("[ClickHouse] create goose store failed: %v\n", err) + } + provider, err := goose.NewProvider( "clickhouse", sqlDB, subFS, - goose.WithTableName(clickhouseGooseVersionTable), + goose.WithStore(store), goose.WithDisableGlobalRegistry(true), ) if err != nil { diff --git a/internal/repository/system_config_cache.go b/internal/repository/system_config_cache.go index f983f928..f149f17b 100644 --- a/internal/repository/system_config_cache.go +++ b/internal/repository/system_config_cache.go @@ -30,8 +30,10 @@ type systemConfigInvalidationMessage struct { } var ( - systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize}) - systemConfigListenerOnce sync.Once + systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize}) + systemConfigListenerOnce sync.Once + systemConfigListenerCtx context.Context + systemConfigListenerCancel context.CancelFunc ) func ensureSystemConfigCacheListener() { @@ -43,12 +45,19 @@ func startSystemConfigCacheInvalidationListener() { return } + systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background()) + go func() { - pubsub := db.Redis.Subscribe(context.Background(), SystemConfigInvalidationChannel) + pubsub := db.Redis.Subscribe(systemConfigListenerCtx, SystemConfigInvalidationChannel) defer func() { _ = pubsub.Close() }() + go func() { + <-systemConfigListenerCtx.Done() + _ = pubsub.Close() + }() + for msg := range pubsub.Channel() { var payload systemConfigInvalidationMessage if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil { @@ -64,6 +73,15 @@ func startSystemConfigCacheInvalidationListener() { }() } +// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard. +func StopSystemConfigCacheListener() { + if systemConfigListenerCancel != nil { + systemConfigListenerCancel() + systemConfigListenerCancel = nil + } + systemConfigListenerOnce = sync.Once{} +} + func cloneSystemConfig(sc model.SystemConfig) model.SystemConfig { return sc } @@ -117,4 +135,4 @@ func InvalidateAllSystemConfigCaches(ctx context.Context) error { // ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache. func ResetSystemConfigRAMCacheForTest() { systemConfigRAMCache.InvalidateAll() -} \ No newline at end of file +}