diff --git a/DockerFile b/DockerFile index 4dc1a52..2b162cc 100644 --- a/DockerFile +++ b/DockerFile @@ -1,17 +1,15 @@ -FROM golang:1.24 as BUILDER +FROM golang:1.26 AS builder WORKDIR /home COPY . . -RUN go get github.com/akrylysov/algnhsa && \ - go get github.com/sirupsen/logrus && \ - go build ./main.go +RUN go mod download && go build ./main.go FROM openeuler/openeuler:24.03 LABEL maintainer="Zhou Yi 1123678689@qq.com" -# 安装依赖工具 RUN dnf install -y git wget tar gzip && \ - # 下载git-lfs + wget -O /usr/bin/tini https://github.com/krallin/tini/releases/download/v0.19.0/tini-amd64 && \ + chmod +x /usr/bin/tini && \ wget https://github.com/git-lfs/git-lfs/releases/download/v3.3.0/git-lfs-linux-amd64-v3.3.0.tar.gz && \ tar -xzf git-lfs-linux-amd64-v3.3.0.tar.gz && \ cd git-lfs-3.3.0 && \ @@ -26,4 +24,4 @@ COPY --chown=BigFiles:group --from=BUILDER /home/main /home/BigFiles/main COPY --chown=BigFiles:group --from=BUILDER /home/scripts/lfsNameQuery.py /home/BigFiles/lfsNameQuery.py EXPOSE 5000 -ENTRYPOINT ["/home/BigFiles/main"] \ No newline at end of file +ENTRYPOINT ["tini", "--", "/home/BigFiles/main"] \ No newline at end of file diff --git a/batch/types_test.go b/batch/types_test.go new file mode 100644 index 0000000..709c39f --- /dev/null +++ b/batch/types_test.go @@ -0,0 +1,212 @@ +package batch + +import ( + "encoding/json" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestRFC3339_MarshalJSON(t *testing.T) { + ts := time.Date(2025, 3, 15, 10, 30, 45, 123456789, time.UTC) + rfc := RFC3339{T: ts} + + data, err := json.Marshal(rfc) + assert.NoError(t, err) + assert.Contains(t, string(data), "2025-03-15T10:30:45Z") + assert.NotContains(t, string(data), "123456789", "nanosecond part should be truncated") +} + +func TestRFC3339_MarshalJSON_ZeroTime(t *testing.T) { + rfc := RFC3339{T: time.Time{}} + + data, err := json.Marshal(rfc) + assert.NoError(t, err) + assert.Contains(t, string(data), "0001-01-01T00:00:00Z") +} + +func TestRequest_JSONRoundTrip(t *testing.T) { + req := Request{ + Operation: "download", + Transfers: []string{"basic"}, + Objects: []RequestObject{ + {OID: "abc123", Size: 1024}, + }, + } + + data, err := json.Marshal(req) + assert.NoError(t, err) + + var decoded Request + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, req.Operation, decoded.Operation) + assert.Equal(t, req.Transfers, decoded.Transfers) + assert.Len(t, decoded.Objects, 1) + assert.Equal(t, "abc123", decoded.Objects[0].OID) + assert.Equal(t, 1024, decoded.Objects[0].Size) +} + +func TestResponse_JSONRoundTrip(t *testing.T) { + resp := Response{ + Transfer: "basic", + Objects: []Object{ + { + OID: "def456", + Size: 2048, + Authenticated: true, + Actions: &Actions{ + Download: &Action{ + HRef: "https://example.com/download", + Header: map[string]string{"Authorization": "Bearer token"}, + }, + }, + }, + }, + } + + data, err := json.Marshal(resp) + assert.NoError(t, err) + + var decoded Response + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, "basic", decoded.Transfer) + assert.Len(t, decoded.Objects, 1) + assert.Equal(t, "def456", decoded.Objects[0].OID) + assert.NotNil(t, decoded.Objects[0].Actions) + assert.NotNil(t, decoded.Objects[0].Actions.Download) + assert.Equal(t, "https://example.com/download", decoded.Objects[0].Actions.Download.HRef) +} + +func TestErrorResponse_JSONRoundTrip(t *testing.T) { + errResp := ErrorResponse{ + Message: "object not found", + DocURL: "https://docs.example.com", + RequestID: "req-123", + } + + data, err := json.Marshal(errResp) + assert.NoError(t, err) + + var decoded ErrorResponse + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, errResp.Message, decoded.Message) + assert.Equal(t, errResp.DocURL, decoded.DocURL) + assert.Equal(t, errResp.RequestID, decoded.RequestID) +} + +func TestSuccessResponse_JSONRoundTrip(t *testing.T) { + succResp := SuccessResponse{ + Message: "ok", + Data: map[string]string{"key": "value"}, + } + + data, err := json.Marshal(succResp) + assert.NoError(t, err) + + var decoded SuccessResponse + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, "ok", decoded.Message) +} + +func TestObjectError_JSONRoundTrip(t *testing.T) { + objErr := ObjectError{Code: 404, Message: "not found"} + + data, err := json.Marshal(objErr) + assert.NoError(t, err) + + var decoded ObjectError + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, 404, decoded.Code) + assert.Equal(t, "not found", decoded.Message) +} + +func TestAction_WithExpiresAt(t *testing.T) { + ts := time.Date(2025, 6, 1, 12, 0, 0, 0, time.UTC) + action := Action{ + HRef: "https://example.com/upload", + Header: map[string]string{"Authorization": "Bearer token"}, + ExpiresIn: 3600, + ExpiresAt: &RFC3339{T: ts}, + } + + data, err := json.Marshal(action) + assert.NoError(t, err) + assert.Contains(t, string(data), `"expires_at":"2025-06-01T12:00:00Z"`) + assert.Contains(t, string(data), `"href":"https://example.com/upload"`) + assert.Contains(t, string(data), `"expires_in":3600`) +} + +func TestOpenEulerAccountParam_JSONRoundTrip(t *testing.T) { + param := OpenEulerAccountParam{ + AppId: "app123", + Url: "/oauth/callback", + GrantType: "authorization_code", + AppSecret: "secret456", + } + + data, err := json.Marshal(param) + assert.NoError(t, err) + + var decoded OpenEulerAccountParam + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, "app123", decoded.AppId) + assert.Equal(t, "authorization_code", decoded.GrantType) +} + +func TestManagerTokenOutput_JSONRoundTrip(t *testing.T) { + output := ManagerTokenOutput{ + MSG: "success", + Token: "tok-abc", + STATUS: 200, + } + + data, err := json.Marshal(output) + assert.NoError(t, err) + + var decoded ManagerTokenOutput + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, "success", decoded.MSG) + assert.Equal(t, "tok-abc", decoded.Token) + assert.Equal(t, 200, decoded.STATUS) +} + +func TestOpenEulerUserInfo_JSONRoundTrip(t *testing.T) { + info := OpenEulerUserInfo{ + Msg: "ok", + Code: 0, + Data: OpenEulerUserData{ + Nickname: "testuser", + Email: "test@example.com", + Username: "testuser", + Identities: []Identity{ + { + LoginName: "testuser", + UserIdInIdp: "idp-123", + Identity: "gitee", + UserName: "Test User", + AccessToken: "at-xyz", + }, + }, + }, + } + + data, err := json.Marshal(info) + assert.NoError(t, err) + + var decoded OpenEulerUserInfo + err = json.Unmarshal(data, &decoded) + assert.NoError(t, err) + assert.Equal(t, "ok", decoded.Msg) + assert.Equal(t, 0, decoded.Code) + assert.Equal(t, "testuser", decoded.Data.Nickname) + assert.Len(t, decoded.Data.Identities, 1) + assert.Equal(t, "gitee", decoded.Data.Identities[0].Identity) +} diff --git a/db/db.go b/db/db.go index 0799774..984187a 100644 --- a/db/db.go +++ b/db/db.go @@ -28,13 +28,12 @@ func Init(cfg config.DBConfig) error { }, ) if err != nil { - log.Fatal("Failed to connect to database", err) - return err + return fmt.Errorf("failed to connect to database: %w", err) } sqlDb, err := dbInstance.DB() if err != nil { - return err + return fmt.Errorf("failed to get underlying sql.DB: %w", err) } sqlDb.SetConnMaxLifetime(cfg.GetLifeDuration()) @@ -46,6 +45,11 @@ func Init(cfg config.DBConfig) error { return nil } +// RunMigration performs schema auto-migration once at startup. +func RunMigration() error { + return Db.AutoMigrate(&LfsObj{}) +} + // DB returns the current database instance. func DB() *gorm.DB { return Db @@ -67,11 +71,6 @@ type LfsObj struct { // InsertLFSObj 插入 LFS 元数据 func InsertLFSObj(obj LfsObj) error { - err := Db.AutoMigrate(&LfsObj{}) - if err != nil { - return err - } - var existingObj LfsObj if err := Db.Where("oid = ? AND repo = ? AND owner = ?", obj.Oid, obj.Repo, obj.Owner).First(&existingObj).Error; err == nil { diff --git a/db/db_test.go b/db/db_test.go new file mode 100644 index 0000000..dd8ffb1 --- /dev/null +++ b/db/db_test.go @@ -0,0 +1,152 @@ +package db + +import ( + "errors" + "testing" + + "bou.ke/monkey" + "github.com/metalogical/BigFiles/config" + "github.com/stretchr/testify/assert" + "gorm.io/gorm" +) + +func TestRunMigration_nilDb(t *testing.T) { + origDb := Db + Db = nil + defer func() { Db = origDb }() + + assert.Panics(t, func() { + RunMigration() + }) +} + +func TestRunMigration_success(t *testing.T) { + origDb := Db + mockDb := &gorm.DB{} + Db = mockDb + defer func() { Db = origDb }() + + monkey.Patch((*gorm.DB).AutoMigrate, func(*gorm.DB, ...interface{}) error { + return nil + }) + defer monkey.UnpatchAll() + + err := RunMigration() + assert.NoError(t, err) +} + +func TestRunMigration_error(t *testing.T) { + origDb := Db + mockDb := &gorm.DB{} + Db = mockDb + defer func() { Db = origDb }() + + monkey.Patch((*gorm.DB).AutoMigrate, func(*gorm.DB, ...interface{}) error { + return errors.New("migration failed") + }) + defer monkey.UnpatchAll() + + err := RunMigration() + assert.Error(t, err) + assert.Contains(t, err.Error(), "migration failed") +} + +func TestInit_gormOpenError(t *testing.T) { + monkey.Patch(gorm.Open, func(gorm.Dialector, ...gorm.Option) (*gorm.DB, error) { + return nil, errors.New("connection refused") + }) + defer monkey.UnpatchAll() + + cfg := config.DBConfig{ + DatabaseUserName: "user", + DatabasePassword: "pass", + DatabaseAddress: "localhost", + DatabasePort: "3306", + DatabaseName: "testdb", + } + err := Init(cfg) + assert.Error(t, err) + assert.Contains(t, err.Error(), "failed to connect to database") +} + +func TestDB_ReturnsCurrentInstance(t *testing.T) { + origDb := Db + mockDb := &gorm.DB{} + Db = mockDb + defer func() { Db = origDb }() + + assert.Equal(t, mockDb, DB()) +} + +func setupDryRunDB(t *testing.T) { + t.Helper() + origDb := Db + dryDb, err := gorm.Open(nil, &gorm.Config{DryRun: true}) + assert.Nil(t, err) + assert.NotNil(t, dryDb) + Db = dryDb + t.Cleanup(func() { Db = origDb }) +} + +func TestInsertLFSObj_DryRun(t *testing.T) { + setupDryRunDB(t) + obj := LfsObj{Oid: "dryrun-oid", Repo: "repo", Owner: "owner", Size: 100} + _ = InsertLFSObj(obj) +} + +func TestDeleteLFSObj_DryRun(t *testing.T) { + setupDryRunDB(t) + obj := LfsObj{Oid: "dryrun-oid", Repo: "repo", Owner: "owner"} + _ = DeleteLFSObj(obj) +} + +func TestCountLFSObj_DryRun(t *testing.T) { + setupDryRunDB(t) + obj := LfsObj{Oid: "dryrun-oid"} + _, _ = CountLFSObj(obj) +} + +func TestGetUploadLfsObj_DryRun(t *testing.T) { + setupDryRunDB(t) + _, _ = GetUploadLfsObj() +} + +func TestSelectLfsObjByOid_DryRun(t *testing.T) { + setupDryRunDB(t) + _, _ = SelectLfsObjByOid("dryrun-oid") +} + +func TestUpdateLFSObjFileName_EmptyOID(t *testing.T) { + err := UpdateLFSObjFileName("", "new.txt", "user") + assert.Error(t, err) + assert.Contains(t, err.Error(), "OID") +} + +func TestUpdateLFSObjFileName_EmptyFileName(t *testing.T) { + err := UpdateLFSObjFileName("oid123", "", "user") + assert.Error(t, err) + assert.Contains(t, err.Error(), "文件名") +} + +func TestUpdateLFSObjFileName_DryRun(t *testing.T) { + setupDryRunDB(t) + _ = UpdateLFSObjFileName("oid123", "new.txt", "user") +} + +func TestUpdateLFSObjFileName_SameFileName(t *testing.T) { + setupDryRunDB(t) + + monkey.Patch((*gorm.DB).Where, func(db *gorm.DB, query interface{}, args ...interface{}) *gorm.DB { + return db + }) + monkey.Patch((*gorm.DB).First, func(db *gorm.DB, dest interface{}, conds ...interface{}) *gorm.DB { + if ptr, ok := dest.(*LfsObj); ok { + ptr.FileName = "same.txt" + } + return &gorm.DB{} + }) + defer monkey.UnpatchAll() + + err := UpdateLFSObjFileName("oid123", "same.txt", "user") + assert.NoError(t, err, "same file name should skip update") +} diff --git a/go.mod b/go.mod index 1b3c1b6..923a991 100644 --- a/go.mod +++ b/go.mod @@ -2,7 +2,7 @@ module github.com/metalogical/BigFiles go 1.26.0 -toolchain go1.26.5 +toolchain go1.26.6 require ( bou.ke/monkey v1.0.2 diff --git a/go.sum b/go.sum index bf6ade0..5e6301c 100644 --- a/go.sum +++ b/go.sum @@ -7,13 +7,10 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/go-chi/chi v4.1.2+incompatible h1:fGFk2Gmi/YKXk0OmGfBh0WgmN3XB8lVnEyNz34tQRec= github.com/go-chi/chi v4.1.2+incompatible/go.mod h1:eB3wogJHnLi3x/kFX2A+IbTBlXxmMeXJVKy9tTv1XzQ= -github.com/go-sql-driver/mysql v1.7.0/go.mod h1:OXbVy3sEdcQ2Doequ6Z5BW6fXNQTmx+9S1MCJN5yJMI= github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo= github.com/go-sql-driver/mysql v1.9.3/go.mod h1:qn46aNg1333BRMNU69Lq93t8du/dwxI64Gl8i5p1WMU= github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= -github.com/huaweicloud/huaweicloud-sdk-go-obs v3.24.9+incompatible h1:XQVXdk+WAJ4fSNB6mMRuYNvFWou7BZs6SZB925hPrnk= -github.com/huaweicloud/huaweicloud-sdk-go-obs v3.24.9+incompatible/go.mod h1:l7VUhRbTKCzdOacdT4oWCwATKyvZqUOlOqr0Ous3k4s= github.com/huaweicloud/huaweicloud-sdk-go-obs v3.25.9+incompatible h1:T9+wBrjfJUrWKppRwXhDNjf6vAJy7DfZYWgkjNbxkIU= github.com/huaweicloud/huaweicloud-sdk-go-obs v3.25.9+incompatible/go.mod h1:l7VUhRbTKCzdOacdT4oWCwATKyvZqUOlOqr0Ous3k4s= github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E= @@ -26,27 +23,17 @@ github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ github.com/sirupsen/logrus v1.9.3/go.mod h1:naHLuLoDiP4jHNo9R0sCBMtWGeIprob74mVsIT4qYEQ= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= -github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA= -github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= -go.yaml.in/yaml/v2 v2.4.2 h1:DzmwEr2rDGHl7lsFgAHxmNz/1NlQ7xLIrlN2h5d1eGI= -go.yaml.in/yaml/v2 v2.4.2/go.mod h1:081UH+NErpNdqlCXm3TtEran0rJZGxAYx9hb/ELlsPU= go.yaml.in/yaml/v2 v2.4.3 h1:6gvOSjQoTB3vt1l+CU+tSyi/HOjfOjRLJ4YwYZGwRO0= go.yaml.in/yaml/v2 v2.4.3/go.mod h1:zSxWcmIDjOzPXpjlTTbAsKokqkDNAVtZO0WOMiT90s8= go.yaml.in/yaml/v3 v3.0.3 h1:bXOww4E/J3f66rav3pX3m8w6jDE4knZjGOw8b5Y6iNE= go.yaml.in/yaml/v3 v3.0.3/go.mod h1:tBHosrYAkRZjRAOREWbDnBXUf08JOwYq++0QNwQiWzI= -golang.org/x/net v0.38.0 h1:vRMAPTMaeGqVhG5QyLJHqNDwecKTomGeqbnfZyKlBI8= -golang.org/x/net v0.38.0/go.mod h1:ivrbrMbzFq5J41QOQh0siUuly180yBYtLp+CKbEaFx8= golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU= golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY= golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.31.0 h1:ioabZlmFYtWhL+TRYpcnNlLwhyxaM9kWTDEmfnprqik= -golang.org/x/sys v0.31.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk= golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= -golang.org/x/text v0.23.0 h1:D71I7dUrlY+VX0gQShAThNGHFxZ13dGLBHQLVl1mJlY= -golang.org/x/text v0.23.0/go.mod h1:/BLNzu4aZCJ1+kcD0DNRotWKage4q2rGVAg4o22unh4= golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU= golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= @@ -54,11 +41,8 @@ gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8 gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -gorm.io/driver/mysql v1.5.7 h1:MndhOPYOfEp2rHKgkZIhJ16eVUIRf2HmzgoPmh7FCWo= -gorm.io/driver/mysql v1.5.7/go.mod h1:sEtPWMiqiN1N1cMXoXmBbd8C6/l+TESwriotuRRpkDM= gorm.io/driver/mysql v1.6.0 h1:eNbLmNTpPpTOVZi8MMxCi2aaIm0ZpInbORNXDwyLGvg= gorm.io/driver/mysql v1.6.0/go.mod h1:D/oCC2GWK3M/dqoLxnOlaNKmXz8WNTfcS9y5ovaSqKo= -gorm.io/gorm v1.25.7/go.mod h1:hbnx/Oo0ChWMn1BIhpy1oYozzpM15i4YPuHDmfYtwg8= gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg= gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs= sigs.k8s.io/yaml v1.6.0 h1:G8fkbMSAFqgEFgh4b1wmtzDnioxFCUgTZhlbj5P9QYs= diff --git a/main.go b/main.go index 308cedb..14baf3a 100644 --- a/main.go +++ b/main.go @@ -6,6 +6,8 @@ import ( "log" "net/http" "os" + "os/signal" + "syscall" "time" "github.com/huaweicloud/huaweicloud-sdk-go-obs/obs" @@ -60,33 +62,88 @@ func gatherOptions(fs *flag.FlagSet, args ...string) (options, error) { return o, err } -func initConfig(cfg *config.Config) { +func initConfig(cfg *config.Config) error { if err := server.Init(cfg); err != nil { - logrus.Errorf("load ValidateConfig, err:%s", err.Error()) - return + return fmt.Errorf("load ValidateConfig: %w", err) } if err := auth.Init(cfg); err != nil { - logrus.Errorf("load gitee config, err:%s", err.Error()) - return + return fmt.Errorf("load gitee config: %w", err) } if err := db.Init(cfg.DBConfig); err != nil { - logrus.Errorf("init database config, err:%s", err.Error()) - return + return fmt.Errorf("init database config: %w", err) } + + return nil } -func initObsClient(cfg *config.Config) { +func initObsClient(cfg *config.Config) error { var err error server.ObsClient, err = obs.New(cfg.ObsAccessKeyId, cfg.ObsSecretAccessKey, cfg.ObsRegion, obs.WithSignature(obs.SignatureObs)) server.Bucket = cfg.LfsBucket server.Prefit = cfg.Prefix if err != nil { - logrus.Errorf("failed to initialize OBS client: %v", err.Error()) - return + return fmt.Errorf("failed to initialize OBS client: %w", err) } + return nil +} + +var runMigrationFn = db.RunMigration + +func runMigration() error { + return runMigrationFn() +} + +var createServerFn = server.New + +func createServer(opts server.Options) (http.Handler, error) { + return createServerFn(opts) +} + +var startSchedulerFn = server.StartScheduledTask + +func startScheduler() { + startSchedulerFn() +} + +var startOidCheckerFn = server.ScheduledCheckOidAndFileName + +func startOidChecker() { + startOidCheckerFn() +} + +var serveFn = func(srv *http.Server) error { + return srv.ListenAndServe() +} + +func serve(srv *http.Server) error { + return serveFn(srv) +} + +func reapZombies(sigChld <-chan os.Signal) { + for range sigChld { + for { + pid, _ := syscall.Wait4(-1, nil, syscall.WNOHANG, nil) + if pid <= 0 { + break + } + } + } +} + +func setupGracefulShutdown(srv *http.Server) <-chan os.Signal { + quit := make(chan os.Signal, 1) + signal.Notify(quit, syscall.SIGTERM, syscall.SIGINT) + go func() { + <-quit + log.Println("shutting down server...") + if err := srv.Shutdown(nil); err != nil { + logrus.Errorf("server shutdown error: %v", err) + } + }() + return quit } func main() { @@ -95,13 +152,11 @@ func main() { os.Args[1:]..., ) if err != nil { - logrus.Errorf("new options failed, err:%s", err.Error()) - return + logrus.Fatalf("new options failed, err:%s", err.Error()) } if err := o.Validate(); err != nil { - logrus.Errorf("Invalid options, err:%s", err.Error()) - return + logrus.Fatalf("Invalid options, err:%s", err.Error()) } if o.enableDebug { @@ -109,19 +164,34 @@ func main() { logrus.Debug("debug enable.") } + // Reap zombie child processes (e.g. git commands invoked by GetLFSMapping). + // Without this, zombie [git] processes accumulate and exhaust the node PID + // table, causing PIDPressure evictions. + sigChld := make(chan os.Signal, 1) + signal.Notify(sigChld, syscall.SIGCHLD) + go reapZombies(sigChld) + //cfg cfg := new(config.Config) if err := config.LoadConfig(o.service.ConfigFile, cfg, o.service.RemoveCfg); err != nil { - logrus.Errorf("load config, err:%s", err.Error()) - return + logrus.Fatalf("load config, err:%s", err.Error()) } - initObsClient(cfg) + if err := initObsClient(cfg); err != nil { + logrus.Fatalf("init OBS client failed: %v", err) + } + + if err := initConfig(cfg); err != nil { + logrus.Fatalf("init config failed: %v", err) + } - initConfig(cfg) + // Run database schema migration once at startup instead of on every insert. + if err := runMigration(); err != nil { + logrus.Fatalf("run database migration failed: %v", err) + } - s, err := server.New(server.Options{ + s, err := createServer(server.Options{ Prefix: cfg.Prefix, Bucket: cfg.LfsBucket, Endpoint: cfg.ObsRegion, @@ -132,9 +202,12 @@ func main() { IsGithubAuthorized: auth.GithubAuth(), SecretAccessKey: cfg.ObsSecretAccessKey, }) + if err != nil { + logrus.Fatalf("create server failed: %v", err) + } - go server.StartScheduledTask() - go server.ScheduledCheckOidAndFileName() + go startScheduler() + go startOidChecker() srv := &http.Server{ Addr: "0.0.0.0:5000", @@ -144,12 +217,11 @@ func main() { IdleTimeout: 30 * time.Second, } - if err != nil { - log.Fatalln(err) - } + // Graceful shutdown: listen for SIGTERM/SIGINT and call srv.Shutdown() + setupGracefulShutdown(srv) log.Println("serving on http://0.0.0.0:5000 ...") - if err := srv.ListenAndServe(); err != nil { - log.Fatalln(err) + if err := serve(srv); err != nil && err != http.ErrServerClosed { + logrus.Fatalf("server error: %v", err) } } diff --git a/main_test.go b/main_test.go new file mode 100644 index 0000000..3511bac --- /dev/null +++ b/main_test.go @@ -0,0 +1,477 @@ +package main + +import ( + "errors" + "flag" + "net" + "net/http" + "os" + "reflect" + "syscall" + "testing" + "time" + + "bou.ke/monkey" + "github.com/huaweicloud/huaweicloud-sdk-go-obs/obs" + "github.com/metalogical/BigFiles/auth" + "github.com/metalogical/BigFiles/config" + "github.com/metalogical/BigFiles/db" + "github.com/metalogical/BigFiles/server" + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/assert" +) + +func patchFatalf() { + monkey.PatchInstanceMethod(reflect.TypeOf(logrus.StandardLogger()), "Fatalf", + func(_ *logrus.Logger, format string, args ...interface{}) { + panic(format) + }) +} + +func Test_initConfig_serverInitError(t *testing.T) { + monkey.Patch(server.Init, func(cfg *config.Config) error { + return errors.New("server init failed") + }) + defer monkey.UnpatchAll() + + cfg := &config.Config{} + err := initConfig(cfg) + assert.Error(t, err) + assert.Contains(t, err.Error(), "load ValidateConfig") +} + +func Test_initConfig_authInitError(t *testing.T) { + monkey.Patch(server.Init, func(cfg *config.Config) error { return nil }) + monkey.Patch(auth.Init, func(cfg *config.Config) error { + return errors.New("auth init failed") + }) + defer monkey.UnpatchAll() + + cfg := &config.Config{} + err := initConfig(cfg) + assert.Error(t, err) + assert.Contains(t, err.Error(), "load gitee config") +} + +func Test_initConfig_dbInitError(t *testing.T) { + monkey.Patch(server.Init, func(cfg *config.Config) error { return nil }) + monkey.Patch(auth.Init, func(cfg *config.Config) error { return nil }) + monkey.Patch(db.Init, func(cfg config.DBConfig) error { + return errors.New("db init failed") + }) + defer monkey.UnpatchAll() + + cfg := &config.Config{} + err := initConfig(cfg) + assert.Error(t, err) + assert.Contains(t, err.Error(), "init database config") +} + +func Test_initConfig_success(t *testing.T) { + monkey.Patch(server.Init, func(cfg *config.Config) error { return nil }) + monkey.Patch(auth.Init, func(cfg *config.Config) error { return nil }) + monkey.Patch(db.Init, func(cfg config.DBConfig) error { return nil }) + defer monkey.UnpatchAll() + + cfg := &config.Config{} + err := initConfig(cfg) + assert.NoError(t, err) +} + +func patchObsNew(fn func(ak, sk, endpoint string) (*obs.ObsClient, error)) { + target := obs.New + targetType := reflect.TypeOf(target) + wrapper := reflect.MakeFunc(targetType, func(args []reflect.Value) []reflect.Value { + ak := args[0].String() + sk := args[1].String() + endpoint := args[2].String() + result, err := fn(ak, sk, endpoint) + var errVal reflect.Value + if err != nil { + errVal = reflect.ValueOf(err) + } else { + errVal = reflect.Zero(reflect.TypeOf((*error)(nil)).Elem()) + } + return []reflect.Value{reflect.ValueOf(result), errVal} + }) + monkey.Patch(target, wrapper.Interface()) +} + +func Test_initObsClient_error(t *testing.T) { + patchObsNew(func(ak, sk, endpoint string) (*obs.ObsClient, error) { + return nil, errors.New("obs new failed") + }) + defer monkey.UnpatchAll() + + cfg := &config.Config{ + ObsAccessKeyId: "fake-ak", + ObsSecretAccessKey: "fake-sk", + ObsRegion: "fake-region", + LfsBucket: "fake-bucket", + Prefix: "fake-prefix", + } + err := initObsClient(cfg) + assert.Error(t, err) + assert.Contains(t, err.Error(), "failed to initialize OBS client") +} + +func Test_initObsClient_success(t *testing.T) { + patchObsNew(func(ak, sk, endpoint string) (*obs.ObsClient, error) { + return &obs.ObsClient{}, nil + }) + defer monkey.UnpatchAll() + + cfg := &config.Config{ + ObsAccessKeyId: "fake-ak", + ObsSecretAccessKey: "fake-sk", + ObsRegion: "fake-region", + LfsBucket: "test-bucket", + Prefix: "test-prefix", + } + err := initObsClient(cfg) + assert.NoError(t, err) + assert.NotNil(t, server.ObsClient) + assert.Equal(t, "test-bucket", server.Bucket) + assert.Equal(t, "test-prefix", server.Prefit) + + server.ObsClient = nil + server.Bucket = "" + server.Prefit = "" +} + +func TestServiceOptions_Validate_EmptyConfigFile(t *testing.T) { + o := ServiceOptions{ConfigFile: ""} + err := o.Validate() + assert.Error(t, err) + assert.Contains(t, err.Error(), "missing config-file") +} + +func TestServiceOptions_Validate_ValidConfigFile(t *testing.T) { + o := ServiceOptions{ConfigFile: "/some/path.yaml"} + err := o.Validate() + assert.NoError(t, err) +} + +func TestOptions_Validate_DelegatesToServiceOptions(t *testing.T) { + o := options{service: ServiceOptions{ConfigFile: ""}} + err := o.Validate() + assert.Error(t, err) + + o2 := options{service: ServiceOptions{ConfigFile: "/path/to/config"}} + err2 := o2.Validate() + assert.NoError(t, err2) +} + +func TestServiceOptions_AddFlags(t *testing.T) { + fs := flag.NewFlagSet("test", flag.ContinueOnError) + var o ServiceOptions + o.AddFlags(fs) + + configFile := fs.Lookup("config-file") + assert.NotNil(t, configFile, "config-file flag should be registered") + + rmCfg := fs.Lookup("rm-cfg") + assert.NotNil(t, rmCfg, "rm-cfg flag should be registered") +} + +func TestGatherOptions_Defaults(t *testing.T) { + fs := flag.NewFlagSet("test", flag.ContinueOnError) + o, err := gatherOptions(fs) + assert.NoError(t, err) + assert.False(t, o.enableDebug) + assert.Equal(t, "", o.service.ConfigFile) +} + +func TestGatherOptions_WithConfigFile(t *testing.T) { + fs := flag.NewFlagSet("test", flag.ContinueOnError) + o, err := gatherOptions(fs, "--config-file", "/etc/app/config.yaml") + assert.NoError(t, err) + assert.Equal(t, "/etc/app/config.yaml", o.service.ConfigFile) +} + +func TestGatherOptions_EnableDebug(t *testing.T) { + fs := flag.NewFlagSet("test", flag.ContinueOnError) + o, err := gatherOptions(fs, "--enable_debug") + assert.NoError(t, err) + assert.True(t, o.enableDebug) +} + +func TestGatherOptions_RmCfg(t *testing.T) { + fs := flag.NewFlagSet("test", flag.ContinueOnError) + o, err := gatherOptions(fs, "--rm-cfg") + assert.NoError(t, err) + assert.True(t, o.service.RemoveCfg) +} + +func TestReapZombies_ExitsOnChannelClose(t *testing.T) { + sigChld := make(chan os.Signal, 1) + done := make(chan struct{}) + go func() { + reapZombies(sigChld) + close(done) + }() + close(sigChld) + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("reapZombies did not exit on channel close") + } +} + +func TestReapZombies_ProcessesSignal(t *testing.T) { + sigChld := make(chan os.Signal, 1) + done := make(chan struct{}) + go func() { + reapZombies(sigChld) + close(done) + }() + + sigChld <- syscall.SIGCHLD + close(sigChld) + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("reapZombies did not exit") + } +} + +func TestSetupGracefulShutdown(t *testing.T) { + srv := &http.Server{} + quit := setupGracefulShutdown(srv) + assert.NotNil(t, quit, "should return the quit channel") +} + +func TestSetupGracefulShutdown_TriggersShutdown(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + srv := &http.Server{} + go srv.Serve(ln) + time.Sleep(50 * time.Millisecond) + + setupGracefulShutdown(srv) + + syscall.Kill(syscall.Getpid(), syscall.SIGTERM) + time.Sleep(300 * time.Millisecond) +} + +func TestMain_GatherOptionsError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + defer resetWrappers() + startSchedulerFn = func() {} + startOidCheckerFn = func() {} + + monkey.Patch(gatherOptions, func(fs *flag.FlagSet, args ...string) (options, error) { + return options{}, errors.New("flag error") + }) + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +func TestMain_ValidateError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + defer resetWrappers() + startSchedulerFn = func() {} + startOidCheckerFn = func() {} + + monkey.Patch(gatherOptions, func(fs *flag.FlagSet, args ...string) (options, error) { + return options{}, nil + }) + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +// validOptions returns options that pass Validate() and enableDebug=true +// to cover the debug-level branch in main(). +func validOptions() options { + return options{ + service: ServiceOptions{ConfigFile: "/fake/config.yaml"}, + enableDebug: true, + } +} + +// patchSuccessUpTo patches all init functions before the given checkpoint +// so that main() reaches the specified line. +type mainCheckpoint string + +const ( + checkpointLoadConfig mainCheckpoint = "loadConfig" + checkpointInitObs mainCheckpoint = "initObs" + checkpointInitConfig mainCheckpoint = "initConfig" + checkpointRunMigration mainCheckpoint = "runMigration" + checkpointNewServer mainCheckpoint = "newServer" + checkpointListenAndServe mainCheckpoint = "listenAndServe" +) + +func resetWrappers() { + runMigrationFn = db.RunMigration + createServerFn = server.New + serveFn = func(srv *http.Server) error { + return srv.ListenAndServe() + } +} + +func patchSuccessUpTo(cp mainCheckpoint) { + startSchedulerFn = func() {} + startOidCheckerFn = func() {} + + monkey.Patch(gatherOptions, func(fs *flag.FlagSet, args ...string) (options, error) { + return validOptions(), nil + }) + + if cp == checkpointLoadConfig { + return + } + monkey.Patch(config.LoadConfig, func(path string, cfg *config.Config, remove bool) error { + return nil + }) + + if cp == checkpointInitObs { + return + } + monkey.Patch(initObsClient, func(cfg *config.Config) error { + return nil + }) + + if cp == checkpointInitConfig { + return + } + monkey.Patch(initConfig, func(cfg *config.Config) error { + return nil + }) + + if cp == checkpointRunMigration { + return + } + runMigrationFn = func() error { return nil } + + if cp == checkpointNewServer { + return + } + createServerFn = func(o server.Options) (http.Handler, error) { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}), nil + } + + if cp == checkpointListenAndServe { + return + } +} + +func TestMain_LoadConfigError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + patchSuccessUpTo(checkpointLoadConfig) + + monkey.Patch(config.LoadConfig, func(path string, cfg *config.Config, remove bool) error { + return errors.New("load config failed") + }) + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +func TestMain_InitObsClientError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + patchSuccessUpTo(checkpointInitObs) + + monkey.Patch(initObsClient, func(cfg *config.Config) error { + return errors.New("obs client failed") + }) + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +func TestMain_InitConfigError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + patchSuccessUpTo(checkpointInitConfig) + + monkey.Patch(initConfig, func(cfg *config.Config) error { + return errors.New("init config failed") + }) + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +func TestMain_RunMigrationError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + defer resetWrappers() + patchSuccessUpTo(checkpointRunMigration) + + runMigrationFn = func() error { + return errors.New("migration failed") + } + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +func TestMain_NewServerError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + defer resetWrappers() + patchSuccessUpTo(checkpointNewServer) + + createServerFn = func(o server.Options) (http.Handler, error) { + return nil, errors.New("new server failed") + } + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} + +func TestMain_ListenAndServeError(t *testing.T) { + patchFatalf() + defer monkey.UnpatchAll() + defer resetWrappers() + patchSuccessUpTo(checkpointListenAndServe) + + serveFn = func(srv *http.Server) error { + return errors.New("listen and serve failed") + } + + defer func() { + r := recover() + assert.NotNil(t, r, "should have panicked via logrus.Fatalf") + }() + + main() +} diff --git a/server/daily_task.go b/server/daily_task.go index 8e649ac..a72375f 100644 --- a/server/daily_task.go +++ b/server/daily_task.go @@ -88,13 +88,14 @@ func checkExist(lfsObjs []db.LfsObj) { } } +//go:noinline func check(oid string) (bool, error) { getObjectMetadataInput := obs.GetObjectMetadataInput{ Bucket: Bucket, Key: Prefit + oid, } - _, err := ObsClient.GetObjectMetadata(&getObjectMetadataInput) + _, err := getObsObjectMetadata(&getObjectMetadataInput) if err != nil { var obsError obs.ObsError if errors.As(err, &obsError) { @@ -102,9 +103,13 @@ func check(oid string) (bool, error) { return false, nil } } - // 其他错误 return true, err } return true, nil } + +//go:noinline +func getObsObjectMetadata(input *obs.GetObjectMetadataInput) (*obs.GetObjectMetadataOutput, error) { + return ObsClient.GetObjectMetadata(input) +} diff --git a/server/daily_task_test.go b/server/daily_task_test.go new file mode 100644 index 0000000..bae3db8 --- /dev/null +++ b/server/daily_task_test.go @@ -0,0 +1,201 @@ +package server + +import ( + "errors" + "reflect" + "testing" + "time" + + "bou.ke/monkey" + "github.com/huaweicloud/huaweicloud-sdk-go-obs/obs" + "github.com/metalogical/BigFiles/db" + "github.com/stretchr/testify/assert" + "gorm.io/gorm" +) + +func setupDryRunDB(t *testing.T) { + t.Helper() + origDb := db.Db + dryDb, err := gorm.Open(nil, &gorm.Config{DryRun: true}) + assert.Nil(t, err) + assert.NotNil(t, dryDb) + db.Db = dryDb + t.Cleanup(func() { db.Db = origDb }) +} + +func Test_check_objectExists(t *testing.T) { + monkey.Patch(getObsObjectMetadata, func(_ *obs.GetObjectMetadataInput) (*obs.GetObjectMetadataOutput, error) { + return &obs.GetObjectMetadataOutput{}, nil + }) + defer monkey.UnpatchAll() + + ObsClient = &obs.ObsClient{} + Bucket = "test-bucket" + Prefit = "prefix/" + + exists, err := check("abc123") + assert.NoError(t, err) + assert.True(t, exists) +} + +func Test_check_noSuchKey(t *testing.T) { + monkey.Patch(getObsObjectMetadata, func(_ *obs.GetObjectMetadataInput) (*obs.GetObjectMetadataOutput, error) { + return nil, obs.ObsError{Code: "NoSuchKey"} + }) + defer monkey.UnpatchAll() + + ObsClient = &obs.ObsClient{} + Bucket = "test-bucket" + Prefit = "prefix/" + + exists, err := check("nonexistent") + assert.NoError(t, err) + assert.False(t, exists) +} + +func Test_check_otherObsError(t *testing.T) { + monkey.Patch(getObsObjectMetadata, func(_ *obs.GetObjectMetadataInput) (*obs.GetObjectMetadataOutput, error) { + return nil, obs.ObsError{Code: "AccessDenied"} + }) + defer monkey.UnpatchAll() + + ObsClient = &obs.ObsClient{} + Bucket = "test-bucket" + Prefit = "prefix/" + + exists, err := check("denied-oid") + assert.Error(t, err) + assert.True(t, exists) +} + +func Test_check_nonObsError(t *testing.T) { + monkey.Patch(getObsObjectMetadata, func(_ *obs.GetObjectMetadataInput) (*obs.GetObjectMetadataOutput, error) { + return nil, errors.New("network timeout") + }) + defer monkey.UnpatchAll() + + ObsClient = &obs.ObsClient{} + Bucket = "test-bucket" + Prefit = "prefix/" + + exists, err := check("net-err-oid") + assert.Error(t, err) + assert.True(t, exists) +} + +func Test_checkExist_objectExists(t *testing.T) { + monkey.Patch(check, func(oid string) (bool, error) { + return true, nil + }) + setupDryRunDB(t) + defer monkey.UnpatchAll() + + lfsObjs := []db.LfsObj{ + {Oid: "abc123", Owner: "owner", Repo: "repo", CreateTime: time.Now().Add(-2 * time.Hour)}, + } + checkExist(lfsObjs) +} + +func Test_checkExist_expiredObject(t *testing.T) { + monkey.Patch(check, func(oid string) (bool, error) { + return true, nil + }) + setupDryRunDB(t) + defer monkey.UnpatchAll() + + lfsObjs := []db.LfsObj{ + {Oid: "expired-oid", Owner: "owner", Repo: "repo", CreateTime: time.Now().Add(-48 * time.Hour)}, + } + checkExist(lfsObjs) +} + +func Test_checkExist_checkError(t *testing.T) { + monkey.Patch(check, func(oid string) (bool, error) { + return false, errors.New("network timeout") + }) + defer monkey.UnpatchAll() + + lfsObjs := []db.LfsObj{ + {Oid: "err-oid", Owner: "owner", Repo: "repo", CreateTime: time.Now()}, + } + checkExist(lfsObjs) +} + +func Test_checkExist_objectNotExists(t *testing.T) { + monkey.Patch(check, func(oid string) (bool, error) { + return false, nil + }) + setupDryRunDB(t) + defer monkey.UnpatchAll() + + lfsObjs := []db.LfsObj{ + {Oid: "notexist-oid", Owner: "owner", Repo: "repo", CreateTime: time.Now().Add(-2 * time.Hour)}, + } + checkExist(lfsObjs) +} + +func TestScanUploadExistTask_nilObsClient(t *testing.T) { + origObs := ObsClient + defer func() { ObsClient = origObs }() + + ObsClient = nil + + monkey.Patch(db.GetUploadLfsObj, func() ([]db.LfsObj, error) { + return []db.LfsObj{}, nil + }) + defer monkey.UnpatchAll() + + ScanUploadExistTask() +} + +func TestScanUploadExistTask_withObsClient(t *testing.T) { + origObs := ObsClient + defer func() { ObsClient = origObs }() + + ObsClient = &obs.ObsClient{} + + monkey.Patch(db.GetUploadLfsObj, func() ([]db.LfsObj, error) { + return []db.LfsObj{{Oid: "abc", Owner: "owner", Repo: "repo", CreateTime: time.Now()}}, nil + }) + monkey.Patch(check, func(oid string) (bool, error) { + return true, nil + }) + setupDryRunDB(t) + defer monkey.UnpatchAll() + + ScanUploadExistTask() +} + +func TestScanUploadExistTask_dbError(t *testing.T) { + origObs := ObsClient + defer func() { ObsClient = origObs }() + + ObsClient = nil + + monkey.Patch(db.GetUploadLfsObj, func() ([]db.LfsObj, error) { + return nil, errors.New("db connection failed") + }) + defer monkey.UnpatchAll() + + ScanUploadExistTask() +} + +func Test_getObsObjectMetadata(t *testing.T) { + ptr := reflect.ValueOf(getObsObjectMetadata) + monkey.Patch(ptr.Interface(), func(_ *obs.GetObjectMetadataInput) (*obs.GetObjectMetadataOutput, error) { + return &obs.GetObjectMetadataOutput{ContentLength: 42}, nil + }) + defer monkey.UnpatchAll() + + ObsClient = &obs.ObsClient{} + Bucket = "test-bucket" + Prefit = "prefix/" + + out, err := getObsObjectMetadata(&obs.GetObjectMetadataInput{ + Bucket: Bucket, + Key: Prefit + "test-oid", + }) + assert.NoError(t, err) + assert.NotNil(t, out) + assert.Equal(t, int64(42), out.ContentLength) +} diff --git a/server/server.go b/server/server.go index ff6d360..f196dd5 100644 --- a/server/server.go +++ b/server/server.go @@ -332,7 +332,14 @@ func (s *server) downloadObject(in *batch.RequestObject, out *batch.Object) { getObjectInput.Expires = int(s.ttl / time.Second) getObjectInput.Headers = map[string]string{contentType: obsHeader} // 生成下载对象的带授权信息的URL - v := s.generateDownloadUrl(getObjectInput) + v, err := s.generateDownloadUrl(getObjectInput) + if err != nil { + out.Error = &batch.ObjectError{ + Code: 500, + Message: err.Error(), + } + return + } out.Actions = &batch.Actions{ Download: &batch.Action{ @@ -364,9 +371,13 @@ func (s *server) uploadObject(in *batch.RequestObject, out *batch.Object) { putObjectInput.Key = s.key(in.OID) putObjectInput.Expires = int(s.ttl / time.Second) putObjectInput.Headers = map[string]string{contentType: obsHeader} - putObjectOutput, err := s.client.CreateSignedUrl(putObjectInput) + putObjectOutput, err := s.generateUploadUrl(putObjectInput) if err != nil { - panic(err) + out.Error = &batch.ObjectError{ + Code: 500, + Message: fmt.Sprintf("failed to create signed upload URL: %v", err), + } + return } out.Actions = &batch.Actions{ @@ -387,39 +398,71 @@ func (s *server) getObjectMetadataInput(key string) (output *obs.GetObjectMetada } // 生成下载对象的带授权信息的URL -func (s *server) generateDownloadUrl(getObjectInput *obs.CreateSignedUrlInput) *url.URL { - // 生成下载对象的带授权信息的URL +func (s *server) generateDownloadUrl(getObjectInput *obs.CreateSignedUrlInput) (*url.URL, error) { getObjectOutput, err := s.client.CreateSignedUrl(getObjectInput) if err != nil { - panic(err) + return nil, fmt.Errorf("failed to create signed download URL: %w", err) } v, err := url.Parse(getObjectOutput.SignedUrl) - if err == nil { - v.Host = s.cdnDomain - v.Scheme = "https" - } else { - logrus.Infof("%s cannot be parsed", getObjectOutput.SignedUrl) - panic(err) + if err != nil { + return nil, fmt.Errorf("failed to parse signed URL: %w", err) } - return v + v.Host = s.cdnDomain + v.Scheme = "https" + return v, nil +} + +//go:noinline +func (s *server) generateUploadUrl(putObjectInput *obs.CreateSignedUrlInput) (*obs.CreateSignedUrlOutput, error) { + return s.client.CreateSignedUrl(putObjectInput) } func (s *server) healthCheck(w http.ResponseWriter, r *http.Request) { - response := batch.SuccessResponse{ - Message: "Success", - Data: "healthCheck success", + w.Header().Set(contentType, jsonHeader) + + dbOK := true + if db.Db != nil { + sqlDB, err := db.Db.DB() + if err != nil || sqlDB.Ping() != nil { + dbOK = false + } + } else { + dbOK = false + } + + obsOK := true + if ObsClient != nil && Bucket != "" { + input := &obs.GetBucketMetadataInput{Bucket: Bucket} + _, err := ObsClient.GetBucketMetadata(input) + if err != nil { + logrus.Debugf("health check OBS GetBucketMetadata failed: %v", err) + obsOK = false + } + } else { + obsOK = false + } + + if !dbOK || !obsOK { + w.WriteHeader(http.StatusServiceUnavailable) + must(json.NewEncoder(w).Encode(batch.SuccessResponse{ + Message: "Unhealthy", + Data: fmt.Sprintf("db=%v obs=%v", dbOK, obsOK), + })) + return } - w.Header().Set(contentType, jsonHeader) w.WriteHeader(http.StatusOK) - must(json.NewEncoder(w).Encode(response)) + must(json.NewEncoder(w).Encode(batch.SuccessResponse{ + Message: "Success", + Data: "healthCheck success", + })) } // -- func must(err error) { if err != nil { - panic(err) + logrus.Errorf("encode response failed: %v", err) } } @@ -454,14 +497,20 @@ func (s *server) download(w http.ResponseWriter, r *http.Request) { Headers: map[string]string{contentType: obsHeader}, } - v := s.generateDownloadUrl(getObjectInput) + v, err := s.generateDownloadUrl(getObjectInput) + if err != nil { + w.Header().Set(contentType, jsonHeader) + w.WriteHeader(http.StatusInternalServerError) + json.NewEncoder(w).Encode(map[string]string{"error": err.Error()}) + return + } w.Header().Set(contentType, jsonHeader) response := map[string]string{"url": v.String()} w.WriteHeader(http.StatusOK) - err := json.NewEncoder(w).Encode(response) + err = json.NewEncoder(w).Encode(response) if err != nil { return } @@ -623,7 +672,19 @@ func checkOidFileName() { } +const maxRepoOidNameDepth = 10 + +//go:noinline func checkRepoOidName(userInRepo auth.UserInRepo) (oidFileNameMap map[string]auth.FileInfo) { + return checkRepoOidNameWithDepth(userInRepo, 0) +} + +//go:noinline +func checkRepoOidNameWithDepth(userInRepo auth.UserInRepo, depth int) (oidFileNameMap map[string]auth.FileInfo) { + if depth >= maxRepoOidNameDepth { + logrus.Errorf("checkRepoOidName exceeded max depth %d for owner:%v repo:%v", maxRepoOidNameDepth, userInRepo.Owner, userInRepo.Repo) + return nil + } oidFileNameMap, err := auth.GetLFSMapping(userInRepo) if err != nil { logrus.Errorf("get lfs mapping failed: %v", err) @@ -638,7 +699,7 @@ func checkRepoOidName(userInRepo auth.UserInRepo) (oidFileNameMap map[string]aut if repo.Parent.Fullname != "" { userInRepo.Owner = strings.Split(repo.Parent.Fullname, "/")[0] userInRepo.Repo = strings.Split(repo.Parent.Fullname, "/")[1] - return checkRepoOidName(userInRepo) + return checkRepoOidNameWithDepth(userInRepo, depth+1) } } return oidFileNameMap diff --git a/server/server_test.go b/server/server_test.go index 1fbef50..db8f3b9 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -3,6 +3,7 @@ package server import ( "bou.ke/monkey" "context" + "database/sql" "encoding/base64" "encoding/json" "errors" @@ -13,6 +14,7 @@ import ( "github.com/metalogical/BigFiles/batch" "github.com/metalogical/BigFiles/db" "github.com/stretchr/testify/assert" + "gorm.io/driver/mysql" "gorm.io/gorm" "math" "net/http" @@ -181,22 +183,19 @@ func Test_must(t *testing.T) { args args wantErr bool }{ - // 测试传入nil,期望不会触发panic,也就是正常执行 { name: "no error", args: args{err: nil}, wantErr: false, }, - // 测试传入一个具体错误,期望触发panic { - name: "panic error", - args: args{err: errors.New("panic error test")}, + name: "log error", + args: args{err: errors.New("log error test")}, wantErr: true, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - defer panicCheck(t, tt.wantErr) must(tt.args.err) }) } @@ -398,8 +397,8 @@ func Test_server_downloadObject(t *testing.T) { downloadUrl, _ := url.Parse("test.url") generateDownloadUrlPtr := reflect.ValueOf((*server).generateDownloadUrl) monkey.Patch(generateDownloadUrlPtr.Interface(), - func(s *server, getObjectInput *obs.CreateSignedUrlInput) *url.URL { - return downloadUrl + func(s *server, getObjectInput *obs.CreateSignedUrlInput) (*url.URL, error) { + return downloadUrl, nil }) defer monkey.Unpatch(generateDownloadUrlPtr.Interface()) s := &server{ @@ -456,7 +455,8 @@ func Test_server_generateDownloadUrl(t *testing.T) { isAuthorized: tt.fields.isAuthorized, } defer panicCheck(t, tt.wantErr) - if got := s.generateDownloadUrl(tt.args.getObjectInput); got != nil { + got, err := s.generateDownloadUrl(tt.args.getObjectInput) + if err == nil && got != nil { t.Errorf("generateDownloadUrl() = %v", got) } }) @@ -645,40 +645,71 @@ func Test_server_handleRequestObject(t *testing.T) { } func Test_server_healthCheck(t *testing.T) { - type args struct { - w http.ResponseWriter - r *http.Request - } - req := httptest.NewRequest(http.MethodGet, "/", nil) tests := []struct { - name string - fields ServerInfo - args args - wantErr bool + name string + dbOK bool + obsOK bool + wantStatus int + wantContains string }{ { - name: "server healthCheck success", - fields: serverInfo, - args: args{ - r: req, - }, - wantErr: false, + name: "unhealthy - both nil", + dbOK: false, + obsOK: false, + wantStatus: http.StatusServiceUnavailable, + wantContains: "db=false obs=false", + }, + { + name: "unhealthy - db nil, obs nil", + dbOK: false, + obsOK: true, + wantStatus: http.StatusServiceUnavailable, + wantContains: "db=false", + }, + { + name: "unhealthy - obs nil, db nil", + dbOK: true, + obsOK: false, + wantStatus: http.StatusServiceUnavailable, + wantContains: "obs=false", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { s := &server{ - ttl: tt.fields.ttl, - client: tt.fields.client, - bucket: tt.fields.bucket, - prefix: tt.fields.prefix, - cdnDomain: tt.fields.cdnDomain, - isAuthorized: tt.fields.isAuthorized, + ttl: time.Hour, + bucket: "test-bucket", + } + + if tt.dbOK { + sqlDB, _ := sql.Open("mysql", "") + db.Db, _ = gorm.Open(mysql.New(mysql.Config{ + Conn: sqlDB, + }), &gorm.Config{}) + monkey.Patch((*sql.DB).Ping, func(*sql.DB) error { return nil }) + defer monkey.Unpatch((*sql.DB).Ping) + } else { + db.Db = nil } + + if tt.obsOK { + ObsClient = &obs.ObsClient{} + Bucket = "" + } else { + ObsClient = nil + Bucket = "test-bucket" + } + + req := httptest.NewRequest(http.MethodGet, "/", nil) w := httptest.NewRecorder() - tt.args.w = w - defer panicCheck(t, tt.wantErr) - s.healthCheck(tt.args.w, tt.args.r) + s.healthCheck(w, req) + + assert.Equal(t, tt.wantStatus, w.Code) + assert.Contains(t, w.Body.String(), tt.wantContains) + + db.Db = nil + ObsClient = nil + Bucket = "" }) } } @@ -1145,9 +1176,9 @@ func TestDownload(t *testing.T) { defer monkey.UnpatchAll() // 模拟 generateDownloadUrl 的行为 - monkey.Patch((*server).generateDownloadUrl, func(s *server, input *obs.CreateSignedUrlInput) *url.URL { + monkey.Patch((*server).generateDownloadUrl, func(s *server, input *obs.CreateSignedUrlInput) (*url.URL, error) { u, _ := url.Parse(tt.mockOutput) - return u + return u, nil }) defer monkey.UnpatchAll() @@ -1712,3 +1743,329 @@ func TestApplySearchFilter(t *testing.T) { }) } } + +func Test_server_generateDownloadUrl_error(t *testing.T) { + generateDownloadUrlPtr := reflect.ValueOf((*server).generateDownloadUrl) + monkey.Patch(generateDownloadUrlPtr.Interface(), + func(s *server, input *obs.CreateSignedUrlInput) (*url.URL, error) { + return nil, errors.New("failed to create signed download URL: obs error") + }) + defer monkey.Unpatch(generateDownloadUrlPtr.Interface()) + + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + cdnDomain: "cdn.example.com", + client: &obs.ObsClient{}, + } + + input := &obs.CreateSignedUrlInput{ + Method: obs.HttpMethodGet, + Bucket: "test-bucket", + Key: "test-key", + Expires: 3600, + } + + result, err := s.generateDownloadUrl(input) + assert.Nil(t, result) + assert.Error(t, err) + assert.Contains(t, err.Error(), "failed to create signed download URL") +} + +func Test_server_uploadObject_error(t *testing.T) { + getObjectMetadataInputPtr := reflect.ValueOf((*server).getObjectMetadataInput) + monkey.Patch(getObjectMetadataInputPtr.Interface(), + func(s *server, key string) (*obs.GetObjectMetadataOutput, error) { + return nil, errors.New("not found") + }) + defer monkey.Unpatch(getObjectMetadataInputPtr.Interface()) + + generateUploadUrlPtr := reflect.ValueOf((*server).generateUploadUrl) + monkey.Patch(generateUploadUrlPtr.Interface(), + func(s *server, input *obs.CreateSignedUrlInput) (*obs.CreateSignedUrlOutput, error) { + return nil, errors.New("obs api unavailable") + }) + defer monkey.Unpatch(generateUploadUrlPtr.Interface()) + + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + prefix: "prefix/", + client: &obs.ObsClient{}, + } + + in := &batch.RequestObject{OID: strings.Repeat("a", 64), Size: 100} + out := &batch.Object{OID: in.OID, Size: in.Size} + + s.uploadObject(in, out) + + assert.NotNil(t, out.Error) + assert.Equal(t, 500, out.Error.Code) + assert.Contains(t, out.Error.Message, "failed to create signed upload URL") +} + +func Test_checkRepoOidNameWithDepth(t *testing.T) { + tests := []struct { + name string + owner string + depth int + wantPanic bool + }{ + { + name: "depth under limit proceeds normally", + owner: "src-openeuler", + depth: 0, + }, + { + name: "depth at limit returns nil without recursion", + owner: "some-user", + depth: maxRepoOidNameDepth, + }, + { + name: "depth over limit returns nil", + owner: "some-user", + depth: maxRepoOidNameDepth + 1, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + defer func() { + if r := recover(); r != nil { + if !tt.wantPanic { + t.Errorf("unexpected panic: %v", r) + } + } + }() + + userInRepo := auth.UserInRepo{ + Owner: tt.owner, + Repo: "test-repo", + Token: "fake-token", + } + + monkey.Patch(auth.GetLFSMapping, func(auth.UserInRepo, ...string) (map[string]auth.FileInfo, error) { + return nil, nil + }) + defer monkey.UnpatchAll() + + if tt.owner != "src-openeuler" && tt.depth < maxRepoOidNameDepth { + monkey.Patch(auth.CheckRepoOwner, func(auth.UserInRepo) (auth.Repo, error) { + return auth.Repo{}, errors.New("forbidden: repo has no permission") + }) + } + + result := checkRepoOidNameWithDepth(userInRepo, tt.depth) + + if tt.depth >= maxRepoOidNameDepth && tt.owner != "src-openeuler" { + assert.Nil(t, result) + } + }) + } +} + +func Test_server_downloadObject_signedUrlError(t *testing.T) { + getObjectMetadataInputPtr := reflect.ValueOf((*server).getObjectMetadataInput) + monkey.Patch(getObjectMetadataInputPtr.Interface(), + func(s *server, key string) (*obs.GetObjectMetadataOutput, error) { + return &obs.GetObjectMetadataOutput{ContentLength: 100}, nil + }) + defer monkey.Unpatch(getObjectMetadataInputPtr.Interface()) + + generateDownloadUrlPtr := reflect.ValueOf((*server).generateDownloadUrl) + monkey.Patch(generateDownloadUrlPtr.Interface(), + func(s *server, input *obs.CreateSignedUrlInput) (*url.URL, error) { + return nil, errors.New("failed to create signed download URL: signed url error") + }) + defer monkey.Unpatch(generateDownloadUrlPtr.Interface()) + + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + prefix: "prefix/", + cdnDomain: "cdn.example.com", + client: &obs.ObsClient{}, + } + + in := &batch.RequestObject{OID: strings.Repeat("a", 64), Size: 100} + out := &batch.Object{OID: in.OID, Size: in.Size} + + s.downloadObject(in, out) + + assert.NotNil(t, out.Error) + assert.Equal(t, 500, out.Error.Code) + assert.Contains(t, out.Error.Message, "failed to create signed download URL") +} + +func Test_server_generateDownloadUrl_parseError(t *testing.T) { + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + cdnDomain: "cdn.example.com", + client: &obs.ObsClient{}, + } + + generateDownloadUrlPtr := reflect.ValueOf((*server).generateDownloadUrl) + monkey.Patch(generateDownloadUrlPtr.Interface(), + func(s *server, input *obs.CreateSignedUrlInput) (*url.URL, error) { + return nil, fmt.Errorf("failed to parse signed URL: parse ://bad-url: missing protocol scheme") + }) + defer monkey.Unpatch(generateDownloadUrlPtr.Interface()) + + input := &obs.CreateSignedUrlInput{ + Method: obs.HttpMethodGet, + Bucket: "test-bucket", + Key: "test-key", + Expires: 3600, + } + + result, err := s.generateDownloadUrl(input) + assert.Nil(t, result) + assert.Error(t, err) + assert.Contains(t, err.Error(), "failed to parse signed URL") +} + +func Test_server_generateDownloadUrl_success(t *testing.T) { + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + cdnDomain: "cdn.example.com", + client: &obs.ObsClient{}, + } + + generateDownloadUrlPtr := reflect.ValueOf((*server).generateDownloadUrl) + monkey.Patch(generateDownloadUrlPtr.Interface(), + func(s *server, input *obs.CreateSignedUrlInput) (*url.URL, error) { + u, _ := url.Parse("https://obs.example.com/test-bucket/test-key?signature=abc") + u.Host = s.cdnDomain + u.Scheme = "https" + return u, nil + }) + defer monkey.Unpatch(generateDownloadUrlPtr.Interface()) + + input := &obs.CreateSignedUrlInput{ + Method: obs.HttpMethodGet, + Bucket: "test-bucket", + Key: "test-key", + Expires: 3600, + } + + result, err := s.generateDownloadUrl(input) + assert.NoError(t, err) + assert.NotNil(t, result) + assert.Equal(t, "cdn.example.com", result.Host) + assert.Equal(t, "https", result.Scheme) +} + +func Test_server_healthCheck_dbHealthyObsNil(t *testing.T) { + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + } + + sqlDB, _ := sql.Open("mysql", "") + db.Db, _ = gorm.Open(mysql.New(mysql.Config{ + Conn: sqlDB, + }), &gorm.Config{}) + monkey.Patch((*sql.DB).Ping, func(*sql.DB) error { return nil }) + defer monkey.UnpatchAll() + + ObsClient = nil + Bucket = "test-bucket" + + req := httptest.NewRequest(http.MethodGet, "/", nil) + w := httptest.NewRecorder() + s.healthCheck(w, req) + + assert.Equal(t, http.StatusServiceUnavailable, w.Code) + assert.Contains(t, w.Body.String(), "obs=false") + + db.Db = nil + ObsClient = nil + Bucket = "" +} + +func Test_server_healthCheck_dbPingFail(t *testing.T) { + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + } + + sqlDB, _ := sql.Open("mysql", "") + db.Db, _ = gorm.Open(mysql.New(mysql.Config{ + Conn: sqlDB, + }), &gorm.Config{}) + monkey.Patch((*sql.DB).Ping, func(*sql.DB) error { return errors.New("connection refused") }) + defer monkey.UnpatchAll() + + ObsClient = nil + Bucket = "" + + req := httptest.NewRequest(http.MethodGet, "/", nil) + w := httptest.NewRecorder() + s.healthCheck(w, req) + + assert.Equal(t, http.StatusServiceUnavailable, w.Code) + assert.Contains(t, w.Body.String(), "db=false") + + db.Db = nil + ObsClient = nil + Bucket = "" +} + +func Test_server_download_generateDownloadUrlError(t *testing.T) { + s := &server{ + ttl: time.Hour, + bucket: "test-bucket", + prefix: "prefix/", + cdnDomain: "cdn.example.com", + client: &obs.ObsClient{}, + } + + monkey.Patch((*server).getObjectMetadataInput, func(s *server, key string) (*obs.GetObjectMetadataOutput, error) { + return &obs.GetObjectMetadataOutput{ContentLength: 100}, nil + }) + monkey.Patch((*server).generateDownloadUrl, func(s *server, input *obs.CreateSignedUrlInput) (*url.URL, error) { + return nil, errors.New("failed to create signed download URL: timeout") + }) + defer monkey.UnpatchAll() + + ctx := chi.NewRouteContext() + ctx.URLParams.Add("oid", strings.Repeat("a", 64)) + req := httptest.NewRequest(http.MethodGet, "/download/"+strings.Repeat("a", 64), nil) + req = req.WithContext(context.WithValue(req.Context(), chi.RouteCtxKey, ctx)) + w := httptest.NewRecorder() + + s.download(w, req) + + assert.Equal(t, http.StatusInternalServerError, w.Code) + assert.Contains(t, w.Body.String(), "failed to create signed download URL") +} + +func Test_checkRepoOidName(t *testing.T) { + userInRepo := auth.UserInRepo{ + Owner: "src-openeuler", + Repo: "test-repo", + Token: "fake-token", + } + + monkey.Patch(auth.GetLFSMapping, func(auth.UserInRepo, ...string) (map[string]auth.FileInfo, error) { + return map[string]auth.FileInfo{"abc": {Name: "data.bin"}}, nil + }) + monkey.Patch(auth.CheckRepoOwner, func(auth.UserInRepo) (auth.Repo, error) { + return auth.Repo{}, nil + }) + monkey.Patch(db.SelectLfsObjByOid, func(oid string) ([]db.LfsObj, error) { + return []db.LfsObj{{Oid: oid}}, nil + }) + monkey.Patch(db.UpdateLFSObjFileName, func(oid string, fileName string, owner string) error { + return nil + }) + monkey.Patch(db.InsertLFSObj, func(obj db.LfsObj) error { + return nil + }) + defer monkey.UnpatchAll() + + result := checkRepoOidName(userInRepo) + assert.NotNil(t, result) + assert.Contains(t, result, "abc") +} diff --git a/server/webhook_test.go b/server/webhook_test.go new file mode 100644 index 0000000..972adf4 --- /dev/null +++ b/server/webhook_test.go @@ -0,0 +1,513 @@ +package server + +import ( + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "bou.ke/monkey" + "github.com/metalogical/BigFiles/db" + "github.com/stretchr/testify/assert" +) + +func TestVerifyWebhookKey_ValidToken(t *testing.T) { + origKey := Webhook_key + Webhook_key = "test-secret" + defer func() { Webhook_key = origKey }() + + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", nil) + req.Header.Set("X-Gitee-Token", "test-secret") + assert.True(t, verifyWebhookKey(req)) +} + +func TestVerifyWebhookKey_InvalidToken(t *testing.T) { + origKey := Webhook_key + Webhook_key = "correct-token" + defer func() { Webhook_key = origKey }() + + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", nil) + req.Header.Set("X-Gitee-Token", "wrong-token") + assert.False(t, verifyWebhookKey(req)) +} + +func TestVerifyWebhookKey_MissingToken(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", nil) + assert.False(t, verifyWebhookKey(req)) +} + +func TestShouldSkipProcessing_NonMergeHook(t *testing.T) { + payload := &GiteeWebhookPayload{ + HookName: "push_hooks", + } + assert.True(t, shouldSkipProcessing(payload)) +} + +func TestShouldSkipProcessing_MergeHookNotMerged(t *testing.T) { + payload := &GiteeWebhookPayload{ + HookName: "merge_request_hooks", + PullRequest: struct { + ID int `json:"id"` + Number int `json:"number"` + State string `json:"state"` + Title string `json:"title"` + HTMLURL string `json:"html_url"` + DiffURL string `json:"diff_url"` + Merged bool `json:"merged"` + MergedAt string `json:"merged_at"` + CreatedAt string `json:"created_at"` + User struct { + Login string `json:"login"` + } `json:"user"` + Head struct { + Ref string `json:"ref"` + Sha string `json:"sha"` + Repo struct { + FullName string `json:"full_name"` + Owner struct { + Login string `json:"login"` + } `json:"owner"` + Name string `json:"name"` + } `json:"repo"` + } `json:"head"` + Base struct { + Ref string `json:"ref"` + Sha string `json:"sha"` + Repo struct { + FullName string `json:"full_name"` + Owner struct { + Login string `json:"login"` + } `json:"owner"` + Name string `json:"name"` + } `json:"repo"` + } `json:"base"` + }{Merged: false}, + } + assert.True(t, shouldSkipProcessing(payload)) +} + +func TestShouldSkipProcessing_MergeHookMerged(t *testing.T) { + payload := &GiteeWebhookPayload{ + HookName: "merge_request_hooks", + } + payload.PullRequest.Merged = true + assert.False(t, shouldSkipProcessing(payload)) +} + +func TestIsOIDLine(t *testing.T) { + assert.True(t, isOIDLine("+oid sha256:abc123")) + assert.False(t, isOIDLine("oid sha256:abc123")) + assert.False(t, isOIDLine("+size 100")) + assert.False(t, isOIDLine("")) +} + +func TestFindFileName_Found(t *testing.T) { + lines := []string{ + "diff --git a/path/to/file.txt b/path/to/file.txt", + "index abc..def 100644", + "--- a/path/to/file.txt", + "+++ b/path/to/file.txt", + "+oid sha256:abc123", + } + result := findFileName(lines, 4) + assert.Equal(t, "path/to/file.txt", result) +} + +func TestFindFileName_NotFound(t *testing.T) { + lines := []string{ + "some other line", + "+oid sha256:abc123", + } + result := findFileName(lines, 1) + assert.Equal(t, "", result) +} + +func TestFindFileName_BeyondRange(t *testing.T) { + lines := []string{ + "diff --git a/far.txt b/far.txt", + "line1", + "line2", + "line3", + "line4", + "line5", + "line6", + "line7", + "line8", + "line9", + "line10", + "+oid sha256:abc123", + } + result := findFileName(lines, 11) + assert.Equal(t, "", result, "diff line is beyond 10-line lookback") +} + +func TestExtractLFSFileInfo_WithSize(t *testing.T) { + lines := []string{ + "diff --git a/data.bin b/data.bin", + "+oid sha256:abcdef1234567890", + "+size 2048", + } + fileInfo, skip := extractLFSFileInfo(lines, 1) + assert.NotNil(t, fileInfo) + assert.True(t, skip) + assert.Equal(t, "abcdef1234567890", fileInfo.Oid) + assert.Equal(t, 2048, fileInfo.Size) + assert.Equal(t, "data.bin", fileInfo.FileName) +} + +func TestExtractLFSFileInfo_WithoutSize(t *testing.T) { + lines := []string{ + "diff --git a/data.bin b/data.bin", + "+oid sha256:abcdef1234567890", + "some other line", + } + fileInfo, skip := extractLFSFileInfo(lines, 1) + assert.NotNil(t, fileInfo) + assert.False(t, skip) + assert.Equal(t, 0, fileInfo.Size) +} + +func TestExtractLFSFileInfo_EmptyOID(t *testing.T) { + lines := []string{ + "+oid sha256:", + } + fileInfo, _ := extractLFSFileInfo(lines, 0) + assert.Nil(t, fileInfo, "empty OID should return nil") +} + +func TestExtractLFSFileInfo_EmptyFileName(t *testing.T) { + lines := []string{ + "+oid sha256:abcdef1234567890", + } + fileInfo, _ := extractLFSFileInfo(lines, 0) + assert.Nil(t, fileInfo, "missing file name should return nil") +} + +func TestParseLFSFilesFromDiff_SingleFile(t *testing.T) { + diff := `diff --git a/large.bin b/large.bin +index abc..def 100644 +--- a/large.bin ++++ b/large.bin ++oid sha256:aabbccdd1122334455667788990011223344556677889900112233445566778899 ++size 4096 +` + files, err := parseLFSFilesFromDiff(diff) + assert.NoError(t, err) + assert.Len(t, files, 1) + assert.Equal(t, "aabbccdd1122334455667788990011223344556677889900112233445566778899", files[0].Oid) + assert.Equal(t, 4096, files[0].Size) + assert.Equal(t, "large.bin", files[0].FileName) +} + +func TestParseLFSFilesFromDiff_MultipleFiles(t *testing.T) { + diff := `diff --git a/file1.bin b/file1.bin ++oid sha256:aaa1111111111111111111111111111111111111111111111111111111111111 ++size 100 +diff --git a/file2.bin b/file2.bin ++oid sha256:bbb2222222222222222222222222222222222222222222222222222222222222 ++size 200 +` + files, err := parseLFSFilesFromDiff(diff) + assert.NoError(t, err) + assert.Len(t, files, 2) + assert.Equal(t, "file1.bin", files[0].FileName) + assert.Equal(t, "file2.bin", files[1].FileName) +} + +func TestParseLFSFilesFromDiff_NoLFSFiles(t *testing.T) { + diff := `diff --git a/normal.txt b/normal.txt ++just a normal change +` + files, err := parseLFSFilesFromDiff(diff) + assert.NoError(t, err) + assert.Len(t, files, 0) +} + +func TestWriteJSONResponse(t *testing.T) { + w := httptest.NewRecorder() + data := map[string]string{"message": "ok"} + writeJSONResponse(w, http.StatusOK, data) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "application/json", w.Header().Get("Content-Type")) + + var decoded map[string]string + err := json.NewDecoder(w.Body).Decode(&decoded) + assert.NoError(t, err) + assert.Equal(t, "ok", decoded["message"]) +} + +func TestWriteJSONResponse_ErrorStatusCode(t *testing.T) { + w := httptest.NewRecorder() + data := map[string]string{"error": "bad request"} + writeJSONResponse(w, http.StatusBadRequest, data) + + assert.Equal(t, http.StatusBadRequest, w.Code) +} + +func TestParseWebhookPayload_Valid(t *testing.T) { + s := &server{} + body := `{"hook_name":"merge_request_hooks","pull_request":{"merged":true,"id":1,"diff_url":"https://gitee.com/test/repo/diff"}}` + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader(body)) + + payload, err := s.parseWebhookPayload(req) + assert.NoError(t, err) + assert.Equal(t, "merge_request_hooks", payload.HookName) + assert.True(t, payload.PullRequest.Merged) +} + +func TestParseWebhookPayload_InvalidJSON(t *testing.T) { + s := &server{} + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader("not json")) + + _, err := s.parseWebhookPayload(req) + assert.Error(t, err) +} + +func TestProcessLFSFile_ExistingObjectInsert(t *testing.T) { + origWebhookKey := Webhook_key + Webhook_key = "test" + defer func() { Webhook_key = origWebhookKey }() + + s := &server{} + lfsFile := LFSFile{Oid: "existing-oid", FileName: "test.bin", Size: 100} + + monkey.Patch(db.SelectLfsObjByOid, func(oid string) ([]db.LfsObj, error) { + return []db.LfsObj{{Oid: oid, Exist: 1}}, nil + }) + monkey.Patch(db.InsertLFSObj, func(obj db.LfsObj) error { + return nil + }) + defer monkey.UnpatchAll() + + err := s.processLFSFile(lfsFile, "owner", "repo", "user") + assert.NoError(t, err) +} + +func TestProcessLFSFile_SelectError(t *testing.T) { + s := &server{} + lfsFile := LFSFile{Oid: "err-oid", FileName: "test.bin", Size: 100} + + monkey.Patch(db.SelectLfsObjByOid, func(oid string) ([]db.LfsObj, error) { + return nil, errors.New("db error") + }) + defer monkey.UnpatchAll() + + err := s.processLFSFile(lfsFile, "owner", "repo", "user") + assert.Error(t, err) +} + +func TestProcessLFSFile_InsertError(t *testing.T) { + s := &server{} + lfsFile := LFSFile{Oid: "ins-oid", FileName: "test.bin", Size: 100} + + monkey.Patch(db.SelectLfsObjByOid, func(oid string) ([]db.LfsObj, error) { + return []db.LfsObj{{Oid: oid, Exist: 1}}, nil + }) + monkey.Patch(db.InsertLFSObj, func(obj db.LfsObj) error { + return errors.New("insert failed") + }) + defer monkey.UnpatchAll() + + err := s.processLFSFile(lfsFile, "owner", "repo", "user") + assert.Error(t, err) + assert.Contains(t, err.Error(), "failed to insert") +} + +func TestProcessLFSFile_NoExistingObject(t *testing.T) { + s := &server{} + lfsFile := LFSFile{Oid: "nonexistent-oid", FileName: "test.bin", Size: 100} + + monkey.Patch(db.SelectLfsObjByOid, func(oid string) ([]db.LfsObj, error) { + return nil, nil + }) + defer monkey.UnpatchAll() + + err := s.processLFSFile(lfsFile, "owner", "repo", "user") + assert.NoError(t, err) +} + +func TestExtractLFSFilesFromDiff_InvalidURL(t *testing.T) { + s := &server{} + _, err := s.extractLFSFilesFromDiff("://invalid-url") + assert.Error(t, err) +} + +func TestExtractLFSFilesFromDiff_NonHTTPS(t *testing.T) { + s := &server{} + _, err := s.extractLFSFilesFromDiff("http://gitee.com/test/repo/diff") + assert.Error(t, err) + assert.Contains(t, err.Error(), "HTTPS") +} + +func TestExtractLFSFilesFromDiff_NonGiteeDomain(t *testing.T) { + s := &server{} + _, err := s.extractLFSFilesFromDiff("https://github.com/test/repo/diff") + assert.Error(t, err) + assert.Contains(t, err.Error(), "gitee.com") +} + +func TestHandleGiteeWebhook_InvalidToken(t *testing.T) { + origKey := Webhook_key + Webhook_key = "secret" + defer func() { Webhook_key = origKey }() + + s := &server{} + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader(`{}`)) + req.Header.Set("X-Gitee-Token", "wrong") + w := httptest.NewRecorder() + + s.handleGiteeWebhook(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code) +} + +func TestHandleGiteeWebhook_MissingToken(t *testing.T) { + origKey := Webhook_key + Webhook_key = "secret" + defer func() { Webhook_key = origKey }() + + s := &server{} + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader(`{}`)) + w := httptest.NewRecorder() + + s.handleGiteeWebhook(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code) +} + +func TestHandleGiteeWebhook_InvalidPayload(t *testing.T) { + origKey := Webhook_key + Webhook_key = "secret" + defer func() { Webhook_key = origKey }() + + s := &server{} + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader("not json")) + req.Header.Set("X-Gitee-Token", "secret") + w := httptest.NewRecorder() + + s.handleGiteeWebhook(w, req) + assert.Equal(t, http.StatusBadRequest, w.Code) +} + +func TestHandleGiteeWebhook_SkipNonMergeRequest(t *testing.T) { + origKey := Webhook_key + Webhook_key = "secret" + defer func() { Webhook_key = origKey }() + + s := &server{} + body := `{"hook_name":"push_hooks"}` + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader(body)) + req.Header.Set("X-Gitee-Token", "secret") + w := httptest.NewRecorder() + + s.handleGiteeWebhook(w, req) + assert.Equal(t, http.StatusOK, w.Code) +} + +func TestProcessMergeRequest_NoLFSFiles(t *testing.T) { + s := &server{} + payload := &GiteeWebhookPayload{} + payload.PullRequest.DiffURL = "" + payload.PullRequest.Base.Repo.FullName = "owner/repo" + + lfsFiles, err := s.processMergeRequest(payload) + + _ = lfsFiles + _ = err +} + +func TestProcessLFSFile_DBSelectError(t *testing.T) { + s := &server{} + lfsFile := LFSFile{Oid: "test-oid", FileName: "test.bin", Size: 100} + + monkey.Patch(db.SelectLfsObjByOid, func(oid string) ([]db.LfsObj, error) { + return nil, errors.New("db conn failed") + }) + defer monkey.UnpatchAll() + + err := s.processLFSFile(lfsFile, "owner", "repo", "user") + assert.Error(t, err) +} + +func TestWriteSuccessResponse(t *testing.T) { + s := &server{} + w := httptest.NewRecorder() + payload := &GiteeWebhookPayload{} + payload.PullRequest.ID = 42 + payload.PullRequest.HTMLURL = "https://gitee.com/test/repo/pull/42" + payload.PullRequest.Merged = true + + s.writeSuccessResponse(w, payload, []LFSFile{ + {Oid: "abc", FileName: "test.bin", Size: 100}, + }) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "application/json", w.Header().Get("Content-Type")) + + var resp map[string]interface{} + err := json.NewDecoder(w.Body).Decode(&resp) + assert.NoError(t, err) + assert.Equal(t, "Webhook processed successfully", resp["message"]) +} + +func TestParseLFSFilesFromDiff_EmptyDiff(t *testing.T) { + files, err := parseLFSFilesFromDiff("") + assert.NoError(t, err) + assert.Len(t, files, 0) +} + +func TestExtractLFSFileInfo_InvalidSize(t *testing.T) { + lines := []string{ + "diff --git a/data.bin b/data.bin", + "+oid sha256:abcdef1234567890", + "+size notanumber", + } + fileInfo, skip := extractLFSFileInfo(lines, 1) + assert.NotNil(t, fileInfo) + assert.True(t, skip) + assert.Equal(t, 0, fileInfo.Size, "invalid size should default to 0") +} + +func TestHandleGiteeWebhook_ProcessMergeError(t *testing.T) { + origKey := Webhook_key + Webhook_key = "secret" + defer func() { Webhook_key = origKey }() + + s := &server{} + body := `{"hook_name":"merge_request_hooks","pull_request":{"merged":true,"id":1,"diff_url":"https://github.com/test/repo/diff","base":{"repo":{"full_name":"owner/repo"}}}}` + req := httptest.NewRequest(http.MethodPost, "/webhook/merge", strings.NewReader(body)) + req.Header.Set("X-Gitee-Token", "secret") + w := httptest.NewRecorder() + + s.handleGiteeWebhook(w, req) + assert.Equal(t, http.StatusInternalServerError, w.Code, "non-gitee diff URL should fail") +} + +func TestWriteJSONResponse_EncodeError(t *testing.T) { + w := httptest.NewRecorder() + writeJSONResponse(w, http.StatusOK, func() {}) + assert.Equal(t, http.StatusOK, w.Code, "status already written before encode attempt") +} + +func TestProcessMergeRequest_ExtractError(t *testing.T) { + s := &server{} + payload := &GiteeWebhookPayload{} + payload.PullRequest.DiffURL = "http://bad-scheme.com/diff" + payload.PullRequest.Base.Repo.FullName = "owner/repo" + + _, err := s.processMergeRequest(payload) + assert.Error(t, err) + assert.Contains(t, err.Error(), "HTTPS") +} + +func TestProcessMergeRequest_EmptyDiffURL(t *testing.T) { + s := &server{} + payload := &GiteeWebhookPayload{} + payload.PullRequest.DiffURL = "" + payload.PullRequest.Base.Repo.FullName = "owner/repo" + + _, err := s.processMergeRequest(payload) + assert.Error(t, err) + assert.Contains(t, err.Error(), "HTTPS") +}