diff --git a/.github/workflows/codecov.yml b/.github/workflows/codecov.yml index d56691bfd..4c31de583 100644 --- a/.github/workflows/codecov.yml +++ b/.github/workflows/codecov.yml @@ -12,11 +12,11 @@ jobs: - uses: actions/setup-go@v4 with: go-version: 'stable' - - name: Install dependencies 📦 + - name: Install dependencies run: go mod tidy - - name: Run tests with coverage ✅ + - name: Run tests with coverage run: go test -v -coverprofile="coverage.out" ./... - - name: Upload coverage report to Codecov 📝 + - name: Upload coverage report to Codecov uses: codecov/codecov-action@v3 with: file: ./coverage.out diff --git a/.github/workflows/filesystem.yml b/.github/workflows/filesystem.yml deleted file mode 100644 index f0d831813..000000000 --- a/.github/workflows/filesystem.yml +++ /dev/null @@ -1,40 +0,0 @@ -name: Test -on: - push: - branches: - - master - pull_request: - paths: - - 'filesystem/**' -env: - AWS_ACCESS_KEY_ID: ${{ secrets.AWS_ACCESS_KEY_ID }} - AWS_ACCESS_KEY_SECRET: ${{ secrets.AWS_ACCESS_KEY_SECRET }} - AWS_DEFAULT_REGION: ${{ secrets.AWS_DEFAULT_REGION }} - AWS_BUCKET: ${{ secrets.AWS_BUCKET }} - AWS_URL: ${{ secrets.AWS_URL }} - ALIYUN_ACCESS_KEY_ID: ${{ secrets.ALIYUN_ACCESS_KEY_ID }} - ALIYUN_ACCESS_KEY_SECRET: ${{ secrets.ALIYUN_ACCESS_KEY_SECRET }} - ALIYUN_BUCKET: ${{ secrets.ALIYUN_BUCKET }} - ALIYUN_URL: ${{ secrets.ALIYUN_URL }} - ALIYUN_ENDPOINT: ${{ secrets.ALIYUN_ENDPOINT }} - TENCENT_ACCESS_KEY_ID: ${{ secrets.TENCENT_ACCESS_KEY_ID }} - TENCENT_ACCESS_KEY_SECRET: ${{ secrets.TENCENT_ACCESS_KEY_SECRET }} - TENCENT_BUCKET: ${{ secrets.TENCENT_BUCKET }} - TENCENT_URL: ${{ secrets.TENCENT_URL }} - MINIO_ACCESS_KEY_ID: ${{ secrets.MINIO_ACCESS_KEY_ID }} - MINIO_ACCESS_KEY_SECRET: ${{ secrets.MINIO_ACCESS_KEY_SECRET }} - MINIO_BUCKET: ${{ secrets.MINIO_BUCKET }} - MINIO_URL: ${{ secrets.MINIO_URL }} - MINIO_ENDPOINT: ${{ secrets.MINIO_ENDPOINT }} -jobs: - filesystem: - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@v4 - - uses: actions/setup-go@v4 - with: - go-version: 'stable' - - name: Install dependencies 📦 - run: go mod tidy - - name: Run tests ✅ - run: go test ./filesystem/... diff --git a/.github/workflows/mail.yml b/.github/workflows/mail.yml index a82d20147..307699177 100644 --- a/.github/workflows/mail.yml +++ b/.github/workflows/mail.yml @@ -23,7 +23,7 @@ jobs: - uses: actions/setup-go@v4 with: go-version: 'stable' - - name: Install dependencies 📦 + - name: Install dependencies run: go mod tidy - - name: Run tests ✅ + - name: Run tests run: go test ./mail/... diff --git a/README.md b/README.md index 10d8e4599..c4b9953e9 100644 --- a/README.md +++ b/README.md @@ -42,7 +42,7 @@ Example [https://github.com/goravel/example](https://github.com/goravel/example) | [Grpc](https://www.goravel.dev/the-basics/grpc.html) | [Artisan Console](https://www.goravel.dev/digging-deeper/artisan-console.html) | [Task Scheduling](https://www.goravel.dev/digging-deeper/task-scheduling.html) | [Queue](https://www.goravel.dev/digging-deeper/queues.html) | | [Event](https://www.goravel.dev/digging-deeper/event.html) | [FileStorage](https://www.goravel.dev/digging-deeper/filesystem.html) | [Mail](https://www.goravel.dev/digging-deeper/mail.html) | [Validation](https://www.goravel.dev/the-basics/validation.html) | | [Mock](https://www.goravel.dev/digging-deeper/mock.html) | [Hash](https://www.goravel.dev/security/hashing.html) | [Crypt](https://www.goravel.dev/security/encryption.html) | [Carbon](https://www.goravel.dev/digging-deeper/helpers.html) | -| [Package Development](https://www.goravel.dev/digging-deeper/package-development.html) | [Testing](/testing/getting-started.html) | | | +| [Package Development](https://www.goravel.dev/digging-deeper/package-development.html) | [Testing](https://www.goravel.dev/testing/getting-started.html) | | | ## Roadmap @@ -69,6 +69,14 @@ This project exists thanks to all the people who contribute, to participate in t + + + +## Sponsor + +Better development of the project is inseparable from your support, reward us by [Open Collective](https://opencollective.com/goravel). + +

## Group diff --git a/README_zh.md b/README_zh.md index 84b1cd574..e96bdaed9 100644 --- a/README_zh.md +++ b/README_zh.md @@ -33,14 +33,14 @@ Laravel! ## 主要功能 -| | | | | -| ---------- | -------------- | -------------- | -------------- | -| [自定义配置](https://www.goravel.dev/zh/getting-started/configuration.html) | [HTTP 服务](https://www.goravel.dev/zh/the-basics/routing.html) | [用户认证](https://www.goravel.dev/zh/security/authentication.html) | [用户授权](https://www.goravel.dev/zh/security/authorization.html) | -| [数据库 ORM](https://www.goravel.dev/zh/ORM/getting-started.html) | [数据库迁移](https://www.goravel.dev/zh/ORM/migrations.html) | [日志](https://www.goravel.dev/zh/the-basics/logging.html) | [缓存](https://www.goravel.dev/zh/digging-deeper/cache.html) | -| [Grpc](https://www.goravel.dev/zh/the-basics/grpc.html) | [Artisan 命令行](https://www.goravel.dev/zh/digging-deeper/artisan-console.html) | [任务调度](https://www.goravel.dev/zh/digging-deeper/task-scheduling.html) | [队列](https://www.goravel.dev/zh/digging-deeper/queues.html) | -| [事件系统](https://www.goravel.dev/zh/digging-deeper/event.html) | [文件存储](https://www.goravel.dev/zh/digging-deeper/filesystem.html) | [邮件](https://www.goravel.dev/zh/digging-deeper/mail.html) | [表单验证](https://www.goravel.dev/zh/the-basics/validation.html) | -| [Mock](https://www.goravel.dev/zh/digging-deeper/mock.html) | [Hash](https://www.goravel.dev/zh/security/hashing.html) | [Crypt](https://www.goravel.dev/zh/security/encryption.html) | [Carbon](https://www.goravel.dev/zh/digging-deeper/helpers.html) | -| [扩展包开发](https://www.goravel.dev/zh/digging-deeper/package-development.html) | [测试](/testing/getting-started.html) | | | +| | | | | +|-----------------------------------------------------------------------------|-------------------------------------------------------------------------------|------------------------------------------------------------------------|------------------------------------------------------------------| +| [自定义配置](https://www.goravel.dev/zh/getting-started/configuration.html) | [HTTP 服务](https://www.goravel.dev/zh/the-basics/routing.html) | [用户认证](https://www.goravel.dev/zh/security/authentication.html) | [用户授权](https://www.goravel.dev/zh/security/authorization.html) | +| [数据库 ORM](https://www.goravel.dev/zh/ORM/getting-started.html) | [数据库迁移](https://www.goravel.dev/zh/ORM/migrations.html) | [日志](https://www.goravel.dev/zh/the-basics/logging.html) | [缓存](https://www.goravel.dev/zh/digging-deeper/cache.html) | +| [Grpc](https://www.goravel.dev/zh/the-basics/grpc.html) | [Artisan 命令行](https://www.goravel.dev/zh/digging-deeper/artisan-console.html) | [任务调度](https://www.goravel.dev/zh/digging-deeper/task-scheduling.html) | [队列](https://www.goravel.dev/zh/digging-deeper/queues.html) | +| [事件系统](https://www.goravel.dev/zh/digging-deeper/event.html) | [文件存储](https://www.goravel.dev/zh/digging-deeper/filesystem.html) | [邮件](https://www.goravel.dev/zh/digging-deeper/mail.html) | [表单验证](https://www.goravel.dev/zh/the-basics/validation.html) | +| [Mock](https://www.goravel.dev/zh/digging-deeper/mock.html) | [Hash](https://www.goravel.dev/zh/security/hashing.html) | [Crypt](https://www.goravel.dev/zh/security/encryption.html) | [Carbon](https://www.goravel.dev/zh/digging-deeper/helpers.html) | +| [扩展包开发](https://www.goravel.dev/zh/digging-deeper/package-development.html) | [测试](https://www.goravel.dev/zh/testing/getting-started.html) | | | ## 路线图 @@ -67,6 +67,14 @@ Laravel! + + + +## 打赏 + +开源项目的发展离不开您的支持,感谢微信打赏。 + +

## 群组 diff --git a/auth/auth.go b/auth/auth.go index 01e60ec30..a6b118058 100644 --- a/auth/auth.go +++ b/auth/auth.go @@ -229,9 +229,12 @@ func (a *Auth) Logout(ctx http.Context) error { } func (a *Auth) makeAuthContext(ctx http.Context, claims *Claims, token string) { - ctx.WithValue(ctxKey, Guards{ - a.guard: {claims, token}, - }) + guards, ok := ctx.Value(ctxKey).(Guards) + if !ok { + guards = make(Guards) + } + guards[a.guard] = &Guard{claims, token} + ctx.WithValue(ctxKey, guards) } func (a *Auth) tokenIsDisabled(token string) bool { diff --git a/auth/auth_test.go b/auth/auth_test.go index 50202c3a3..f94b1658d 100644 --- a/auth/auth_test.go +++ b/auth/auth_test.go @@ -21,7 +21,7 @@ import ( "github.com/goravel/framework/support/carbon" ) -var guard = "user" +var testUserGuard = "user" type User struct { orm.Model @@ -107,7 +107,7 @@ func (s *AuthTestSuite) SetupTest() { s.mockConfig = &configmock.Config{} s.mockOrm = &ormmock.Orm{} s.mockDB = &ormmock.Query{} - s.auth = NewAuth(guard, s.mockCache, s.mockConfig, s.mockOrm) + s.auth = NewAuth(testUserGuard, s.mockCache, s.mockConfig, s.mockOrm) } func (s *AuthTestSuite) TestLoginUsingID_EmptySecret() { @@ -256,7 +256,7 @@ func (s *AuthTestSuite) TestParse_TokenExpired() { payload, err := s.auth.Parse(ctx, token) s.Equal(&authcontract.Payload{ - Guard: guard, + Guard: testUserGuard, Key: "1", ExpireAt: jwt.NewNumericDate(expireAt).Local(), IssuedAt: jwt.NewNumericDate(issuedAt).Local(), @@ -269,7 +269,7 @@ func (s *AuthTestSuite) TestParse_TokenExpired() { } func (s *AuthTestSuite) TestParse_InvalidCache() { - auth := NewAuth(guard, nil, s.mockConfig, s.mockOrm) + auth := NewAuth(testUserGuard, nil, s.mockConfig, s.mockOrm) ctx := Background() payload, err := auth.Parse(ctx, "1") s.Nil(payload) @@ -288,7 +288,7 @@ func (s *AuthTestSuite) TestParse_Success() { payload, err := s.auth.Parse(ctx, token) s.Equal(&authcontract.Payload{ - Guard: guard, + Guard: testUserGuard, Key: "1", ExpireAt: jwt.NewNumericDate(carbon.Now().AddMinutes(2).ToStdTime()).Local(), IssuedAt: jwt.NewNumericDate(carbon.Now().ToStdTime()).Local(), @@ -311,7 +311,7 @@ func (s *AuthTestSuite) TestParse_SuccessWithPrefix() { payload, err := s.auth.Parse(ctx, "Bearer "+token) s.Equal(&authcontract.Payload{ - Guard: guard, + Guard: testUserGuard, Key: "1", ExpireAt: jwt.NewNumericDate(carbon.Now().AddMinutes(2).ToStdTime()).Local(), IssuedAt: jwt.NewNumericDate(carbon.Now().ToStdTime()).Local(), @@ -464,6 +464,62 @@ func (s *AuthTestSuite) TestUser_Success() { s.Nil(err) s.mockConfig.AssertExpectations(s.T()) + s.mockCache.AssertExpectations(s.T()) + s.mockOrm.AssertExpectations(s.T()) + s.mockDB.AssertExpectations(s.T()) +} + +func (s *AuthTestSuite) TestUser_Success_MultipleParse() { + testAdminGuard := "admin" + + s.mockConfig.On("GetString", "jwt.secret").Return("Goravel").Twice() + s.mockConfig.On("GetInt", "jwt.ttl").Return(2).Once() + + ctx := Background() + token1, err := s.auth.LoginUsingID(ctx, 1) + s.Nil(err) + + s.mockConfig.On("GetString", "jwt.secret").Return("Goravel").Twice() + s.mockConfig.On("GetInt", "jwt.ttl").Return(2).Once() + + ctx = Background() + token2, err := s.auth.Guard(testAdminGuard).LoginUsingID(ctx, 2) + s.Nil(err) + + s.mockCache.On("GetBool", "jwt:disabled:"+token1, false).Return(false).Once() + + payload, err := s.auth.Parse(ctx, token1) + s.Nil(err) + s.NotNil(payload) + s.Equal(testUserGuard, payload.Guard) + s.Equal("1", payload.Key) + + s.mockCache.On("GetBool", "jwt:disabled:"+token2, false).Return(false).Once() + + payload, err = s.auth.Guard(testAdminGuard).Parse(ctx, token2) + s.Nil(err) + s.NotNil(payload) + s.Equal(testAdminGuard, payload.Guard) + s.Equal("2", payload.Key) + + var user1 User + s.mockOrm.On("Query").Return(s.mockDB) + s.mockDB.On("FindOrFail", &user1, clause.Eq{Column: clause.PrimaryColumn, Value: "1"}).Return(nil).Once() + + err = s.auth.User(ctx, &user1) + s.Nil(err) + + var user2 User + s.mockOrm.On("Query").Return(s.mockDB) + s.mockDB.On("FindOrFail", &user2, clause.Eq{Column: clause.PrimaryColumn, Value: "2"}).Return(nil).Once() + + err = s.auth.Guard(testAdminGuard).User(ctx, &user2) + s.Nil(err) + + s.mockConfig.AssertExpectations(s.T()) + s.mockCache.AssertExpectations(s.T()) + s.mockOrm.AssertExpectations(s.T()) + s.mockDB.AssertExpectations(s.T()) } func (s *AuthTestSuite) TestRefresh_NotParse() { @@ -541,7 +597,7 @@ func (s *AuthTestSuite) TestRefresh_Success() { } func (s *AuthTestSuite) TestLogout_CacheUnsupported() { - s.auth = NewAuth(guard, nil, s.mockConfig, s.mockOrm) + s.auth = NewAuth(testUserGuard, nil, s.mockConfig, s.mockOrm) s.mockConfig.On("GetString", "jwt.secret").Return("Goravel").Once() s.mockConfig.On("GetInt", "jwt.ttl").Return(2).Once() @@ -644,3 +700,19 @@ func (s *AuthTestSuite) TestLogout_Error_TTL_Is_0() { s.mockConfig.AssertExpectations(s.T()) } + +func (s *AuthTestSuite) TestMakeAuthContext() { + testAdminGuard := "admin" + + ctx := Background() + s.auth.makeAuthContext(ctx, nil, "1") + guards, ok := ctx.Value(ctxKey).(Guards) + s.True(ok) + s.Equal(&Guard{nil, "1"}, guards[testUserGuard]) + + s.auth.Guard(testAdminGuard).(*Auth).makeAuthContext(ctx, nil, "2") + guards, ok = ctx.Value(ctxKey).(Guards) + s.True(ok) + s.Equal(&Guard{nil, "1"}, guards[testUserGuard]) + s.Equal(&Guard{nil, "2"}, guards[testAdminGuard]) +} diff --git a/config/application.go b/config/application.go index e472d4544..a7745c38c 100644 --- a/config/application.go +++ b/config/application.go @@ -19,35 +19,31 @@ type Application struct { } func NewApplication(envPath string) *Application { - if !file.Exists(envPath) { - color.Redln("Please create " + envPath + " and initialize it first.") - color.Warnln("Example command: \ncp .env.example .env && go run . artisan key:generate") - os.Exit(0) - } - app := &Application{} app.vip = viper.New() - app.vip.SetConfigType("env") - app.vip.SetConfigFile(envPath) + app.vip.AutomaticEnv() - if err := app.vip.ReadInConfig(); err != nil { - color.Redln("Invalid Config error: " + err.Error()) - os.Exit(0) - } + if file.Exists(envPath) { + app.vip.SetConfigType("env") + app.vip.SetConfigFile(envPath) - app.vip.SetEnvPrefix("goravel") - app.vip.AutomaticEnv() + if err := app.vip.ReadInConfig(); err != nil { + color.Redln("Invalid Config error: " + err.Error()) + os.Exit(0) + } + } appKey := app.Env("APP_KEY") - if support.Env != support.EnvArtisan { + if !support.IsKeyGenerateCommand { if appKey == nil { color.Redln("Please initialize APP_KEY first.") - color.Warnln("Example command: \ngo run . artisan key:generate") + color.Println("Create a .env file and run command: go run . artisan key:generate") + color.Println("Or set a system variable: APP_KEY={32-bit number} go run .") os.Exit(0) } if len(appKey.(string)) != 32 { - color.Redln("Invalid APP_KEY, please reset it.") + color.Redln("Invalid APP_KEY, the length must be 32, please reset it.") color.Warnln("Example command: \ngo run . artisan key:generate") os.Exit(0) } diff --git a/config/application_test.go b/config/application_test.go index 372c771fe..ee331eeac 100644 --- a/config/application_test.go +++ b/config/application_test.go @@ -12,22 +12,32 @@ import ( type ApplicationTestSuite struct { suite.Suite - config *Application - config2 *Application + config *Application + customConfig *Application } func TestApplicationTestSuite(t *testing.T) { - assert.Nil(t, file.Create(".env", "APP_KEY=12345678901234567890123456789012")) + assert.Nil(t, file.Create(".env", ` +APP_KEY=12345678901234567890123456789012 +APP_DEBUG=true +DB_PORT=3306 +`)) temp, err := os.CreateTemp("", "goravel.env") assert.Nil(t, err) + defer temp.Close() defer os.Remove(temp.Name()) - _, err = temp.Write([]byte("APP_KEY=12345678901234567890123456789012")) + + _, err = temp.Write([]byte(` +APP_KEY=12345678901234567890123456789012 +APP_DEBUG=true +DB_PORT=3306 +`)) assert.Nil(t, err) assert.Nil(t, temp.Close()) suite.Run(t, &ApplicationTestSuite{ - config: NewApplication(".env"), - config2: NewApplication(temp.Name()), + config: NewApplication(".env"), + customConfig: NewApplication(temp.Name()), }) assert.Nil(t, file.Remove(".env")) @@ -38,42 +48,44 @@ func (s *ApplicationTestSuite) SetupTest() { } func (s *ApplicationTestSuite) TestEnv() { + s.Equal("12345678901234567890123456789012", s.config.Env("APP_KEY").(string)) s.Equal("goravel", s.config.Env("APP_NAME", "goravel").(string)) - s.Equal("127.0.0.1", s.config.Env("DB_HOST", "127.0.0.1").(string)) - s.Equal("goravel", s.config2.Env("APP_NAME", "goravel").(string)) - s.Equal("127.0.0.1", s.config2.Env("DB_HOST", "127.0.0.1").(string)) + s.Equal("12345678901234567890123456789012", s.customConfig.Env("APP_KEY").(string)) + s.Equal("goravel", s.customConfig.Env("APP_NAME", "goravel").(string)) } func (s *ApplicationTestSuite) TestAdd() { s.config.Add("app", map[string]any{ "env": "local", }) - s.config2.Add("app", map[string]any{ + s.customConfig.Add("app", map[string]any{ "env": "local", }) s.Equal("local", s.config.GetString("app.env")) - s.Equal("local", s.config2.GetString("app.env")) + s.Equal("local", s.customConfig.GetString("app.env")) s.config.Add("path.with.dot.case1", "value1") - s.config2.Add("path.with.dot.case1", "value1") + s.customConfig.Add("path.with.dot.case1", "value1") s.Equal("value1", s.config.GetString("path.with.dot.case1")) - s.Equal("value1", s.config2.GetString("path.with.dot.case1")) + s.Equal("value1", s.customConfig.GetString("path.with.dot.case1")) s.config.Add("path.with.dot.case2", "value2") - s.config2.Add("path.with.dot.case2", "value2") + s.customConfig.Add("path.with.dot.case2", "value2") s.Equal("value2", s.config.GetString("path.with.dot.case2")) - s.Equal("value2", s.config2.GetString("path.with.dot.case2")) + s.Equal("value2", s.customConfig.GetString("path.with.dot.case2")) s.config.Add("path.with.dot", map[string]any{"case3": "value3"}) - s.config2.Add("path.with.dot", map[string]any{"case3": "value3"}) + s.customConfig.Add("path.with.dot", map[string]any{"case3": "value3"}) s.Equal("value3", s.config.GetString("path.with.dot.case3")) - s.Equal("value3", s.config2.GetString("path.with.dot.case3")) + s.Equal("value3", s.customConfig.GetString("path.with.dot.case3")) } func (s *ApplicationTestSuite) TestGet() { + s.Equal("12345678901234567890123456789012", s.config.Get("APP_KEY").(string)) s.Equal("goravel", s.config.Get("APP_NAME", "goravel").(string)) - s.Equal("goravel", s.config2.Get("APP_NAME", "goravel").(string)) + s.Equal("12345678901234567890123456789012", s.customConfig.Get("APP_KEY").(string)) + s.Equal("goravel", s.customConfig.Get("APP_NAME", "goravel").(string)) } func (s *ApplicationTestSuite) TestGetString() { @@ -85,11 +97,11 @@ func (s *ApplicationTestSuite) TestGetString() { }, }, }) - s.config2.Add("database", map[string]any{ - "default": s.config2.Env("DB_CONNECTION", "mysql"), + s.customConfig.Add("database", map[string]any{ + "default": s.customConfig.Env("DB_CONNECTION", "mysql"), "connections": map[string]any{ "mysql": map[string]any{ - "host": s.config2.Env("DB_HOST", "127.0.0.1"), + "host": s.customConfig.Env("DB_HOST", "127.0.0.1"), }, }, }) @@ -97,17 +109,31 @@ func (s *ApplicationTestSuite) TestGetString() { s.Equal("goravel", s.config.GetString("APP_NAME", "goravel")) s.Equal("127.0.0.1", s.config.GetString("database.connections.mysql.host")) s.Equal("mysql", s.config.GetString("database.default")) - s.Equal("goravel", s.config2.GetString("APP_NAME", "goravel")) - s.Equal("127.0.0.1", s.config2.GetString("database.connections.mysql.host")) - s.Equal("mysql", s.config2.GetString("database.default")) + s.Equal("goravel", s.customConfig.GetString("APP_NAME", "goravel")) + s.Equal("127.0.0.1", s.customConfig.GetString("database.connections.mysql.host")) + s.Equal("mysql", s.customConfig.GetString("database.default")) } func (s *ApplicationTestSuite) TestGetInt() { - s.Equal(s.config.GetInt("DB_PORT", 3306), 3306) - s.Equal(s.config2.GetInt("DB_PORT", 3306), 3306) + s.Equal(3306, s.config.GetInt("DB_PORT")) + s.Equal(3306, s.customConfig.GetInt("DB_PORT")) } func (s *ApplicationTestSuite) TestGetBool() { - s.Equal(true, s.config.GetBool("APP_DEBUG", true)) - s.Equal(true, s.config2.GetBool("APP_DEBUG", true)) + s.Equal(true, s.config.GetBool("APP_DEBUG")) + s.Equal(true, s.customConfig.GetBool("APP_DEBUG")) +} + +func TestOsVariables(t *testing.T) { + assert.Nil(t, os.Setenv("APP_KEY", "12345678901234567890123456789013")) + assert.Nil(t, os.Setenv("APP_NAME", "goravel")) + assert.Nil(t, os.Setenv("APP_PORT", "3306")) + assert.Nil(t, os.Setenv("APP_DEBUG", "true")) + + config := NewApplication(".env") + + assert.Equal(t, "12345678901234567890123456789013", config.GetString("APP_KEY")) + assert.Equal(t, "goravel", config.GetString("APP_NAME")) + assert.Equal(t, 3306, config.GetInt("APP_PORT")) + assert.True(t, config.GetBool("APP_DEBUG")) } diff --git a/contracts/auth/access/gate.go b/contracts/auth/access/gate.go index 151134c84..c27290e0b 100644 --- a/contracts/auth/access/gate.go +++ b/contracts/auth/access/gate.go @@ -4,18 +4,29 @@ import "context" //go:generate mockery --name=Gate type Gate interface { + // WithContext returns a new Gate instance with the given context. WithContext(ctx context.Context) Gate + // Allows determines if the given ability should be granted for the current user. Allows(ability string, arguments map[string]any) bool + // Denies determines if the given ability should be denied for the current user. Denies(ability string, arguments map[string]any) bool + // Inspect the given ability against the current user. Inspect(ability string, arguments map[string]any) Response + // Define a new ability. Define(ability string, callback func(ctx context.Context, arguments map[string]any) Response) + // Any one of the given abilities should be granted for the current user. Any(abilities []string, arguments map[string]any) bool + // None of the given abilities should be granted for the current user. None(abilities []string, arguments map[string]any) bool + // Before register a callback to run before all Gate checks. Before(callback func(ctx context.Context, ability string, arguments map[string]any) Response) + // After register a callback to run after all Gate checks. After(callback func(ctx context.Context, ability string, arguments map[string]any, result Response) Response) } type Response interface { + // Allowed to determine if the response was allowed. Allowed() bool + // Message to get the response message. Message() string } diff --git a/contracts/auth/auth.go b/contracts/auth/auth.go index b6d1c6ec2..8b368e336 100644 --- a/contracts/auth/auth.go +++ b/contracts/auth/auth.go @@ -8,12 +8,19 @@ import ( //go:generate mockery --name=Auth type Auth interface { + // Guard attempts to get the guard against the local cache. Guard(name string) Auth + // Parse the given token. Parse(ctx http.Context, token string) (*Payload, error) + // User returns the current authenticated user. User(ctx http.Context, user any) error + // Login logs a user into the application. Login(ctx http.Context, user any) (token string, err error) + // LoginUsingID logs the given user ID into the application. LoginUsingID(ctx http.Context, id any) (token string, err error) + // Refresh the token for the current user. Refresh(ctx http.Context) (token string, err error) + // Logout logs the user out of the application. Logout(ctx http.Context) error } diff --git a/contracts/cache/cache.go b/contracts/cache/cache.go index 697365b78..9f35d80fa 100644 --- a/contracts/cache/cache.go +++ b/contracts/cache/cache.go @@ -13,40 +13,52 @@ type Cache interface { //go:generate mockery --name=Driver type Driver interface { - //Add Driver an item in the cache if the key does not exist. + // Add an item in the cache if the key does not exist. Add(key string, value any, t time.Duration) bool + // Decrement decrements the value of an item in the cache. Decrement(key string, value ...int) (int, error) - //Forever Driver an item in the cache indefinitely. + // Forever add an item in the cache indefinitely. Forever(key string, value any) bool - //Forget Remove an item from the cache. + // Forget removes an item from the cache. Forget(key string) bool - //Flush Remove all items from the cache. + // Flush remove all items from the cache. Flush() bool - //Get Retrieve an item from the cache by key. + // Get retrieve an item from the cache by key. Get(key string, def ...any) any + // GetBool retrieves an item from the cache by key as a boolean. GetBool(key string, def ...bool) bool + // GetInt retrieves an item from the cache by key as an integer. GetInt(key string, def ...int) int + // GetInt64 retrieves an item from the cache by key as a 64-bit integer. GetInt64(key string, def ...int64) int64 + // GetString retrieves an item from the cache by key as a string. GetString(key string, def ...string) string - //Has Check an item exists in the cache. + // Has check an item exists in the cache. Has(key string) bool + // Increment increments the value of an item in the cache. Increment(key string, value ...int) (int, error) + // Lock get a lock instance. Lock(key string, t ...time.Duration) Lock - //Put Driver an item in the cache for a given time. + // Put Driver an item in the cache for a given time. Put(key string, value any, t time.Duration) error - //Pull Retrieve an item from the cache and delete it. + // Pull retrieve an item from the cache and delete it. Pull(key string, def ...any) any - //Remember Get an item from the cache, or execute the given Closure and store the result. + // Remember gets an item from the cache, or execute the given Closure and store the result. Remember(key string, ttl time.Duration, callback func() (any, error)) (any, error) - //RememberForever Get an item from the cache, or execute the given Closure and store the result forever. + // RememberForever get an item from the cache, or execute the given Closure and store the result forever. RememberForever(key string, callback func() (any, error)) (any, error) + // WithContext returns a new Cache instance with the given context. WithContext(ctx context.Context) Driver } //go:generate mockery --name=Lock type Lock interface { + // Block attempt to acquire the lock for the given number of seconds. Block(t time.Duration, callback ...func()) bool + // Get attempts to acquire the lock. Get(callback ...func()) bool + // Release the lock. Release() bool + // ForceRelease releases the lock in disregard of ownership. ForceRelease() bool } diff --git a/contracts/config/config.go b/contracts/config/config.go index 5486b0145..f5ed30d37 100644 --- a/contracts/config/config.go +++ b/contracts/config/config.go @@ -2,16 +2,16 @@ package config //go:generate mockery --name=Config type Config interface { - //Env Get config from env. + // Env get config from env. Env(envName string, defaultValue ...any) any - //Add config to application. + // Add config to application. Add(name string, configuration any) - //Get config from application. + // Get config from application. Get(path string, defaultValue ...any) any - //GetString Get string type config from application. + // GetString get string type config from application. GetString(path string, defaultValue ...any) string - //GetInt Get int type config from application. + // GetInt get int type config from application. GetInt(path string, defaultValue ...any) int - //GetBool Get bool type config from application. + // GetBool get bool type config from application. GetBool(path string, defaultValue ...any) bool } diff --git a/contracts/console/artisan.go b/contracts/console/artisan.go index bce905e62..a1ee5357e 100644 --- a/contracts/console/artisan.go +++ b/contracts/console/artisan.go @@ -2,15 +2,15 @@ package console //go:generate mockery --name=Artisan type Artisan interface { - //Register commands. + // Register commands. Register(commands []Command) - //Call Run an Artisan console command by name. + // Call run an Artisan console command by name. Call(command string) - //CallAndExit Run an Artisan console command by name and exit. + // CallAndExit run an Artisan console command by name and exit. CallAndExit(command string) - //Run a command. args include: ["./main", "artisan", "command"] + // Run a command. args include: ["./main", "artisan", "command"] Run(args []string, exitIfArtisan bool) } diff --git a/contracts/console/command.go b/contracts/console/command.go index 1d983c6f1..cdc89f176 100644 --- a/contracts/console/command.go +++ b/contracts/console/command.go @@ -5,27 +5,38 @@ import ( ) type Command interface { - //Signature The name and signature of the console command. + // Signature set the unique signature for the command. Signature() string - //Description The console command description. + // Description the console command description. Description() string - //Extend The console command extend. + // Extend the console command extend. Extend() command.Extend - //Handle Execute the console command. + // Handle execute the console command. Handle(ctx Context) error } //go:generate mockery --name=Context type Context interface { + // Argument get the value of a command argument. Argument(index int) string + // Arguments get all the arguments passed to command. Arguments() []string + // Option gets the value of a command option. Option(key string) string + // OptionSlice looks up the value of a local StringSliceFlag, returns nil if not found OptionSlice(key string) []string + // OptionBool looks up the value of a local BoolFlag, returns false if not found OptionBool(key string) bool + // OptionFloat64 looks up the value of a local Float64Flag, returns zero if not found OptionFloat64(key string) float64 + // OptionFloat64Slice looks up the value of a local Float64SliceFlag, returns nil if not found OptionFloat64Slice(key string) []float64 + // OptionInt looks up the value of a local IntFlag, returns zero if not found OptionInt(key string) int + // OptionIntSlice looks up the value of a local IntSliceFlag, returns nil if not found OptionIntSlice(key string) []int + // OptionInt64 looks up the value of a local Int64Flag, returns zero if not found OptionInt64(key string) int64 + // OptionInt64Slice looks up the value of a local Int64SliceFlag, returns nil if not found OptionInt64Slice(key string) []int64 } diff --git a/contracts/console/command/command.go b/contracts/console/command/command.go index b689e26db..b911eb086 100644 --- a/contracts/console/command/command.go +++ b/contracts/console/command/command.go @@ -18,6 +18,7 @@ type Extend struct { } type Flag interface { + // Type gets a flag type. Type() string } diff --git a/contracts/database/factory/factory.go b/contracts/database/factory/factory.go index e3233e93f..21ee685fe 100644 --- a/contracts/database/factory/factory.go +++ b/contracts/database/factory/factory.go @@ -1,9 +1,11 @@ package factory type Factory interface { + // Definition defines the model's default state. Definition() map[string]any } type Model interface { + // Factory creates a new factory instance for the model. Factory() Factory } diff --git a/contracts/database/orm/events.go b/contracts/database/orm/events.go index 67f8c39ce..8c00b48a9 100644 --- a/contracts/database/orm/events.go +++ b/contracts/database/orm/events.go @@ -19,15 +19,23 @@ const EventForceDeleting EventType = "force_deleting" const EventForceDeleted EventType = "force_deleted" type Event interface { + // Context returns the event context. Context() context.Context + // GetAttribute returns the attribute value for the given key. GetAttribute(key string) any + // GetOriginal returns the original attribute value for the given key. GetOriginal(key string, def ...any) any + // IsDirty returns true if the given column is dirty. IsDirty(columns ...string) bool + // IsClean returns true if the given column is clean. IsClean(columns ...string) bool + // Query returns the query instance. Query() Query + // SetAttribute sets the attribute value for the given key. SetAttribute(key string, value any) } type DispatchesEvents interface { + // DispatchesEvents returns the event handlers. DispatchesEvents() map[EventType]func(Event) error } diff --git a/contracts/database/orm/factory.go b/contracts/database/orm/factory.go index aa28dfe67..7e679eae4 100644 --- a/contracts/database/orm/factory.go +++ b/contracts/database/orm/factory.go @@ -2,8 +2,12 @@ package orm //go:generate mockery --name=Factory type Factory interface { + // Count sets the number of models that should be generated. Count(count int) Factory + // Create creates a model and persists it to the database. Create(value any, attributes ...map[string]any) error + // CreateQuietly creates a model and persists it to the database without firing any model events. CreateQuietly(value any, attributes ...map[string]any) error + // Make creates a model and returns it, but does not persist it to the database. Make(value any, attributes ...map[string]any) error } diff --git a/contracts/database/orm/observer.go b/contracts/database/orm/observer.go index af876c00c..dc0317be3 100644 --- a/contracts/database/orm/observer.go +++ b/contracts/database/orm/observer.go @@ -1,15 +1,26 @@ package orm type Observer interface { + // Retrieved called when the model is retrieved from the database. Retrieved(Event) error + // Creating called when the model is being created. Creating(Event) error + // Created called when the model has been created. Created(Event) error + // Updating called when the model is being updated. Updating(Event) error + // Updated called when the model has been updated. Updated(Event) error + // Saving called when the model is being saved. Saving(Event) error + // Saved called when the model has been saved. Saved(Event) error + // Deleting called when the model is being deleted. Deleting(Event) error + // Deleted called when the model has been deleted. Deleted(Event) error + // ForceDeleting called when the model is being force deleted. ForceDeleting(Event) error + // ForceDeleted called when the model has been force deleted. ForceDeleted(Event) error } diff --git a/contracts/database/orm/orm.go b/contracts/database/orm/orm.go index b719fe698..8e18e7faf 100644 --- a/contracts/database/orm/orm.go +++ b/contracts/database/orm/orm.go @@ -7,89 +7,157 @@ import ( //go:generate mockery --name=Orm type Orm interface { + // Connection gets an Orm instance from the connection pool. Connection(name string) Orm + // DB gets the underlying database connection. DB() (*sql.DB, error) + // Query gets a new query builder instance. Query() Query + // Factory gets a new factory instance for the given model name. Factory() Factory + // Observe registers an observer with the Orm. Observe(model any, observer Observer) + // Transaction runs a callback wrapped in a database transaction. Transaction(txFunc func(tx Transaction) error) error + // WithContext sets the context to be used by the Orm. WithContext(ctx context.Context) Orm } //go:generate mockery --name=Transaction type Transaction interface { Query + // Commit commits the changes in a transaction. Commit() error + // Rollback rolls back the changes in a transaction. Rollback() error } //go:generate mockery --name=Query type Query interface { + // Association gets an association instance by name. Association(association string) Association + // Begin begins a new transaction Begin() (Transaction, error) + // Driver gets the driver for the query. Driver() Driver + // Count retrieve the "count" result of the query. Count(count *int64) error + // Create inserts new record into the database. Create(value any) error + // Cursor returns a cursor, use scan to iterate over the returned rows. Cursor() (chan Cursor, error) + // Delete deletes records matching given conditions, if the conditions are empty will delete all records. Delete(value any, conds ...any) (*Result, error) + // Distinct specifies distinct fields to query. Distinct(args ...any) Query + // Exec executes raw sql Exec(sql string, values ...any) (*Result, error) + // Find finds records that match given conditions. Find(dest any, conds ...any) error + // FindOrFail finds records that match given conditions or throws an error. FindOrFail(dest any, conds ...any) error + // First finds record that match given conditions. First(dest any) error + // FirstOrCreate finds the first record that matches the given attributes + // or create a new one with those attributes if none was found. FirstOrCreate(dest any, conds ...any) error + // FirstOr finds the first record that matches the given conditions or + // execute the callback and return its result if no record is found. FirstOr(dest any, callback func() error) error + // FirstOrFail finds the first record that matches the given conditions or throws an error. FirstOrFail(dest any) error + // FirstOrNew finds the first record that matches the given conditions or + // return a new instance of the model initialized with those attributes. FirstOrNew(dest any, attributes any, values ...any) error + // ForceDelete forces delete records matching given conditions. ForceDelete(value any, conds ...any) (*Result, error) + // Get retrieves all rows from the database. Get(dest any) error + // Group specifies the group method on the query. Group(name string) Query + // Having specifying HAVING conditions for the query. Having(query any, args ...any) Query + // Join specifying JOIN conditions for the query. Join(query string, args ...any) Query + // Limit the number of records returned. Limit(limit int) Query + // Load loads a relationship for the model. Load(dest any, relation string, args ...any) error + // LoadMissing loads a relationship for the model that is not already loaded. LoadMissing(dest any, relation string, args ...any) error + // LockForUpdate locks the selected rows in the table for updating. LockForUpdate() Query + // Model sets the model instance to be queried. Model(value any) Query + // Offset specifies the number of records to skip before starting to return the records. Offset(offset int) Query + // Omit specifies columns that should be omitted from the query. Omit(columns ...string) Query + // Order specifies the order in which the results should be returned. Order(value any) Query + // OrWhere add an "or where" clause to the query. OrWhere(query any, args ...any) Query + // Paginate the given query into a simple paginator. Paginate(page, limit int, dest any, total *int64) error + // Pluck retrieves a single column from the database. Pluck(column string, dest any) error + // Raw creates a raw query. Raw(sql string, values ...any) Query + // Save updates value in a database Save(value any) error + // SaveQuietly updates value in a database without firing events SaveQuietly(value any) error + // Scan scans the query result and populates the destination object. Scan(dest any) error + // Scopes applies one or more query scopes. Scopes(funcs ...func(Query) Query) Query + // Select specifies fields that should be retrieved from the database. Select(query any, args ...any) Query + // SharedLock locks the selected rows in the table. SharedLock() Query + // Sum calculates the sum of a column's values and populates the destination object. Sum(column string, dest any) error + // Table specifies the table for the query. Table(name string, args ...any) Query + // Update updates records with the given column and values Update(column any, value ...any) (*Result, error) + // UpdateOrCreate finds the first record that matches the given attributes + // or create a new one with those attributes if none was found. UpdateOrCreate(dest any, attributes any, values any) error + // Where add a "where" clause to the query. Where(query any, args ...any) Query + // WithoutEvents disables event firing for the query. WithoutEvents() Query + // WithTrashed allows soft deleted models to be included in the results. WithTrashed() Query + // With returns a new query instance with the given relationships eager loaded. With(query string, args ...any) Query } //go:generate mockery --name=Association type Association interface { + // Find finds records that match given conditions. Find(out any, conds ...any) error + // Append appending a model to the association. Append(values ...any) error + // Replace replaces the association with the given value. Replace(values ...any) error + // Delete deletes the given value from the association. Delete(values ...any) error + // Clear clears the association. Clear() error + // Count returns the number of records in the association. Count() int64 } type ConnectionModel interface { + // Connection gets the connection name for the model. Connection() string } //go:generate mockery --name=Cursor type Cursor interface { + // Scan scans the current row into the given destination. Scan(value any) error } diff --git a/contracts/database/seeder/seeder.go b/contracts/database/seeder/seeder.go index 013d7303e..8e6f6c1b3 100644 --- a/contracts/database/seeder/seeder.go +++ b/contracts/database/seeder/seeder.go @@ -4,24 +4,19 @@ package seeder type Facade interface { // Register registers seeders. Register(seeders []Seeder) - // GetSeeder gets a seeder instance from the seeders. GetSeeder(name string) Seeder - - // All seeders + // GetSeeders gets all the seeders GetSeeders() []Seeder - // Call executes the specified seeder(s). Call(seeders []Seeder) error - // CallOnce executes the specified seeder(s) only once. CallOnce(seeders []Seeder) error } type Seeder interface { - // Signature The name and signature of the seeder. + // Signature the unique signature of the seeder. Signature() string - // Run executes the seeder logic. Run() error } diff --git a/contracts/event/events.go b/contracts/event/events.go index 68f58ec95..4c47913cb 100644 --- a/contracts/event/events.go +++ b/contracts/event/events.go @@ -2,23 +2,31 @@ package event //go:generate mockery --name=Instance type Instance interface { + // Register event listeners to the application. Register(map[Event][]Listener) + // Job create a new event task. Job(event Event, args []Arg) Task + // GetEvents gets all registered events. GetEvents() map[Event][]Listener } type Event interface { + // Handle the event. Handle(args []Arg) ([]Arg, error) } type Listener interface { + // Signature returns the unique identifier for the listener. Signature() string + // Queue configure the event queue options. Queue(args ...any) Queue + // Handle the event. Handle(args ...any) error } //go:generate mockery --name=Task type Task interface { + // Dispatch an event and call the listeners. Dispatch() error } diff --git a/contracts/filesystem/mocks/Driver.go b/contracts/filesystem/mocks/Driver.go index 8d3c7712f..b3824b6db 100644 --- a/contracts/filesystem/mocks/Driver.go +++ b/contracts/filesystem/mocks/Driver.go @@ -1,4 +1,4 @@ -// Code generated by mockery v2.20.0. DO NOT EDIT. +// Code generated by mockery v2.33.2. DO NOT EDIT. package mocks @@ -206,6 +206,32 @@ func (_m *Driver) Get(file string) (string, error) { return r0, r1 } +// GetBytes provides a mock function with given fields: file +func (_m *Driver) GetBytes(file string) ([]byte, error) { + ret := _m.Called(file) + + var r0 []byte + var r1 error + if rf, ok := ret.Get(0).(func(string) ([]byte, error)); ok { + return rf(file) + } + if rf, ok := ret.Get(0).(func(string) []byte); ok { + r0 = rf(file) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]byte) + } + } + + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(file) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // LastModified provides a mock function with given fields: file func (_m *Driver) LastModified(file string) (time.Time, error) { ret := _m.Called(file) @@ -450,13 +476,12 @@ func (_m *Driver) WithContext(ctx context.Context) filesystem.Driver { return r0 } -type mockConstructorTestingTNewDriver interface { +// NewDriver creates a new instance of Driver. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewDriver(t interface { mock.TestingT Cleanup(func()) -} - -// NewDriver creates a new instance of Driver. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -func NewDriver(t mockConstructorTestingTNewDriver) *Driver { +}) *Driver { mock := &Driver{} mock.Mock.Test(t) diff --git a/contracts/filesystem/mocks/File.go b/contracts/filesystem/mocks/File.go index 34b75f304..82706f6d4 100644 --- a/contracts/filesystem/mocks/File.go +++ b/contracts/filesystem/mocks/File.go @@ -1,10 +1,12 @@ -// Code generated by mockery v2.20.0. DO NOT EDIT. +// Code generated by mockery v2.33.2. DO NOT EDIT. package mocks import ( filesystem "github.com/goravel/framework/contracts/filesystem" mock "github.com/stretchr/testify/mock" + + time "time" ) // File is an autogenerated mock type for the File type @@ -114,6 +116,78 @@ func (_m *File) HashName(path ...string) string { return r0 } +// LastModified provides a mock function with given fields: +func (_m *File) LastModified() (time.Time, error) { + ret := _m.Called() + + var r0 time.Time + var r1 error + if rf, ok := ret.Get(0).(func() (time.Time, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() time.Time); ok { + r0 = rf() + } else { + r0 = ret.Get(0).(time.Time) + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// MimeType provides a mock function with given fields: +func (_m *File) MimeType() (string, error) { + ret := _m.Called() + + var r0 string + var r1 error + if rf, ok := ret.Get(0).(func() (string, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() string); ok { + r0 = rf() + } else { + r0 = ret.Get(0).(string) + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// Size provides a mock function with given fields: +func (_m *File) Size() (int64, error) { + ret := _m.Called() + + var r0 int64 + var r1 error + if rf, ok := ret.Get(0).(func() (int64, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() int64); ok { + r0 = rf() + } else { + r0 = ret.Get(0).(int64) + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // Store provides a mock function with given fields: path func (_m *File) Store(path string) (string, error) { ret := _m.Called(path) @@ -162,13 +236,12 @@ func (_m *File) StoreAs(path string, name string) (string, error) { return r0, r1 } -type mockConstructorTestingTNewFile interface { +// NewFile creates a new instance of File. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewFile(t interface { mock.TestingT Cleanup(func()) -} - -// NewFile creates a new instance of File. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -func NewFile(t mockConstructorTestingTNewFile) *File { +}) *File { mock := &File{} mock.Mock.Test(t) diff --git a/contracts/filesystem/mocks/Storage.go b/contracts/filesystem/mocks/Storage.go index e42f9000a..70e4b9c2a 100644 --- a/contracts/filesystem/mocks/Storage.go +++ b/contracts/filesystem/mocks/Storage.go @@ -1,4 +1,4 @@ -// Code generated by mockery v2.20.0. DO NOT EDIT. +// Code generated by mockery v2.33.2. DO NOT EDIT. package mocks @@ -222,6 +222,32 @@ func (_m *Storage) Get(file string) (string, error) { return r0, r1 } +// GetBytes provides a mock function with given fields: file +func (_m *Storage) GetBytes(file string) ([]byte, error) { + ret := _m.Called(file) + + var r0 []byte + var r1 error + if rf, ok := ret.Get(0).(func(string) ([]byte, error)); ok { + return rf(file) + } + if rf, ok := ret.Get(0).(func(string) []byte); ok { + r0 = rf(file) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]byte) + } + } + + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(file) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // LastModified provides a mock function with given fields: file func (_m *Storage) LastModified(file string) (time.Time, error) { ret := _m.Called(file) @@ -466,13 +492,12 @@ func (_m *Storage) WithContext(ctx context.Context) filesystem.Driver { return r0 } -type mockConstructorTestingTNewStorage interface { +// NewStorage creates a new instance of Storage. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewStorage(t interface { mock.TestingT Cleanup(func()) -} - -// NewStorage creates a new instance of Storage. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -func NewStorage(t mockConstructorTestingTNewStorage) *Storage { +}) *Storage { mock := &Storage{} mock.Mock.Test(t) diff --git a/contracts/filesystem/storage.go b/contracts/filesystem/storage.go index 33a10bf6d..c56dd2486 100644 --- a/contracts/filesystem/storage.go +++ b/contracts/filesystem/storage.go @@ -8,46 +8,82 @@ import ( //go:generate mockery --name=Storage type Storage interface { Driver + // Disk gets the instance of the given disk. Disk(disk string) Driver } //go:generate mockery --name=Driver type Driver interface { + // AllDirectories gets all the directories within a given directory(recursive). AllDirectories(path string) ([]string, error) + // AllFiles gets all the files from the given directory(recursive). AllFiles(path string) ([]string, error) + // Copy the given file to a new location. Copy(oldFile, newFile string) error + // Delete deletes the given file(s). Delete(file ...string) error + // DeleteDirectory deletes the given directory(recursive). DeleteDirectory(directory string) error + // Directories get all the directories within a given directory. Directories(path string) ([]string, error) + // Exists determines if a file exists. Exists(file string) bool + // Files gets all the files from the given directory. Files(path string) ([]string, error) + // Get gets the contents of a file. Get(file string) (string, error) + // GetBytes gets the contents of a file as a byte array. + GetBytes(file string) ([]byte, error) + // LastModified gets the file's last modified time. LastModified(file string) (time.Time, error) + // MakeDirectory creates a directory. MakeDirectory(directory string) error + // MimeType gets the file's mime type. MimeType(file string) (string, error) + // Missing determines if a file is missing. Missing(file string) bool + // Move a file to a new location. Move(oldFile, newFile string) error + // Path gets the full path for the file. Path(file string) string + // Put writes the contents of a file. Put(file, content string) error + // PutFile upload the given file. PutFile(path string, source File) (string, error) + // PutFileAs upload the given file with a new name. PutFileAs(path string, source File, name string) (string, error) + // Size gets the file size of a given file. Size(file string) (int64, error) + // TemporaryUrl get a temporary URL for the file. TemporaryUrl(file string, time time.Time) (string, error) + // WithContext sets the context to be used by the driver. WithContext(ctx context.Context) Driver + // Url get the URL for the file at the given path. Url(file string) string } //go:generate mockery --name=File type File interface { + // Disk gets the instance of the given disk. Disk(disk string) File + // Extension gets the file extension. Extension() (string, error) + // File gets the file path. File() string + // GetClientOriginalName gets the client original name. GetClientOriginalName() string + // GetClientOriginalExtension gets the client original extension. GetClientOriginalExtension() string + // HashName gets the file's hash name. HashName(path ...string) string + // LastModified gets the file's last modified time. LastModified() (time.Time, error) + // MimeType gets the file's mime type. MimeType() (string, error) + // Size gets the file size. Size() (int64, error) + // Store the file at the given path. Store(path string) (string, error) + // StoreAs store the file at the given path with a new name. StoreAs(path string, name string) (string, error) } diff --git a/contracts/foundation/application.go b/contracts/foundation/application.go index 9f0dd0029..e5e0ddf25 100644 --- a/contracts/foundation/application.go +++ b/contracts/foundation/application.go @@ -7,13 +7,22 @@ import ( //go:generate mockery --name=Application type Application interface { Container + // Boot register and bootstrap configured service providers. Boot() + // Commands register the given commands with the console application. Commands([]console.Command) + // Path gets the path respective to "app" directory. Path(path string) string + // BasePath get the base path of the Goravel installation. BasePath(path string) string + // ConfigPath get the path to the configuration files. ConfigPath(path string) string + // DatabasePath get the path to the database directory. DatabasePath(path string) string + // StoragePath get the path to the storage directory. StoragePath(path string) string + // PublicPath get the path to the public directory. PublicPath(path string) string + // Publishes register the given paths to be published by the "vendor:publish" command. Publishes(packageName string, paths map[string]string, groups ...string) } diff --git a/contracts/foundation/container.go b/contracts/foundation/container.go index 14b222220..e7be9cc4f 100644 --- a/contracts/foundation/container.go +++ b/contracts/foundation/container.go @@ -24,31 +24,58 @@ import ( ) type Container interface { + // Bind registers a binding with the container. Bind(key any, callback func(app Application) (any, error)) + // BindWith registers a binding with the container. BindWith(key any, callback func(app Application, parameters map[string]any) (any, error)) + // Instance registers an existing instance as shared in the container. Instance(key, instance any) + // Make resolves the given type from the container. Make(key any) (any, error) + // MakeArtisan resolves the artisan console instance. MakeArtisan() console.Artisan + // MakeAuth resolves the auth instance. MakeAuth() auth.Auth + // MakeCache resolves the cache instance. MakeCache() cache.Cache + // MakeConfig resolves the config instance. MakeConfig() config.Config + // MakeCrypt resolves the crypt instance. MakeCrypt() crypt.Crypt + // MakeEvent resolves the event instance. MakeEvent() event.Instance + // MakeGate resolves the gate instance. MakeGate() access.Gate + // MakeGrpc resolves the grpc instance. MakeGrpc() grpc.Grpc + // MakeHash resolves the hash instance. MakeHash() hash.Hash + // MakeLog resolves the log instance. MakeLog() log.Log + // MakeMail resolves the mail instance. MakeMail() mail.Mail + // MakeOrm resolves the orm instance. MakeOrm() orm.Orm + // MakeQueue resolves the queue instance. MakeQueue() queue.Queue + // MakeRateLimiter resolves the rate limiter instance. MakeRateLimiter() http.RateLimiter + // MakeRoute resolves the route instance. MakeRoute() route.Route + // MakeSchedule resolves the schedule instance. MakeSchedule() schedule.Schedule + // MakeStorage resolves the storage instance. MakeStorage() filesystem.Storage + // MakeTesting resolves the testing instance. MakeTesting() testing.Testing + // MakeValidation resolves the validation instance. MakeValidation() validation.Validation + // MakeView resolves the view instance. MakeView() http.View + // MakeSeeder resolves the seeder instance. MakeSeeder() seeder.Facade + // MakeWith resolves the given type with the given parameters from the container. MakeWith(key any, parameters map[string]any) (any, error) + // Singleton registers a shared binding in the container. Singleton(key any, callback func(app Application) (any, error)) } diff --git a/contracts/grpc/grpc.go b/contracts/grpc/grpc.go index bdfdc411c..c1f4ad5d2 100644 --- a/contracts/grpc/grpc.go +++ b/contracts/grpc/grpc.go @@ -8,9 +8,14 @@ import ( //go:generate mockery --name=Grpc type Grpc interface { + // Run starts the gRPC server. Run(host ...string) error + // Server gets the gRPC server instance. Server() *grpc.Server + // Client gets the gRPC client instance. Client(ctx context.Context, name string) (*grpc.ClientConn, error) + // UnaryServerInterceptors sets the gRPC server interceptors. UnaryServerInterceptors([]grpc.UnaryServerInterceptor) + // UnaryClientInterceptorGroups sets the gRPC client interceptor groups. UnaryClientInterceptorGroups(map[string][]grpc.UnaryClientInterceptor) } diff --git a/contracts/http/context.go b/contracts/http/context.go index b72c8e603..2ec0ee27c 100644 --- a/contracts/http/context.go +++ b/contracts/http/context.go @@ -7,18 +7,27 @@ import ( type Middleware func(Context) type HandlerFunc func(Context) Response type ResourceController interface { + // Index method for controller Index(Context) Response + // Show method for controller Show(Context) Response + // Store method for controller Store(Context) Response + // Update method for controller Update(Context) Response + // Destroy method for controller Destroy(Context) Response } //go:generate mockery --name=Context type Context interface { context.Context + // Context returns the Context Context() context.Context + // WithValue add value associated with key in context WithValue(key string, value any) + // Request returns the ContextRequest Request() ContextRequest + // Response returns the ContextResponse Response() ContextResponse } diff --git a/contracts/http/rate_limiter.go b/contracts/http/rate_limiter.go index f027a43e4..eebffa1b2 100644 --- a/contracts/http/rate_limiter.go +++ b/contracts/http/rate_limiter.go @@ -2,12 +2,17 @@ package http //go:generate mockery --name=RateLimiter type RateLimiter interface { + // For register a new rate limiter. For(name string, callback func(ctx Context) Limit) + // ForWithLimits register a new rate limiter with limits. ForWithLimits(name string, callback func(ctx Context) []Limit) + // Limiter get a rate limiter instance by name. Limiter(name string) func(ctx Context) []Limit } type Limit interface { + // By set the signature key name for the rate limiter. By(key string) Limit + // Response set the response callback that should be used. Response(func(ctx Context)) Limit } diff --git a/contracts/http/request.go b/contracts/http/request.go index 2f2784b41..a45ebb328 100644 --- a/contracts/http/request.go +++ b/contracts/http/request.go @@ -9,56 +9,84 @@ import ( //go:generate mockery --name=ContextRequest type ContextRequest interface { + // Header retrieves the value of the specified HTTP header by its key. + // If the header is not found, it returns the optional default value (if provided). Header(key string, defaultValue ...string) string + // Headers return all the HTTP headers of the request. Headers() http.Header + // Method retrieves the HTTP request method (e.g., GET, POST, PUT). Method() string + // Path retrieves the current path information for the request. Path() string + // Url retrieves the URL (excluding the query string) for the request. Url() string + // FullUrl retrieves the full URL, including the query string, for the request. FullUrl() string + // Ip retrieves the client's IP address. Ip() string + // Host retrieves the host name. Host() string - - // All Retrieve json, form and query + // All retrieves data from JSON, form, and query parameters. All() map[string]any - // Bind Retrieve json and bind to obj + // Bind retrieve json and bind to obj Bind(obj any) error - // Route Retrieve an input item from the request: /users/{id} + // Route retrieves a route parameter from the request path (e.g., /users/{id}). Route(key string) string + // RouteInt retrieves a route parameter from the request path and attempts to parse it as an integer. RouteInt(key string) int + // RouteInt64 retrieves a route parameter from the request path and attempts to parse it as a 64-bit integer. RouteInt64(key string) int64 - // Query Retrieve a query string item form the request: /users?id=1 + // Query retrieves a query string parameter from the request (e.g., /users?id=1). Query(key string, defaultValue ...string) string + // QueryInt retrieves a query string parameter from the request and attempts to parse it as an integer. QueryInt(key string, defaultValue ...int) int + // QueryInt64 retrieves a query string parameter from the request and attempts to parse it as a 64-bit integer. QueryInt64(key string, defaultValue ...int64) int64 + // QueryBool retrieves a query string parameter from the request and attempts to parse it as a boolean. QueryBool(key string, defaultValue ...bool) bool + // QueryArray retrieves a query string parameter from the request and returns it as a slice of strings. QueryArray(key string) []string + // QueryMap retrieves a query string parameter from the request and returns it as a map of key-value pairs. QueryMap(key string) map[string]string + // Queries returns all the query string parameters from the request as a map of key-value pairs. Queries() map[string]string - // Input Retrieve data by order: json, form, query, route + // Input retrieves data from the request in the following order: JSON, form, query, and route parameters. Input(key string, defaultValue ...string) string InputArray(key string, defaultValue ...[]string) []string InputMap(key string, defaultValue ...map[string]string) map[string]string InputInt(key string, defaultValue ...int) int InputInt64(key string, defaultValue ...int64) int64 InputBool(key string, defaultValue ...bool) bool - + // File retrieves a file by its key from the request. File(name string) (filesystem.File, error) + // AbortWithStatus aborts the request with the specified HTTP status code. AbortWithStatus(code int) + // AbortWithStatusJson aborts the request with the specified HTTP status code + // and returns a JSON response object. AbortWithStatusJson(code int, jsonObj any) - + // Next skips the current request handler, allowing the next middleware or handler to be executed. Next() + // Origin retrieves the underlying *http.Request object for advanced request handling. Origin() *http.Request + // Validate performs request data validation using specified rules and options. Validate(rules map[string]string, options ...validation.Option) (validation.Validator, error) + // ValidateRequest validates the request data against a pre-defined FormRequest structure + // and returns validation errors, if any. ValidateRequest(request FormRequest) (validation.Errors, error) } type FormRequest interface { + // Authorize determine if the user is authorized to make this request. Authorize(ctx Context) error + // Rules get the validation rules that apply to the request. Rules(ctx Context) map[string]string + // Messages get the validation messages that apply to the request. Messages(ctx Context) map[string]string + // Attributes get custom attributes for validator errors. Attributes(ctx Context) map[string]string + // PrepareForValidation prepare the data for validation. PrepareForValidation(ctx Context, data validation.Data) error } diff --git a/contracts/http/response.go b/contracts/http/response.go index d7b95f6ce..8fcfe4a50 100644 --- a/contracts/http/response.go +++ b/contracts/http/response.go @@ -14,44 +14,70 @@ type Response interface { //go:generate mockery --name=ContextResponse type ContextResponse interface { + // Data write the given data to the response. Data(code int, contentType string, data []byte) Response + // Download initiates a file download by specifying the file path and the desired filename Download(filepath, filename string) Response + // File serves a file located at the specified file path as the response. File(filepath string) Response + // Header sets an HTTP header field with the given key and value. Header(key, value string) ContextResponse + // Json sends a JSON response with the specified status code and data object. Json(code int, obj any) Response + // Origin returns the ResponseOrigin Origin() ResponseOrigin + // Redirect performs an HTTP redirect to the specified location with the given status code. Redirect(code int, location string) Response + // String writes a string response with the specified status code and format. + // The 'values' parameter can be used to replace placeholders in the format string. String(code int, format string, values ...any) Response + // Success returns ResponseSuccess Success() ResponseSuccess + // Status sets the HTTP response status code and returns the ResponseStatus. Status(code int) ResponseStatus + // View returns ResponseView View() ResponseView + // Writer returns the underlying http.ResponseWriter associated with the response. Writer() http.ResponseWriter + // Flush flushes any buffered data to the client. Flush() } //go:generate mockery --name=ResponseStatus type ResponseStatus interface { + // Data write the given data to the Response. Data(contentType string, data []byte) Response + // Json sends a JSON Response with the specified data object. Json(obj any) Response + // String writes a string Response with the specified format and values. String(format string, values ...any) Response } //go:generate mockery --name=ResponseSuccess type ResponseSuccess interface { + // Data write the given data to the Response. Data(contentType string, data []byte) Response + // Json sends a JSON Response with the specified data object. Json(obj any) Response + // String writes a string Response with the specified format and values. String(format string, values ...any) Response } //go:generate mockery --name=ResponseOrigin type ResponseOrigin interface { + // Body returns the response's body content as a *bytes.Buffer. Body() *bytes.Buffer + // Header returns the response's HTTP header. Header() http.Header + // Size returns the size, in bytes, of the response's body content. Size() int + // Status returns the HTTP status code of the response. Status() int } type ResponseView interface { + // Make generates a Response for the specified view with optional data. Make(view string, data ...any) Response + // First generates a response for the first available view from the provided list. First(views []string, data ...any) Response } diff --git a/contracts/http/view.go b/contracts/http/view.go index a8416f67d..4c82769d6 100644 --- a/contracts/http/view.go +++ b/contracts/http/view.go @@ -2,8 +2,14 @@ package http //go:generate mockery --name=View type View interface { + // Exists checks if a view with the specified name exists. Exists(view string) bool + // Share associates a key-value pair, where the key is a string and the value is of any type, + // with the current view context. This shared data can be accessed by other parts of the application. Share(key string, value any) + // Shared retrieves the value associated with the given key from the current view context's shared data. + // If the key does not exist, it returns the optional default value (if provided). Shared(key string, def ...any) any + // GetShared returns a map containing all the shared data associated with the current view context. GetShared() map[string]any } diff --git a/contracts/log/log.go b/contracts/log/log.go index 80380ad25..7bffec635 100644 --- a/contracts/log/log.go +++ b/contracts/log/log.go @@ -25,23 +25,36 @@ const ( //go:generate mockery --name=Log type Log interface { + // WithContext adds a context to the logger. WithContext(ctx context.Context) Writer Writer } //go:generate mockery --name=Writer type Writer interface { + // Debug logs a message at DebugLevel. Debug(args ...any) + // Debugf is equivalent to Debug, but with support for fmt.Printf-style arguments. Debugf(format string, args ...any) + // Info logs a message at InfoLevel. Info(args ...any) + // Infof is equivalent to Info, but with support for fmt.Printf-style arguments. Infof(format string, args ...any) + // Warning logs a message at WarningLevel. Warning(args ...any) + // Warningf is equivalent to Warning, but with support for fmt.Printf-style arguments. Warningf(format string, args ...any) + // Error logs a message at ErrorLevel. Error(args ...any) + // Errorf is equivalent to Error, but with support for fmt.Printf-style arguments. Errorf(format string, args ...any) + // Fatal logs a message at FatalLevel. Fatal(args ...any) + // Fatalf is equivalent to Fatal, but with support for fmt.Printf-style arguments. Fatalf(format string, args ...any) + // Panic logs a message at PanicLevel. Panic(args ...any) + // Panicf is equivalent to Panic, but with support for fmt.Printf-style arguments. Panicf(format string, args ...any) // Code set a code or slug that describes the error. // Error messages are intended to be read by humans, but such code is expected to @@ -68,7 +81,7 @@ type Writer interface { //go:generate mockery --name=Logger type Logger interface { - // Handle pass channel config path here + // Handle pass a channel config path here Handle(channel string) (Hook, error) } @@ -76,14 +89,18 @@ type Logger interface { type Hook interface { // Levels monitoring level Levels() []Level - // Fire execute logic when trigger + // Fire executes logic when trigger Fire(Entry) error } //go:generate mockery --name=Entry type Entry interface { + // Context returns the context of the entry. Context() context.Context + // Level returns the level of the entry. Level() Level + // Time returns the timestamp of the entry. Time() time.Time + // Message returns the message of the entry. Message() string } diff --git a/contracts/mail/mail.go b/contracts/mail/mail.go index 6a56ab0d7..aa1bcec2c 100644 --- a/contracts/mail/mail.go +++ b/contracts/mail/mail.go @@ -2,13 +2,21 @@ package mail //go:generate mockery --name=Mail type Mail interface { + // Content set the content of Mail. Content(content Content) Mail + // From set the sender of Mail. From(address From) Mail + // To set the recipients of Mail. To(addresses []string) Mail + // Cc adds a "carbon copy" address to the Mail. Cc(addresses []string) Mail + // Bcc adds a "blind carbon copy" address to the Mail. Bcc(addresses []string) Mail + // Attach attaches files to the Mail. Attach(files []string) Mail + // Send the Mail Send() error + // Queue a given Mail Queue(queue *Queue) error } diff --git a/contracts/queue/job.go b/contracts/queue/job.go index 12a5fe889..9cfcf9f3c 100644 --- a/contracts/queue/job.go +++ b/contracts/queue/job.go @@ -1,7 +1,9 @@ package queue type Job interface { + // Signature set the unique signature of the job. Signature() string + // Handle executes the job. Handle(args ...any) error } diff --git a/contracts/queue/queue.go b/contracts/queue/queue.go index d72eee874..c703e6480 100644 --- a/contracts/queue/queue.go +++ b/contracts/queue/queue.go @@ -3,13 +3,13 @@ package queue //go:generate mockery --name=Queue type Queue interface { Worker(args *Args) Worker - // Register Register jobs + // Register register jobs Register(jobs []Job) - // GetJobs Get all jobs + // GetJobs get all jobs GetJobs() []Job - // Job Add a job to queue + // Job add a job to queue Job(job Job, args []Arg) Task - // Chain Creates a chain of jobs to be processed one by one, passing + // Chain creates a chain of jobs to be processed one by one, passing Chain(jobs []Jobs) Task } diff --git a/contracts/queue/task.go b/contracts/queue/task.go index 4a833f0f1..9e80da11c 100644 --- a/contracts/queue/task.go +++ b/contracts/queue/task.go @@ -6,9 +6,14 @@ import ( //go:generate mockery --name=Task type Task interface { + // Dispatch dispatches the task. Dispatch() error + // DispatchSync dispatches the task synchronously. DispatchSync() error + // Delay dispatches the task after the given delay. Delay(time time.Time) Task + // OnConnection sets the connection of the task. OnConnection(connection string) Task + // OnQueue sets the queue of the task. OnQueue(queue string) Task } diff --git a/contracts/route/route.go b/contracts/route/route.go index 2d7814a56..344df27c7 100644 --- a/contracts/route/route.go +++ b/contracts/route/route.go @@ -11,30 +11,50 @@ type GroupFunc func(router Router) //go:generate mockery --name=Route type Route interface { Router + // Fallback registers a handler to be executed when no other route was matched. Fallback(handler contractshttp.HandlerFunc) + // GlobalMiddleware registers global middleware to be applied to all routes of the router. GlobalMiddleware(middlewares ...contractshttp.Middleware) + // Run starts the HTTP server and listens for incoming connections on the specified host. Run(host ...string) error + // RunTLS starts the HTTPS server with the provided TLS configuration and listens on the specified host. RunTLS(host ...string) error + // RunTLSWithCert starts the HTTPS server with the provided certificate and key files and listens on the specified host and port. RunTLSWithCert(host, certFile, keyFile string) error + // ServeHTTP serves HTTP requests. ServeHTTP(writer http.ResponseWriter, request *http.Request) } //go:generate mockery --name=Router type Router interface { + // Group creates a new router group with the specified handler. Group(handler GroupFunc) + // Prefix adds a common prefix to the routes registered with the router. Prefix(addr string) Router + // Middleware sets the middleware for the router. Middleware(middlewares ...contractshttp.Middleware) Router + // Any registers a new route responding to all verbs. Any(relativePath string, handler contractshttp.HandlerFunc) + // Get registers a new GET route with the router. Get(relativePath string, handler contractshttp.HandlerFunc) + // Post registers a new POST route with the router. Post(relativePath string, handler contractshttp.HandlerFunc) + // Delete registers a new DELETE route with the router. Delete(relativePath string, handler contractshttp.HandlerFunc) + // Patch registers a new PATCH route with the router. Patch(relativePath string, handler contractshttp.HandlerFunc) + // Put registers a new PUT route with the router. Put(relativePath string, handler contractshttp.HandlerFunc) + // Options registers a new OPTIONS route with the router. Options(relativePath string, handler contractshttp.HandlerFunc) + // Resource registers RESTful routes for a resource controller. Resource(relativePath string, controller contractshttp.ResourceController) + // Static registers a new route with path prefix to serve static files from the provided root directory. Static(relativePath, root string) + // StaticFile registers a new route with a specific path to serve a static file from the filesystem. StaticFile(relativePath, filepath string) + // StaticFS registers a new route with a path prefix to serve static files from the provided file system. StaticFS(relativePath string, fs http.FileSystem) } diff --git a/contracts/schedule/event.go b/contracts/schedule/event.go index 0014a53ce..024eabc1d 100644 --- a/contracts/schedule/event.go +++ b/contracts/schedule/event.go @@ -2,33 +2,62 @@ package schedule //go:generate mockery --name=Event type Event interface { + // At schedule the event to run at the specified time. At(time string) Event + // Cron schedule the event using the given Cron expression. Cron(expression string) Event + // Daily schedule the event to run daily. Daily() Event + // DailyAt schedule the event to run daily at a given time (10:00, 19:30, etc). DailyAt(time string) Event + // DelayIfStillRunning if the event is still running, the event will be delayed. DelayIfStillRunning() Event + // EveryMinute schedule the event to run every minute. EveryMinute() Event + // EveryTwoMinutes schedule the event to run every two minutes. EveryTwoMinutes() Event + // EveryThreeMinutes schedule the event to run every three minutes. EveryThreeMinutes() Event + // EveryFourMinutes schedule the event to run every four minutes. EveryFourMinutes() Event + // EveryFiveMinutes schedule the event to run every five minutes. EveryFiveMinutes() Event + // EveryTenMinutes schedule the event to run every ten minutes. EveryTenMinutes() Event + // EveryFifteenMinutes schedule the event to run every fifteen minutes. EveryFifteenMinutes() Event + // EveryThirtyMinutes schedule the event to run every thirty minutes. EveryThirtyMinutes() Event + // EveryTwoHours schedule the event to run every two hours. EveryTwoHours() Event + // EveryThreeHours schedule the event to run every three hours. EveryThreeHours() Event + // EveryFourHours schedule the event to run every four hours. EveryFourHours() Event + // EverySixHours schedule the event to run every six hours. EverySixHours() Event + // GetCron get cron expression. GetCron() string + // GetCommand get the command. GetCommand() string + // GetCallback get callback. GetCallback() func() + // GetName get name. GetName() string + // GetSkipIfStillRunning get skipIfStillRunning bool. GetSkipIfStillRunning() bool + // GetDelayIfStillRunning get delayIfStillRunning bool. GetDelayIfStillRunning() bool + // Hourly schedule the event to run hourly. Hourly() Event + // HourlyAt schedule the event to run hourly at a given offset in the hour. HourlyAt(offset []string) Event + // IsOnOneServer get isOnOneServer bool. IsOnOneServer() bool + // Name set the event name. Name(name string) Event + // OnOneServer only allow the event to run on one server for each cron expression. OnOneServer() Event + // SkipIfStillRunning if the event is still running, the event will be skipped. SkipIfStillRunning() Event } diff --git a/contracts/schedule/schedule.go b/contracts/schedule/schedule.go index 9a3fb5d4a..ef031bf32 100644 --- a/contracts/schedule/schedule.go +++ b/contracts/schedule/schedule.go @@ -2,15 +2,12 @@ package schedule //go:generate mockery --name=Schedule type Schedule interface { - //Call Add a new callback event to the schedule. + // Call add a new callback event to the schedule. Call(callback func()) Event - - //Command Add a new Artisan command event to the schedule. + // Command adds a new Artisan command event to the schedule. Command(command string) Event - - //Register schedules. + // Register schedules. Register(events []Event) - - //Run schedules. + // Run schedules. Run() } diff --git a/contracts/testing/testing.go b/contracts/testing/testing.go index 27b1e7b8d..ec9e09e59 100644 --- a/contracts/testing/testing.go +++ b/contracts/testing/testing.go @@ -8,25 +8,36 @@ import ( ) type Testing interface { + // Docker get the Docker instance. Docker() Docker } type Docker interface { + // Database get a database connection instance. Database(connection ...string) (Database, error) } type Database interface { + // Build the database. Build() error + // Config gets the database configuration. Config() Config + // Clear clears the database. Clear() error + // Image gets the database image. Image(Image) + // Seed runs the database seeds. Seed(seeds ...seeder.Seeder) } type DatabaseDriver interface { + // Config gets the database configuration. Config(resource *dockertest.Resource) Config + // Clear clears the database. Clear(pool *dockertest.Pool, resource *dockertest.Resource) error + // Name gets the database driver name. Name() orm.Driver + // Image gets the database image. Image() *dockertest.RunOptions } diff --git a/contracts/validation/validation.go b/contracts/validation/validation.go index acba21edd..94065e5a0 100644 --- a/contracts/validation/validation.go +++ b/contracts/validation/validation.go @@ -4,33 +4,48 @@ type Option func(map[string]any) //go:generate mockery --name=Validation type Validation interface { + // Make create a new validator instance. Make(data any, rules map[string]string, options ...Option) (Validator, error) + // AddRules add the custom rules. AddRules([]Rule) error + // Rules get the custom rules. Rules() []Rule } //go:generate mockery --name=Validator type Validator interface { + // Bind the data to the validation. Bind(ptr any) error + // Errors get the validation errors. Errors() Errors + // Fails determine if the validation fails. Fails() bool } //go:generate mockery --name=Errors type Errors interface { + // One gets the first error message for a given field. One(key ...string) string + // Get gets all the error messages for a given field. Get(key string) map[string]string + // All gets all the error messages. All() map[string]map[string]string + // Has checks if there are any error messages for a given field. Has(key string) bool } type Data interface { + // Get the value from the given key. Get(key string) (val any, exist bool) + // Set the value for a given key. Set(key string, val any) error } type Rule interface { + // Signature set the unique signature of the rule. Signature() string + // Passes determine if the validation rule passes. Passes(data Data, val any, options ...any) bool + // Message gets the validation error message. Message() string } diff --git a/crypt/aes.go b/crypt/aes.go index 334127528..1d42d967b 100644 --- a/crypt/aes.go +++ b/crypt/aes.go @@ -8,11 +8,11 @@ import ( "errors" "io" - "github.com/bytedance/sonic" "github.com/gookit/color" "github.com/goravel/framework/contracts/config" "github.com/goravel/framework/support" + "github.com/goravel/framework/support/json" ) type AES struct { @@ -60,10 +60,12 @@ func (b *AES) EncryptString(value string) (string, error) { ciphertext := aesgcm.Seal(nil, iv, plaintext, nil) - jsonEncoded, err := sonic.Marshal(map[string][]byte{ + var jsonEncoded []byte + jsonEncoded, err = json.Marshal(map[string][]byte{ "iv": iv, "value": ciphertext, }) + if err != nil { return "", err } @@ -79,7 +81,7 @@ func (b *AES) DecryptString(payload string) (string, error) { } decodeJson := make(map[string][]byte) - err = sonic.Unmarshal(decodePayload, &decodeJson) + err = json.Unmarshal(decodePayload, &decodeJson) if err != nil { return "", err } diff --git a/database/db/dsn.go b/database/db/dsn.go index aa16aa054..c88c73951 100644 --- a/database/db/dsn.go +++ b/database/db/dsn.go @@ -48,8 +48,8 @@ func (d *DsnImpl) Postgresql(config databasecontract.Config) string { sslmode := d.config.GetString("database.connections." + d.connection + ".sslmode") timezone := d.config.GetString("database.connections." + d.connection + ".timezone") - return fmt.Sprintf("host=%s user=%s password=%s dbname=%s port=%d sslmode=%s TimeZone=%s", - host, config.Username, config.Password, config.Database, config.Port, sslmode, timezone) + return fmt.Sprintf("postgres://%s:%s@%s:%d/%s?sslmode=%s&timezone=%s", + config.Username, config.Password, host, config.Port, config.Database, sslmode, timezone) } func (d *DsnImpl) Sqlite(config databasecontract.Config) string { diff --git a/database/db/dsn_test.go b/database/db/dsn_test.go index b0faeaf48..64ce27598 100644 --- a/database/db/dsn_test.go +++ b/database/db/dsn_test.go @@ -60,8 +60,8 @@ func (s *DsnTestSuite) TestPostgresql() { s.mockConfig.On("GetString", fmt.Sprintf("database.connections.%s.sslmode", connection)).Return(sslmode).Once() s.mockConfig.On("GetString", fmt.Sprintf("database.connections.%s.timezone", connection)).Return(timezone).Once() - s.Equal(fmt.Sprintf("host=%s user=%s password=%s dbname=%s port=%d sslmode=%s TimeZone=%s", - testHost, testUsername, testPassword, testDatabase, testPort, sslmode, timezone), dsn.Postgresql(testConfig)) + s.Equal(fmt.Sprintf("postgres://%s:%s@%s:%d/%s?sslmode=%s&timezone=%s", + testUsername, testPassword, testHost, testPort, testDatabase, sslmode, timezone), dsn.Postgresql(testConfig)) } func (s *DsnTestSuite) TestSqlite() { diff --git a/database/gorm/conditions.go b/database/gorm/conditions.go new file mode 100644 index 000000000..163678c6d --- /dev/null +++ b/database/gorm/conditions.go @@ -0,0 +1,57 @@ +package gorm + +import ( + ormcontract "github.com/goravel/framework/contracts/database/orm" +) + +type Conditions struct { + distinct []any + group string + having *Having + join []Join + limit *int + lockForUpdate bool + model any + offset *int + omit []string + order []any + scopes []func(ormcontract.Query) ormcontract.Query + selectColumns *Select + sharedLock bool + table *Table + where []Where + with []With + withoutEvents bool + withTrashed bool +} + +type Having struct { + query any + args []any +} + +type Join struct { + query string + args []any +} + +type Select struct { + query any + args []any +} + +type Table struct { + name string + args []any +} + +type Where struct { + query any + args []any + or bool +} + +type With struct { + query string + args []any +} diff --git a/database/gorm/cursor.go b/database/gorm/cursor.go index e783901f7..3dd708097 100644 --- a/database/gorm/cursor.go +++ b/database/gorm/cursor.go @@ -13,7 +13,8 @@ import ( ) type CursorImpl struct { - row map[string]any + query *QueryImpl + row map[string]any } func (c *CursorImpl) Scan(value any) error { @@ -33,7 +34,21 @@ func (c *CursorImpl) Scan(value any) error { return err } - return decoder.Decode(c.row) + if err := decoder.Decode(c.row); err != nil { + return err + } + + for _, item := range c.query.conditions.with { + // Need to new a query, avoid to clear the conditions + query := c.query.new(c.query.instance) + // The new query must be cleared + query.clearConditions() + if err := query.Load(value, item.query, item.args...); err != nil { + return err + } + } + + return nil } func ToTimeHookFunc() mapstructure.DecodeHookFunc { diff --git a/database/gorm/dialector_test.go b/database/gorm/dialector_test.go index 2e2eee0cc..08d6b7f2d 100644 --- a/database/gorm/dialector_test.go +++ b/database/gorm/dialector_test.go @@ -63,8 +63,8 @@ func (s *DialectorTestSuite) TestPostgresql() { Return("UTC").Once() dialectors, err := dialector.Make([]databasecontract.Config{s.config}) s.Equal(postgres.New(postgres.Config{ - DSN: fmt.Sprintf("host=%s user=%s password=%s dbname=%s port=%d sslmode=%s TimeZone=%s", - s.config.Host, s.config.Username, s.config.Password, s.config.Database, s.config.Port, "disable", "UTC"), + DSN: fmt.Sprintf("postgres://%s:%s@%s:%d/%s?sslmode=%s&timezone=%s", + s.config.Username, s.config.Password, s.config.Host, s.config.Port, s.config.Database, "disable", "UTC"), }), dialectors[0]) s.Nil(err) } diff --git a/database/gorm/event.go b/database/gorm/event.go index 4c257622c..0df1e224f 100644 --- a/database/gorm/event.go +++ b/database/gorm/event.go @@ -28,6 +28,102 @@ func NewEvent(query *QueryImpl, model, dest any) *Event { } } +func (e *Event) ColumnNamesWithDbColumnNames() map[string]string { + if e.columnNamesWithDbColumnNames != nil { + return e.columnNamesWithDbColumnNames + } + + res := make(map[string]string) + var modelType reflect.Type + var modelValue reflect.Value + + if e.model != nil { + modelType = reflect.TypeOf(e.model) + modelValue = reflect.ValueOf(e.model) + } else { + modelType = reflect.TypeOf(e.dest) + modelValue = reflect.ValueOf(e.dest) + } + if modelType.Kind() == reflect.Pointer { + modelType = modelType.Elem() + modelValue = modelValue.Elem() + } + + for i := 0; i < modelType.NumField(); i++ { + if !modelType.Field(i).IsExported() { + continue + } + if modelType.Field(i).Name == "Model" && modelValue.Field(i).Type().Kind() == reflect.Struct { + structField := modelValue.Field(i).Type() + for j := 0; j < structField.NumField(); j++ { + if !structField.Field(i).IsExported() { + continue + } + dbColumn := structNameToDbColumnName(structField.Field(j).Name, structField.Field(j).Tag.Get("gorm")) + res[structField.Field(j).Name] = dbColumn + res[dbColumn] = dbColumn + } + } + + dbColumn := structNameToDbColumnName(modelType.Field(i).Name, modelType.Field(i).Tag.Get("gorm")) + res[modelType.Field(i).Name] = dbColumn + res[dbColumn] = dbColumn + } + + return res +} + +func (e *Event) Context() context.Context { + return e.query.ctx +} + +func (e *Event) DestOfMap() map[string]any { + if e.destOfMap != nil { + return e.destOfMap + } + + var destOfMap map[string]any + if destMap, ok := e.dest.(map[string]any); ok { + destOfMap = destMap + } else { + destType := reflect.TypeOf(e.dest) + if destType.Kind() == reflect.Pointer { + destType = destType.Elem() + } + if destType.Kind() == reflect.Struct { + destOfMap = structToMap(e.dest) + } + } + + e.destOfMap = destOfMap + + return e.destOfMap +} + +func (e *Event) GetAttribute(key string) any { + destOfMap := e.DestOfMap() + value, exist := destOfMap[e.toDBColumnName(key)] + if exist && e.validColumn(key) && e.validValue(key, value) { + return value + } + + return e.GetOriginal(key) +} + +func (e *Event) GetOriginal(key string, def ...any) any { + modelOfMap := e.ModelOfMap() + value, exist := modelOfMap[e.toDBColumnName(key)] + if exist { + return value + } + + if len(def) > 0 { + return def[0] + } + + return nil +} + func (e *Event) IsDirty(columns ...string) bool { destOfMap := e.DestOfMap() @@ -63,12 +159,22 @@ func (e *Event) IsClean(fields ...string) bool { return !e.IsDirty(fields...) } -func (e *Event) Query() orm.Query { - return NewQueryWithWithoutEvents(e.query.instance.Session(&gorm.Session{NewDB: true}), false, e.query.config) +func (e *Event) ModelOfMap() map[string]any { + if e.modelOfMap != nil { + return e.modelOfMap + } + + if e.model == nil { + return map[string]any{} + } + + e.modelOfMap = structToMap(e.model) + + return e.modelOfMap } -func (e *Event) Context() context.Context { - return e.query.instance.Statement.Context +func (e *Event) Query() orm.Query { + return NewQueryImpl(e.query.ctx, e.query.config, e.query.connection, e.query.instance.Session(&gorm.Session{NewDB: true}), nil) } func (e *Event) SetAttribute(key string, value any) { @@ -110,28 +216,35 @@ func (e *Event) SetAttribute(key string, value any) { } } -func (e *Event) GetAttribute(key string) any { - destOfMap := e.DestOfMap() - value, exist := destOfMap[e.toDBColumnName(key)] - if exist && e.validColumn(key) && e.validValue(key, value) { - return value +func (e *Event) dirty(destColumn string, destValue any) bool { + modelOfMap := e.ModelOfMap() + dbDestColumn := e.toDBColumnName(destColumn) + + if modelValue, exist := modelOfMap[dbDestColumn]; exist { + return !reflect.DeepEqual(modelValue, destValue) } - return e.GetOriginal(key) + return true } -func (e *Event) GetOriginal(key string, def ...any) any { - modelOfMap := e.ModelOfMap() - value, exist := modelOfMap[e.toDBColumnName(key)] - if exist { - return value +func (e *Event) equalColumnName(origin, source string) bool { + originDbColumnName := e.toDBColumnName(origin) + sourceDbColumnName := e.toDBColumnName(source) + + if originDbColumnName == "" || sourceDbColumnName == "" { + return false } - if len(def) > 0 { - return def[0] + return originDbColumnName == sourceDbColumnName +} + +func (e *Event) toDBColumnName(name string) string { + dbColumnName, exist := e.ColumnNamesWithDbColumnNames()[name] + if exist { + return dbColumnName } - return nil + return "" } func (e *Event) validColumn(name string) bool { @@ -194,119 +307,6 @@ func (e *Event) validValue(name string, value any) bool { return !valueValue.IsZero() } -func (e *Event) dirty(destColumn string, destValue any) bool { - modelOfMap := e.ModelOfMap() - dbDestColumn := e.toDBColumnName(destColumn) - - if modelValue, exist := modelOfMap[dbDestColumn]; exist { - return !reflect.DeepEqual(modelValue, destValue) - } - - return true -} - -func (e *Event) equalColumnName(origin, source string) bool { - originDbColumnName := e.toDBColumnName(origin) - sourceDbColumnName := e.toDBColumnName(source) - - if originDbColumnName == "" || sourceDbColumnName == "" { - return false - } - - return originDbColumnName == sourceDbColumnName -} - -func (e *Event) toDBColumnName(name string) string { - dbColumnName, exist := e.ColumnNamesWithDbColumnNames()[name] - if exist { - return dbColumnName - } - - return "" -} - -func (e *Event) ModelOfMap() map[string]any { - if e.modelOfMap != nil { - return e.modelOfMap - } - - if e.model == nil { - return map[string]any{} - } - - e.modelOfMap = structToMap(e.model) - - return e.modelOfMap -} - -func (e *Event) DestOfMap() map[string]any { - if e.destOfMap != nil { - return e.destOfMap - } - - var destOfMap map[string]any - if destMap, ok := e.dest.(map[string]any); ok { - destOfMap = destMap - } else { - destType := reflect.TypeOf(e.dest) - if destType.Kind() == reflect.Pointer { - destType = destType.Elem() - } - if destType.Kind() == reflect.Struct { - destOfMap = structToMap(e.dest) - } - } - - e.destOfMap = destOfMap - - return e.destOfMap -} - -func (e *Event) ColumnNamesWithDbColumnNames() map[string]string { - if e.columnNamesWithDbColumnNames != nil { - return e.columnNamesWithDbColumnNames - } - - res := make(map[string]string) - var modelType reflect.Type - var modelValue reflect.Value - - if e.model != nil { - modelType = reflect.TypeOf(e.model) - modelValue = reflect.ValueOf(e.model) - } else { - modelType = reflect.TypeOf(e.dest) - modelValue = reflect.ValueOf(e.dest) - } - if modelType.Kind() == reflect.Pointer { - modelType = modelType.Elem() - modelValue = modelValue.Elem() - } - - for i := 0; i < modelType.NumField(); i++ { - if !modelType.Field(i).IsExported() { - continue - } - if modelType.Field(i).Name == "Model" && modelValue.Field(i).Type().Kind() == reflect.Struct { - structField := modelValue.Field(i).Type() - for j := 0; j < structField.NumField(); j++ { - if !structField.Field(i).IsExported() { - continue - } - dbColumn := structNameToDbColumnName(structField.Field(j).Name, structField.Field(j).Tag.Get("gorm")) - res[structField.Field(j).Name] = dbColumn - res[dbColumn] = dbColumn - } - } - - dbColumn := structNameToDbColumnName(modelType.Field(i).Name, modelType.Field(i).Tag.Get("gorm")) - res[modelType.Field(i).Name] = dbColumn - res[dbColumn] = dbColumn - } - - return res -} - func structToMap(data any) map[string]any { res := make(map[string]any) modelType := reflect.TypeOf(data) diff --git a/database/gorm/event_test.go b/database/gorm/event_test.go index 7f57ff614..c8febb446 100644 --- a/database/gorm/event_test.go +++ b/database/gorm/event_test.go @@ -21,12 +21,14 @@ type TestEventModel struct { var testNow = time.Now().Add(-1 * time.Second) var testEventModel = TestEventModel{Name: "name", Avatar: "avatar", IsAdmin: true, IsManage: 0, AdminAt: testNow, ManageAt: testNow, high: 1} -var testQuery = NewQueryWithWithoutEvents(&gorm.DB{ - Statement: &gorm.Statement{ - Selects: []string{}, - Omits: []string{}, +var testQuery = &QueryImpl{ + instance: &gorm.DB{ + Statement: &gorm.Statement{ + Selects: []string{}, + Omits: []string{}, + }, }, -}, false, nil) +} type EventTestSuite struct { suite.Suite @@ -48,13 +50,15 @@ func (s *EventTestSuite) SetupTest() { func (s *EventTestSuite) TestSetAttribute() { dest := map[string]any{"avatar": "avatar1"} - query := NewQueryWithWithoutEvents(&gorm.DB{ - Statement: &gorm.Statement{ - Selects: []string{}, - Omits: []string{}, - Dest: dest, + query := &QueryImpl{ + instance: &gorm.DB{ + Statement: &gorm.Statement{ + Selects: []string{}, + Omits: []string{}, + Dest: dest, + }, }, - }, false, nil) + } event := NewEvent(query, &testEventModel, dest) @@ -148,23 +152,27 @@ func (s *EventTestSuite) TestValidColumn() { s.True(event.validColumn("manage")) s.False(event.validColumn("age")) - event.query = NewQueryWithWithoutEvents(&gorm.DB{ - Statement: &gorm.Statement{ - Selects: []string{"name"}, - Omits: []string{}, + event.query = &QueryImpl{ + instance: &gorm.DB{ + Statement: &gorm.Statement{ + Selects: []string{"name"}, + Omits: []string{}, + }, }, - }, false, nil) + } s.True(event.validColumn("Name")) s.True(event.validColumn("name")) s.False(event.validColumn("avatar")) s.False(event.validColumn("Avatar")) - event.query = NewQueryWithWithoutEvents(&gorm.DB{ - Statement: &gorm.Statement{ - Selects: []string{}, - Omits: []string{"name"}, + event.query = &QueryImpl{ + instance: &gorm.DB{ + Statement: &gorm.Statement{ + Selects: []string{}, + Omits: []string{"name"}, + }, }, - }, false, nil) + } s.False(event.validColumn("Name")) s.False(event.validColumn("name")) s.True(event.validColumn("avatar")) diff --git a/database/gorm/query.go b/database/gorm/query.go index 873f8c191..09051f49c 100644 --- a/database/gorm/query.go +++ b/database/gorm/query.go @@ -2,6 +2,7 @@ package gorm import ( "context" + "database/sql" "errors" "fmt" "reflect" @@ -21,17 +22,35 @@ import ( "github.com/goravel/framework/support/database" ) -var QuerySet = wire.NewSet(NewQueryImpl, wire.Bind(new(ormcontract.Query), new(*QueryImpl))) +var QuerySet = wire.NewSet(BuildQueryImpl, wire.Bind(new(ormcontract.Query), new(*QueryImpl))) var _ ormcontract.Query = &QueryImpl{} type QueryImpl struct { - config config.Config - ctx context.Context - instance *gormio.DB - withoutEvents bool + conditions Conditions + config config.Config + connection string + ctx context.Context + instance *gormio.DB + queries map[string]*QueryImpl } -func NewQueryImpl(ctx context.Context, config config.Config, gorm Gorm) (*QueryImpl, error) { +func NewQueryImpl(ctx context.Context, config config.Config, connection string, db *gormio.DB, conditions *Conditions) *QueryImpl { + queryImpl := &QueryImpl{ + config: config, + connection: connection, + ctx: ctx, + instance: db, + queries: make(map[string]*QueryImpl), + } + + if conditions != nil { + queryImpl.conditions = *conditions + } + + return queryImpl +} + +func BuildQueryImpl(ctx context.Context, config config.Config, connection string, gorm Gorm) (*QueryImpl, error) { db, err := gorm.Make() if err != nil { return nil, err @@ -40,25 +59,19 @@ func NewQueryImpl(ctx context.Context, config config.Config, gorm Gorm) (*QueryI db = db.WithContext(ctx) } - return &QueryImpl{ - instance: db, - config: config, - ctx: ctx, - }, nil -} - -func NewQueryWithWithoutEvents(instance *gormio.DB, withoutEvents bool, config config.Config) *QueryImpl { - return &QueryImpl{instance: instance, withoutEvents: withoutEvents, config: config, ctx: instance.Statement.Context} + return NewQueryImpl(ctx, config, connection, db, nil), nil } func (r *QueryImpl) Association(association string) ormcontract.Association { - return r.instance.Association(association) + query := r.buildConditions() + + return query.instance.Association(association) } func (r *QueryImpl) Begin() (ormcontract.Transaction, error) { tx := r.instance.Begin() - return NewTransaction(tx, r.config), tx.Error + return NewTransaction(tx, r.config, r.connection), tx.Error } func (r *QueryImpl) Driver() ormcontract.Driver { @@ -66,33 +79,43 @@ func (r *QueryImpl) Driver() ormcontract.Driver { } func (r *QueryImpl) Count(count *int64) error { - return r.instance.Count(count).Error + query := r.buildConditions() + + return query.instance.Count(count).Error } func (r *QueryImpl) Create(value any) error { - if err := r.refreshConnection(value); err != nil { + query, err := r.refreshConnection(value) + if err != nil { return err } - if len(r.instance.Statement.Selects) > 0 && len(r.instance.Statement.Omits) > 0 { + query = query.buildConditions() + + if len(query.instance.Statement.Selects) > 0 && len(query.instance.Statement.Omits) > 0 { return errors.New("cannot set Select and Omits at the same time") } - if len(r.instance.Statement.Selects) > 0 { - return r.selectCreate(value) + if len(query.instance.Statement.Selects) > 0 { + return query.selectCreate(value) } - if len(r.instance.Statement.Omits) > 0 { - return r.omitCreate(value) + if len(query.instance.Statement.Omits) > 0 { + return query.omitCreate(value) } - return r.create(value) + return query.create(value) } func (r *QueryImpl) Cursor() (chan ormcontract.Cursor, error) { + with := r.conditions.with + query := r.buildConditions() + r.conditions.with = with + var err error cursorChan := make(chan ormcontract.Cursor) go func() { - rows, err := r.instance.Rows() + var rows *sql.Rows + rows, err = query.instance.Rows() if err != nil { return } @@ -100,11 +123,11 @@ func (r *QueryImpl) Cursor() (chan ormcontract.Cursor, error) { for rows.Next() { val := make(map[string]any) - err := r.instance.ScanRows(rows, val) + err = query.instance.ScanRows(rows, val) if err != nil { return } - cursorChan <- &CursorImpl{row: val} + cursorChan <- &CursorImpl{query: r, row: val} } close(cursorChan) }() @@ -112,19 +135,22 @@ func (r *QueryImpl) Cursor() (chan ormcontract.Cursor, error) { } func (r *QueryImpl) Delete(dest any, conds ...any) (*ormcontract.Result, error) { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return nil, err } - if err := r.deleting(dest); err != nil { + query = query.buildConditions() + + if err := query.deleting(dest); err != nil { return nil, err } - res := r.instance.Delete(dest, conds...) + res := query.instance.Delete(dest, conds...) if res.Error != nil { return nil, res.Error } - if err := r.deleted(dest); err != nil { + if err := query.deleted(dest); err != nil { return nil, err } @@ -134,13 +160,15 @@ func (r *QueryImpl) Delete(dest any, conds ...any) (*ormcontract.Result, error) } func (r *QueryImpl) Distinct(args ...any) ormcontract.Query { - tx := r.instance.Distinct(args...) + conditions := r.conditions + conditions.distinct = append(conditions.distinct, args...) - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Exec(sql string, values ...any) (*ormcontract.Result, error) { - result := r.instance.Exec(sql, values...) + query := r.buildConditions() + result := query.instance.Exec(sql, values...) return &ormcontract.Result{ RowsAffected: result.RowsAffected, @@ -148,28 +176,36 @@ func (r *QueryImpl) Exec(sql string, values ...any) (*ormcontract.Result, error) } func (r *QueryImpl) Find(dest any, conds ...any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } + + query = query.buildConditions() + if err := filterFindConditions(conds...); err != nil { return err } - if err := r.instance.Find(dest, conds...).Error; err != nil { + if err := query.instance.Find(dest, conds...).Error; err != nil { return err } - return r.retrieved(dest) + return query.retrieved(dest) } func (r *QueryImpl) FindOrFail(dest any, conds ...any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } + + query = query.buildConditions() + if err := filterFindConditions(conds...); err != nil { return err } - res := r.instance.Find(dest, conds...) + res := query.instance.Find(dest, conds...) if err := res.Error; err != nil { return err } @@ -178,14 +214,18 @@ func (r *QueryImpl) FindOrFail(dest any, conds ...any) error { return orm.ErrRecordNotFound } - return r.retrieved(dest) + return query.retrieved(dest) } func (r *QueryImpl) First(dest any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } - res := r.instance.First(dest) + + query = query.buildConditions() + + res := query.instance.First(dest) if res.Error != nil { if errors.Is(res.Error, gormio.ErrRecordNotFound) { return nil @@ -194,15 +234,18 @@ func (r *QueryImpl) First(dest any) error { return res.Error } - return r.retrieved(dest) + return query.retrieved(dest) } func (r *QueryImpl) FirstOr(dest any, callback func() error) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } - err := r.instance.First(dest).Error - if err != nil { + + query = query.buildConditions() + + if err := query.instance.First(dest).Error; err != nil { if errors.Is(err, gormio.ErrRecordNotFound) { return callback() } @@ -210,40 +253,47 @@ func (r *QueryImpl) FirstOr(dest any, callback func() error) error { return err } - return r.retrieved(dest) + return query.retrieved(dest) } func (r *QueryImpl) FirstOrCreate(dest any, conds ...any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } + + query = query.buildConditions() + if len(conds) == 0 { return errors.New("query condition is require") } var res *gormio.DB if len(conds) > 1 { - res = r.instance.Attrs(conds[1]).FirstOrInit(dest, conds[0]) + res = query.instance.Attrs(conds[1]).FirstOrInit(dest, conds[0]) } else { - res = r.instance.FirstOrInit(dest, conds[0]) + res = query.instance.FirstOrInit(dest, conds[0]) } if res.Error != nil { return res.Error } if res.RowsAffected > 0 { - return r.retrieved(dest) + return query.retrieved(dest) } - return r.Create(dest) + return query.Create(dest) } func (r *QueryImpl) FirstOrFail(dest any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } - err := r.instance.First(dest).Error - if err != nil { + + query = query.buildConditions() + + if err := query.instance.First(dest).Error; err != nil { if errors.Is(err, gormio.ErrRecordNotFound) { return orm.ErrRecordNotFound } @@ -251,45 +301,53 @@ func (r *QueryImpl) FirstOrFail(dest any) error { return err } - return r.retrieved(dest) + return query.retrieved(dest) } func (r *QueryImpl) FirstOrNew(dest any, attributes any, values ...any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } + + query = query.buildConditions() + var res *gormio.DB if len(values) > 0 { - res = r.instance.Attrs(values[0]).FirstOrInit(dest, attributes) + res = query.instance.Attrs(values[0]).FirstOrInit(dest, attributes) } else { - res = r.instance.FirstOrInit(dest, attributes) + res = query.instance.FirstOrInit(dest, attributes) } if res.Error != nil { return res.Error } if res.RowsAffected > 0 { - return r.retrieved(dest) + return query.retrieved(dest) } return nil } func (r *QueryImpl) ForceDelete(value any, conds ...any) (*ormcontract.Result, error) { - if err := r.refreshConnection(value); err != nil { + query, err := r.refreshConnection(value) + if err != nil { return nil, err } - if err := r.forceDeleting(value); err != nil { + + query = query.buildConditions() + + if err := query.forceDeleting(value); err != nil { return nil, err } - res := r.instance.Unscoped().Delete(value, conds...) + res := query.instance.Unscoped().Delete(value, conds...) if res.Error != nil { return nil, res.Error } if res.RowsAffected > 0 { - if err := r.forceDeleted(value); err != nil { + if err := query.forceDeleted(value); err != nil { return nil, err } } @@ -304,15 +362,20 @@ func (r *QueryImpl) Get(dest any) error { } func (r *QueryImpl) Group(name string) ormcontract.Query { - tx := r.instance.Group(name) + conditions := r.conditions + conditions.group = name - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Having(query any, args ...any) ormcontract.Query { - tx := r.instance.Having(query, args...) + conditions := r.conditions + conditions.having = &Having{ + query: query, + args: args, + } - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Instance() *gormio.DB { @@ -320,14 +383,20 @@ func (r *QueryImpl) Instance() *gormio.DB { } func (r *QueryImpl) Join(query string, args ...any) ormcontract.Query { - tx := r.instance.Joins(query, args...) - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + conditions := r.conditions + conditions.join = append(conditions.join, Join{ + query: query, + args: args, + }) + + return r.setConditions(conditions) } func (r *QueryImpl) Limit(limit int) ormcontract.Query { - tx := r.instance.Limit(limit) + conditions := r.conditions + conditions.limit = &limit - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Load(model any, relation string, args ...any) error { @@ -345,8 +414,7 @@ func (r *QueryImpl) Load(model any, relation string, args ...any) error { } copyDest := copyStruct(model) - query := r.With(relation, args...) - err := query.Find(model) + err := r.With(relation, args...).Find(model) t := destType.Elem() v := reflect.ValueOf(model).Elem() @@ -397,133 +465,138 @@ func (r *QueryImpl) LoadMissing(model any, relation string, args ...any) error { } func (r *QueryImpl) LockForUpdate() ormcontract.Query { - driver := r.instance.Name() - mysqlDialector := mysql.Dialector{} - postgresqlDialector := postgres.Dialector{} - sqlserverDialector := sqlserver.Dialector{} - - if driver == mysqlDialector.Name() || driver == postgresqlDialector.Name() { - tx := r.instance.Clauses(clause.Locking{Strength: "UPDATE"}) + conditions := r.conditions + conditions.lockForUpdate = true - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) - } else if driver == sqlserverDialector.Name() { - tx := r.instance.Clauses(hints.With("rowlock", "updlock", "holdlock")) - - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) - } - - return r + return r.setConditions(conditions) } func (r *QueryImpl) Model(value any) ormcontract.Query { - if err := r.refreshConnection(value); err != nil { - return nil - } - tx := r.instance.Model(value) + conditions := r.conditions + conditions.model = value - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Offset(offset int) ormcontract.Query { - tx := r.instance.Offset(offset) + conditions := r.conditions + conditions.offset = &offset - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Omit(columns ...string) ormcontract.Query { - tx := r.instance.Omit(columns...) + conditions := r.conditions + conditions.omit = columns - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Order(value any) ormcontract.Query { - tx := r.instance.Order(value) + conditions := r.conditions + conditions.order = append(r.conditions.order, value) - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) OrWhere(query any, args ...any) ormcontract.Query { - tx := r.instance.Or(query, args...) + conditions := r.conditions + conditions.where = append(r.conditions.where, Where{ + query: query, + args: args, + or: true, + }) - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Paginate(page, limit int, dest any, total *int64) error { + query, err := r.refreshConnection(dest) + if err != nil { + return err + } + + query = query.buildConditions() + offset := (page - 1) * limit if total != nil { - if r.instance.Statement.Table == "" && r.instance.Statement.Model == nil { - if err := r.Model(dest).Count(total); err != nil { + if query.conditions.table == nil && query.conditions.model == nil { + if err := query.Model(dest).Count(total); err != nil { return err } } else { - if err := r.Count(total); err != nil { + if err := query.Count(total); err != nil { return err } } } - return r.Offset(offset).Limit(limit).Find(dest) + return query.Offset(offset).Limit(limit).Find(dest) } func (r *QueryImpl) Pluck(column string, dest any) error { - return r.instance.Pluck(column, dest).Error + query := r.buildConditions() + + return query.instance.Pluck(column, dest).Error } func (r *QueryImpl) Raw(sql string, values ...any) ormcontract.Query { - tx := r.instance.Raw(sql, values...) - - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.new(r.instance.Raw(sql, values...)) } func (r *QueryImpl) Save(value any) error { - if err := r.refreshConnection(value); err != nil { + query, err := r.refreshConnection(value) + if err != nil { return err } - if len(r.instance.Statement.Selects) > 0 && len(r.instance.Statement.Omits) > 0 { + + query = query.buildConditions() + + if len(query.instance.Statement.Selects) > 0 && len(query.instance.Statement.Omits) > 0 { return errors.New("cannot set Select and Omits at the same time") } - model := r.instance.Statement.Model + model := query.instance.Statement.Model id := database.GetID(value) update := id != nil - if err := r.saving(model, value); err != nil { + if err := query.saving(model, value); err != nil { return err } if update { - if err := r.updating(model, value); err != nil { + if err := query.updating(model, value); err != nil { return err } } else { - if err := r.creating(value); err != nil { + if err := query.creating(value); err != nil { return err } } - if len(r.instance.Statement.Selects) > 0 { - if err := r.selectSave(value); err != nil { + if len(query.instance.Statement.Selects) > 0 { + if err := query.selectSave(value); err != nil { return err } - } else if len(r.instance.Statement.Omits) > 0 { - if err := r.omitSave(value); err != nil { + } else if len(query.instance.Statement.Omits) > 0 { + if err := query.omitSave(value); err != nil { return err } } else { - if err := r.save(value); err != nil { + if err := query.save(value); err != nil { return err } } if update { - if err := r.updated(model, value); err != nil { + if err := query.updated(model, value); err != nil { return err } } else { - if err := r.created(value); err != nil { + if err := query.created(value); err != nil { return err } } - if err := r.saved(model, value); err != nil { + if err := query.saved(model, value); err != nil { return err } @@ -535,98 +608,98 @@ func (r *QueryImpl) SaveQuietly(value any) error { } func (r *QueryImpl) Scan(dest any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } - return r.instance.Scan(dest).Error -} + query = query.buildConditions() -func (r *QueryImpl) Select(query any, args ...any) ormcontract.Query { - tx := r.instance.Select(query, args...) - - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return query.instance.Scan(dest).Error } func (r *QueryImpl) Scopes(funcs ...func(ormcontract.Query) ormcontract.Query) ormcontract.Query { - var gormFuncs []func(*gormio.DB) *gormio.DB - for _, item := range funcs { - gormFuncs = append(gormFuncs, func(tx *gormio.DB) *gormio.DB { - item(NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config)) + conditions := r.conditions + conditions.scopes = append(r.conditions.scopes, funcs...) - return tx - }) + return r.setConditions(conditions) +} + +func (r *QueryImpl) Select(query any, args ...any) ormcontract.Query { + conditions := r.conditions + conditions.selectColumns = &Select{ + query: query, + args: args, } - tx := r.instance.Scopes(gormFuncs...) + return r.setConditions(conditions) +} - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) +func (r *QueryImpl) SetContext(ctx context.Context) { + r.ctx = ctx + r.instance.Statement.Context = ctx } func (r *QueryImpl) SharedLock() ormcontract.Query { - driver := r.instance.Name() - mysqlDialector := mysql.Dialector{} - postgresqlDialector := postgres.Dialector{} - sqlserverDialector := sqlserver.Dialector{} - - if driver == mysqlDialector.Name() || driver == postgresqlDialector.Name() { - tx := r.instance.Clauses(clause.Locking{Strength: "SHARE"}) - - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) - } else if driver == sqlserverDialector.Name() { - tx := r.instance.Clauses(hints.With("rowlock", "holdlock")) - - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) - } + conditions := r.conditions + conditions.sharedLock = true - return r + return r.setConditions(conditions) } func (r *QueryImpl) Sum(column string, dest any) error { - return r.instance.Select("SUM(" + column + ")").Row().Scan(dest) + query := r.buildConditions() + + return query.instance.Select("SUM(" + column + ")").Row().Scan(dest) } func (r *QueryImpl) Table(name string, args ...any) ormcontract.Query { - tx := r.instance.Table(name, args...) + conditions := r.conditions + conditions.table = &Table{ + name: name, + args: args, + } - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } func (r *QueryImpl) Update(column any, value ...any) (*ormcontract.Result, error) { + query := r.buildConditions() + if _, ok := column.(string); !ok && len(value) > 0 { return nil, errors.New("parameter error, please check the document") } var singleUpdate bool - model := r.instance.Statement.Model + model := query.instance.Statement.Model if model != nil { id := database.GetID(model) singleUpdate = id != nil } if c, ok := column.(string); ok && len(value) > 0 { - r.instance.Statement.Dest = map[string]any{c: value[0]} + query.instance.Statement.Dest = map[string]any{c: value[0]} } if len(value) == 0 { - r.instance.Statement.Dest = column + query.instance.Statement.Dest = column } if singleUpdate { - if err := r.saving(model, r.instance.Statement.Dest); err != nil { + if err := query.saving(model, query.instance.Statement.Dest); err != nil { return nil, err } - if err := r.updating(model, r.instance.Statement.Dest); err != nil { + if err := query.updating(model, query.instance.Statement.Dest); err != nil { return nil, err } } - res, err := r.updates(r.instance.Statement.Dest) + res, err := query.updates(query.instance.Statement.Dest) if singleUpdate && err == nil { - if err := r.updated(model, r.instance.Statement.Dest); err != nil { + if err := query.updated(model, query.instance.Statement.Dest); err != nil { return nil, err } - if err := r.saved(model, r.instance.Statement.Dest); err != nil { + if err := query.saved(model, query.instance.Statement.Dest); err != nil { return nil, err } } @@ -635,104 +708,344 @@ func (r *QueryImpl) Update(column any, value ...any) (*ormcontract.Result, error } func (r *QueryImpl) UpdateOrCreate(dest any, attributes any, values any) error { - if err := r.refreshConnection(dest); err != nil { + query, err := r.refreshConnection(dest) + if err != nil { return err } - res := r.instance.Assign(values).FirstOrInit(dest, attributes) + + query = query.buildConditions() + + res := query.instance.Assign(values).FirstOrInit(dest, attributes) if res.Error != nil { return res.Error } if res.RowsAffected > 0 { - return r.Save(dest) + return query.Save(dest) } - return r.Create(dest) + return query.Create(dest) } func (r *QueryImpl) Where(query any, args ...any) ormcontract.Query { - tx := r.instance.Where(query, args...) + conditions := r.conditions + conditions.where = append(r.conditions.where, Where{ + query: query, + args: args, + }) - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) +} + +func (r *QueryImpl) With(query string, args ...any) ormcontract.Query { + conditions := r.conditions + conditions.with = append(r.conditions.with, With{ + query: query, + args: args, + }) + + return r.setConditions(conditions) } func (r *QueryImpl) WithoutEvents() ormcontract.Query { - return NewQueryWithWithoutEvents(r.instance, true, r.config) + conditions := r.conditions + conditions.withoutEvents = true + + return r.setConditions(conditions) } func (r *QueryImpl) WithTrashed() ormcontract.Query { - tx := r.instance.Unscoped() + conditions := r.conditions + conditions.withTrashed = true - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return r.setConditions(conditions) } -func (r *QueryImpl) With(query string, args ...any) ormcontract.Query { - if len(args) == 1 { - switch arg := args[0].(type) { - case func(ormcontract.Query) ormcontract.Query: - newArgs := []any{ - func(tx *gormio.DB) *gormio.DB { - query := arg(NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config)) - - return query.(*QueryImpl).instance - }, - } +func (r *QueryImpl) buildConditions() *QueryImpl { + query := r.buildModel() + db := query.instance + db = query.buildDistinct(db) + db = query.buildGroup(db) + db = query.buildHaving(db) + db = query.buildJoin(db) + db = query.buildLockForUpdate(db) + db = query.buildLimit(db) + db = query.buildOrder(db) + db = query.buildOffset(db) + db = query.buildOmit(db) + db = query.buildScopes(db) + db = query.buildSelectColumns(db) + db = query.buildSharedLock(db) + db = query.buildTable(db) + db = query.buildWith(db) + db = query.buildWithTrashed(db) + db = query.buildWhere(db) + + return query.new(db) +} + +func (r *QueryImpl) buildDistinct(db *gormio.DB) *gormio.DB { + if len(r.conditions.distinct) == 0 { + return db + } - tx := r.instance.Preload(query, newArgs...) + db = db.Distinct(r.conditions.distinct...) + r.conditions.distinct = nil - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) - } + return db +} + +func (r *QueryImpl) buildGroup(db *gormio.DB) *gormio.DB { + if r.conditions.group == "" { + return db } - tx := r.instance.Preload(query, args...) + db = db.Group(r.conditions.group) + r.conditions.group = "" - return NewQueryWithWithoutEvents(tx, r.withoutEvents, r.config) + return db } -func (r *QueryImpl) refreshConnection(value any) error { - model, ok := value.(ormcontract.ConnectionModel) - if !ok { +func (r *QueryImpl) buildHaving(db *gormio.DB) *gormio.DB { + if r.conditions.having == nil { + return db + } + + db = db.Having(r.conditions.having.query, r.conditions.having.args...) + r.conditions.having = nil + + return db +} + +func (r *QueryImpl) buildJoin(db *gormio.DB) *gormio.DB { + if r.conditions.join == nil { + return db + } + + for _, item := range r.conditions.join { + db = db.Joins(item.query, item.args...) + } + + r.conditions.join = nil + + return db +} + +func (r *QueryImpl) buildLimit(db *gormio.DB) *gormio.DB { + if r.conditions.limit == nil { + return db + } + + db = db.Limit(*r.conditions.limit) + r.conditions.limit = nil + + return db +} + +func (r *QueryImpl) buildLockForUpdate(db *gormio.DB) *gormio.DB { + if !r.conditions.lockForUpdate { + return db + } + + driver := r.instance.Name() + mysqlDialector := mysql.Dialector{} + postgresqlDialector := postgres.Dialector{} + sqlserverDialector := sqlserver.Dialector{} + + if driver == mysqlDialector.Name() || driver == postgresqlDialector.Name() { + return db.Clauses(clause.Locking{Strength: "UPDATE"}) + } else if driver == sqlserverDialector.Name() { + return db.Clauses(hints.With("rowlock", "updlock", "holdlock")) + } + + r.conditions.lockForUpdate = false + + return db +} + +func (r *QueryImpl) buildModel() *QueryImpl { + if r.conditions.model == nil { + return r + } + + query, err := r.refreshConnection(r.conditions.model) + if err != nil { return nil } - conn := model.Connection() - if conn == "" { - conn = r.config.GetString("database.default") + + return query.new(query.instance.Model(r.conditions.model)) +} + +func (r *QueryImpl) buildOffset(db *gormio.DB) *gormio.DB { + if r.conditions.offset == nil { + return db } - driver := driver2gorm(r.config.GetString(fmt.Sprintf("database.connections.%s.driver", conn))) - if driver == "" { - return fmt.Errorf("connection %s driver is not supported", conn) + + db = db.Offset(*r.conditions.offset) + r.conditions.offset = nil + + return db +} + +func (r *QueryImpl) buildOmit(db *gormio.DB) *gormio.DB { + if len(r.conditions.omit) == 0 { + return db } - // if a driver is not the same, we need to refresh the connection - if driver != r.instance.Name() { - query, err := InitializeQuery(r.ctx, r.config, conn) - if err != nil { - return err - } - dbInstance := query.instance - stmt := r.instance.Statement - stmt.DB = dbInstance.Statement.DB - stmt.ConnPool = dbInstance.ConnPool - if r.ctx != nil { - dbInstance = dbInstance.WithContext(r.ctx) + + db = db.Omit(r.conditions.omit...) + r.conditions.omit = nil + + return db +} + +func (r *QueryImpl) buildOrder(db *gormio.DB) *gormio.DB { + if len(r.conditions.order) == 0 { + return db + } + + for _, order := range r.conditions.order { + db = db.Order(order) + } + + r.conditions.order = nil + + return db +} + +func (r *QueryImpl) buildSelectColumns(db *gormio.DB) *gormio.DB { + if r.conditions.selectColumns == nil { + return db + } + + db = db.Select(r.conditions.selectColumns.query, r.conditions.selectColumns.args...) + r.conditions.selectColumns = nil + + return db +} + +func (r *QueryImpl) buildScopes(db *gormio.DB) *gormio.DB { + if len(r.conditions.scopes) == 0 { + return db + } + + var gormFuncs []func(*gormio.DB) *gormio.DB + for _, scope := range r.conditions.scopes { + gormFuncs = append(gormFuncs, func(tx *gormio.DB) *gormio.DB { + queryImpl := r.new(tx) + query := scope(queryImpl) + queryImpl = query.(*QueryImpl) + queryImpl = queryImpl.buildConditions() + + return queryImpl.instance + }) + } + + db = db.Scopes(gormFuncs...) + r.conditions.scopes = nil + + return db +} + +func (r *QueryImpl) buildSharedLock(db *gormio.DB) *gormio.DB { + if !r.conditions.sharedLock { + return db + } + + driver := r.instance.Name() + mysqlDialector := mysql.Dialector{} + postgresqlDialector := postgres.Dialector{} + sqlserverDialector := sqlserver.Dialector{} + + if driver == mysqlDialector.Name() || driver == postgresqlDialector.Name() { + return db.Clauses(clause.Locking{Strength: "SHARE"}) + } else if driver == sqlserverDialector.Name() { + return db.Clauses(hints.With("rowlock", "holdlock")) + } + + r.conditions.sharedLock = false + + return db +} + +func (r *QueryImpl) buildTable(db *gormio.DB) *gormio.DB { + if r.conditions.table == nil { + return db + } + + db = db.Table(r.conditions.table.name, r.conditions.table.args...) + r.conditions.table = nil + + return db +} + +func (r *QueryImpl) buildWhere(db *gormio.DB) *gormio.DB { + if len(r.conditions.where) == 0 { + return db + } + + for _, item := range r.conditions.where { + if item.or { + db = db.Or(item.query, item.args...) + } else { + db = db.Where(item.query, item.args...) } - dbInstance.Statement = stmt - r.instance = dbInstance } - return nil + + r.conditions.where = nil + + return db } -func (r *QueryImpl) selectCreate(value any) error { - if len(r.instance.Statement.Selects) > 1 { - for _, val := range r.instance.Statement.Selects { - if val == orm.Associations { - return errors.New("cannot set orm.Associations and other fields at the same time") +func (r *QueryImpl) buildWith(db *gormio.DB) *gormio.DB { + if len(r.conditions.with) == 0 { + return db + } + + for _, item := range r.conditions.with { + isSet := false + if len(item.args) == 1 { + if arg, ok := item.args[0].(func(ormcontract.Query) ormcontract.Query); ok { + newArgs := []any{ + func(tx *gormio.DB) *gormio.DB { + queryImpl := NewQueryImpl(r.ctx, r.config, r.connection, tx, nil) + query := arg(queryImpl) + queryImpl = query.(*QueryImpl) + queryImpl = queryImpl.buildConditions() + + return queryImpl.instance + }, + } + + db = db.Preload(item.query, newArgs...) + isSet = true } } + + if !isSet { + db = db.Preload(item.query, item.args...) + } } - if len(r.instance.Statement.Selects) == 1 && r.instance.Statement.Selects[0] == orm.Associations { - r.instance.Statement.Selects = []string{} + r.conditions.with = nil + + return db +} + +func (r *QueryImpl) buildWithTrashed(db *gormio.DB) *gormio.DB { + if !r.conditions.withTrashed { + return db } + db = db.Unscoped() + r.conditions.withTrashed = false + + return db +} + +func (r *QueryImpl) clearConditions() { + r.conditions = Conditions{} +} + +func (r *QueryImpl) create(value any) error { if err := r.saving(nil, value); err != nil { return err } @@ -740,7 +1053,7 @@ func (r *QueryImpl) selectCreate(value any) error { return err } - if err := r.instance.Create(value).Error; err != nil { + if err := r.instance.Omit(orm.Associations).Create(value).Error; err != nil { return err } @@ -754,6 +1067,81 @@ func (r *QueryImpl) selectCreate(value any) error { return nil } +func (r *QueryImpl) created(dest any) error { + return r.event(ormcontract.EventCreated, nil, dest) +} + +func (r *QueryImpl) creating(dest any) error { + return r.event(ormcontract.EventCreating, nil, dest) +} + +func (r *QueryImpl) deleting(dest any) error { + return r.event(ormcontract.EventDeleting, nil, dest) +} + +func (r *QueryImpl) deleted(dest any) error { + return r.event(ormcontract.EventDeleted, nil, dest) +} + +func (r *QueryImpl) forceDeleting(dest any) error { + return r.event(ormcontract.EventForceDeleting, nil, dest) +} + +func (r *QueryImpl) forceDeleted(dest any) error { + return r.event(ormcontract.EventForceDeleted, nil, dest) +} + +func (r *QueryImpl) event(event ormcontract.EventType, model, dest any) error { + if r.conditions.withoutEvents { + return nil + } + + instance := NewEvent(r, model, dest) + + if dispatchesEvents, exist := dest.(ormcontract.DispatchesEvents); exist { + if event, exist := dispatchesEvents.DispatchesEvents()[event]; exist { + return event(instance) + } + + return nil + } + if model != nil { + if dispatchesEvents, exist := model.(ormcontract.DispatchesEvents); exist { + if event, exist := dispatchesEvents.DispatchesEvents()[event]; exist { + return event(instance) + } + } + + return nil + } + + if observer := observer(dest); observer != nil { + if observerEvent := observerEvent(event, observer); observerEvent != nil { + return observerEvent(instance) + } + + return nil + } + + if model != nil { + if observer := observer(model); observer != nil { + if observerEvent := observerEvent(event, observer); observerEvent != nil { + return observerEvent(instance) + } + + return nil + } + } + + return nil +} + +func (r *QueryImpl) new(db *gormio.DB) *QueryImpl { + query := NewQueryImpl(r.ctx, r.config, r.connection, db, &r.conditions) + + return query +} + func (r *QueryImpl) omitCreate(value any) error { if len(r.instance.Statement.Omits) > 1 { for _, val := range r.instance.Statement.Omits { @@ -794,7 +1182,73 @@ func (r *QueryImpl) omitCreate(value any) error { return nil } -func (r *QueryImpl) create(value any) error { +func (r *QueryImpl) omitSave(value any) error { + for _, val := range r.instance.Statement.Omits { + if val == orm.Associations { + return r.instance.Omit(orm.Associations).Save(value).Error + } + } + + return r.instance.Save(value).Error +} + +func (r *QueryImpl) refreshConnection(model any) (*QueryImpl, error) { + connection, err := getModelConnection(model) + if err != nil { + return nil, err + } + if connection == "" || connection == r.connection { + return r, nil + } + + query, ok := r.queries[connection] + if !ok { + var err error + query, err = InitializeQuery(r.ctx, r.config, connection) + if err != nil { + return nil, err + } + + if r.queries == nil { + r.queries = make(map[string]*QueryImpl) + } + r.queries[connection] = query + } + + query.conditions = r.conditions + + return query, nil +} + +func (r *QueryImpl) retrieved(dest any) error { + return r.event(ormcontract.EventRetrieved, nil, dest) +} + +func (r *QueryImpl) save(value any) error { + return r.instance.Omit(orm.Associations).Save(value).Error +} + +func (r *QueryImpl) saved(model, dest any) error { + return r.event(ormcontract.EventSaved, model, dest) +} + +func (r *QueryImpl) saving(model, dest any) error { + return r.event(ormcontract.EventSaving, model, dest) +} + +func (r *QueryImpl) selectCreate(value any) error { + if len(r.instance.Statement.Selects) > 1 { + for _, val := range r.instance.Statement.Selects { + if val == orm.Associations { + return errors.New("cannot set orm.Associations and other fields at the same time") + } + } + } + + if len(r.instance.Statement.Selects) == 1 && r.instance.Statement.Selects[0] == orm.Associations { + r.instance.Statement.Selects = []string{} + } + if err := r.saving(nil, value); err != nil { return err } @@ -802,7 +1256,7 @@ func (r *QueryImpl) create(value any) error { return err } - if err := r.instance.Omit(orm.Associations).Create(value).Error; err != nil { + if err := r.instance.Create(value).Error; err != nil { return err } @@ -830,18 +1284,19 @@ func (r *QueryImpl) selectSave(value any) error { return nil } -func (r *QueryImpl) omitSave(value any) error { - for _, val := range r.instance.Statement.Omits { - if val == orm.Associations { - return r.instance.Omit(orm.Associations).Save(value).Error - } - } +func (r *QueryImpl) setConditions(conditions Conditions) *QueryImpl { + query := r.new(r.instance) + query.conditions = conditions - return r.instance.Save(value).Error + return query } -func (r *QueryImpl) save(value any) error { - return r.instance.Omit(orm.Associations).Save(value).Error +func (r *QueryImpl) updating(model, dest any) error { + return r.event(ormcontract.EventUpdating, model, dest) +} + +func (r *QueryImpl) updated(model, dest any) error { + return r.event(ormcontract.EventUpdated, model, dest) } func (r *QueryImpl) updates(values any) (*ormcontract.Result, error) { @@ -889,95 +1344,6 @@ func (r *QueryImpl) updates(values any) (*ormcontract.Result, error) { }, result.Error } -func (r *QueryImpl) retrieved(dest any) error { - return r.event(ormcontract.EventRetrieved, nil, dest) -} - -func (r *QueryImpl) updating(model, dest any) error { - return r.event(ormcontract.EventUpdating, model, dest) -} - -func (r *QueryImpl) updated(model, dest any) error { - return r.event(ormcontract.EventUpdated, model, dest) -} - -func (r *QueryImpl) saving(model, dest any) error { - return r.event(ormcontract.EventSaving, model, dest) -} - -func (r *QueryImpl) saved(model, dest any) error { - return r.event(ormcontract.EventSaved, model, dest) -} - -func (r *QueryImpl) creating(dest any) error { - return r.event(ormcontract.EventCreating, nil, dest) -} - -func (r *QueryImpl) created(dest any) error { - return r.event(ormcontract.EventCreated, nil, dest) -} - -func (r *QueryImpl) deleting(dest any) error { - return r.event(ormcontract.EventDeleting, nil, dest) -} - -func (r *QueryImpl) deleted(dest any) error { - return r.event(ormcontract.EventDeleted, nil, dest) -} - -func (r *QueryImpl) forceDeleting(dest any) error { - return r.event(ormcontract.EventForceDeleting, nil, dest) -} - -func (r *QueryImpl) forceDeleted(dest any) error { - return r.event(ormcontract.EventForceDeleted, nil, dest) -} - -func (r *QueryImpl) event(event ormcontract.EventType, model, dest any) error { - if r.withoutEvents { - return nil - } - - instance := NewEvent(r, model, dest) - - if dispatchesEvents, exist := dest.(ormcontract.DispatchesEvents); exist { - if event, exist := dispatchesEvents.DispatchesEvents()[event]; exist { - return event(instance) - } - - return nil - } - if model != nil { - if dispatchesEvents, exist := model.(ormcontract.DispatchesEvents); exist { - if event, exist := dispatchesEvents.DispatchesEvents()[event]; exist { - return event(instance) - } - } - - return nil - } - - if observer := observer(dest); observer != nil { - if observerEvent := observerEvent(event, observer); observerEvent != nil { - return observerEvent(instance) - } - - return nil - } - - if model != nil { - if observer := observer(model); observer != nil { - if observerEvent := observerEvent(event, observer); observerEvent != nil { - return observerEvent(instance) - } - - return nil - } - } - - return nil -} - func filterFindConditions(conds ...any) error { if len(conds) > 0 { switch cond := conds[0].(type) { @@ -999,6 +1365,37 @@ func filterFindConditions(conds ...any) error { return nil } +func getModelConnection(model any) (string, error) { + value1 := reflect.ValueOf(model) + if value1.Kind() == reflect.Ptr && value1.IsNil() { + value1 = reflect.New(value1.Type().Elem()) + } + modelType := reflect.Indirect(value1).Type() + + if modelType.Kind() == reflect.Interface { + modelType = reflect.Indirect(reflect.ValueOf(model)).Elem().Type() + } + + for modelType.Kind() == reflect.Slice || modelType.Kind() == reflect.Array || modelType.Kind() == reflect.Ptr { + modelType = modelType.Elem() + } + + if modelType.Kind() != reflect.Struct { + if modelType.PkgPath() == "" { + return "", errors.New("invalid model") + } + return "", fmt.Errorf("%s: %s.%s", "invalid model", modelType.PkgPath(), modelType.Name()) + } + + modelValue := reflect.New(modelType) + connectionModel, ok := modelValue.Interface().(ormcontract.ConnectionModel) + if !ok { + return "", nil + } + + return connectionModel.Connection(), nil +} + func observer(dest any) ormcontract.Observer { destType := reflect.TypeOf(dest) if destType.Kind() == reflect.Pointer { @@ -1046,18 +1443,3 @@ func observerEvent(event ormcontract.EventType, observer ormcontract.Observer) f return nil } - -func driver2gorm(driver string) string { - switch driver { - case "mysql": - return "mysql" - case "postgresql": - return "postgres" - case "sqlite": - return "sqlite" - case "sqlserver": - return "sqlserver" - default: - return "" - } -} diff --git a/database/gorm/query_test.go b/database/gorm/query_test.go index 3adce6123..b297a455d 100644 --- a/database/gorm/query_test.go +++ b/database/gorm/query_test.go @@ -9,330 +9,25 @@ import ( "testing" "time" - "github.com/spf13/cast" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/suite" _ "gorm.io/driver/postgres" + configmocks "github.com/goravel/framework/contracts/config/mocks" ormcontract "github.com/goravel/framework/contracts/database/orm" databasedb "github.com/goravel/framework/database/db" "github.com/goravel/framework/database/orm" "github.com/goravel/framework/support/file" ) -type contextKey int - -const testContextKey contextKey = 0 - -type User struct { - orm.Model - orm.SoftDeletes - Name string - Avatar string - Address *Address - Books []*Book - House *House `gorm:"polymorphic:Houseable"` - Phones []*Phone `gorm:"polymorphic:Phoneable"` - Roles []*Role `gorm:"many2many:role_user"` - age int -} - -func (u *User) DispatchesEvents() map[ormcontract.EventType]func(ormcontract.Event) error { - return map[ormcontract.EventType]func(ormcontract.Event) error{ - ormcontract.EventCreating: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil { - if name.(string) == "event_creating_name" { - event.SetAttribute("avatar", "event_creating_avatar") - } - if name.(string) == "event_creating_FirstOrCreate_name" { - event.SetAttribute("avatar", "event_creating_FirstOrCreate_avatar") - } - if name.(string) == "event_creating_IsDirty_name" { - if event.IsDirty("name") { - event.SetAttribute("avatar", "event_creating_IsDirty_avatar") - } - } - if name.(string) == "event_context" { - val := event.Context().Value(testContextKey) - event.SetAttribute("avatar", val.(string)) - } - if name.(string) == "event_query" { - _ = event.Query().Create(&User{Name: "event_query1"}) - } - } - - return nil - }, - ormcontract.EventCreated: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil { - if name.(string) == "event_created_name" { - event.SetAttribute("avatar", "event_created_avatar") - } - if name.(string) == "event_created_FirstOrCreate_name" { - event.SetAttribute("avatar", "event_created_FirstOrCreate_avatar") - } - } - - return nil - }, - ormcontract.EventSaving: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil { - if name.(string) == "event_saving_create_name" { - event.SetAttribute("avatar", "event_saving_create_avatar") - } - if name.(string) == "event_saving_save_name" { - event.SetAttribute("avatar", "event_saving_save_avatar") - } - if name.(string) == "event_saving_FirstOrCreate_name" { - event.SetAttribute("avatar", "event_saving_FirstOrCreate_avatar") - } - if name.(string) == "event_save_without_name" { - event.SetAttribute("avatar", "event_save_without_avatar") - } - if name.(string) == "event_save_quietly_name" { - event.SetAttribute("avatar", "event_save_quietly_avatar") - } - if name.(string) == "event_saving_IsDirty_name" { - if event.IsDirty("name") { - event.SetAttribute("avatar", "event_saving_IsDirty_avatar") - } - } - } - - avatar := event.GetAttribute("avatar") - if avatar != nil && avatar.(string) == "event_saving_single_update_avatar" { - event.SetAttribute("avatar", "event_saving_single_update_avatar1") - } - - return nil - }, - ormcontract.EventSaved: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil { - if name.(string) == "event_saved_create_name" { - event.SetAttribute("avatar", "event_saved_create_avatar") - } - if name.(string) == "event_saved_save_name" { - event.SetAttribute("avatar", "event_saved_save_avatar") - } - if name.(string) == "event_saved_FirstOrCreate_name" { - event.SetAttribute("avatar", "event_saved_FirstOrCreate_avatar") - } - if name.(string) == "event_save_without_name" { - event.SetAttribute("avatar", "event_saved_without_avatar") - } - if name.(string) == "event_save_quietly_name" { - event.SetAttribute("avatar", "event_saved_quietly_avatar") - } - } - - avatar := event.GetAttribute("avatar") - if avatar != nil && avatar.(string) == "event_saved_map_update_avatar" { - event.SetAttribute("avatar", "event_saved_map_update_avatar1") - } - - return nil - }, - ormcontract.EventUpdating: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil { - if name.(string) == "event_updating_create_name" { - event.SetAttribute("avatar", "event_updating_create_avatar") - } - if name.(string) == "event_updating_save_name" { - event.SetAttribute("avatar", "event_updating_save_avatar") - } - if name.(string) == "event_updating_single_update_IsDirty_name1" { - if event.IsDirty("name") { - name := event.GetAttribute("name") - if name != "event_updating_single_update_IsDirty_name1" { - return errors.New("error") - } - - event.SetAttribute("avatar", "event_updating_single_update_IsDirty_avatar") - } - } - if name.(string) == "event_updating_map_update_IsDirty_name1" { - if event.IsDirty("name") { - name := event.GetAttribute("name") - if name != "event_updating_map_update_IsDirty_name1" { - return errors.New("error") - } - - event.SetAttribute("avatar", "event_updating_map_update_IsDirty_avatar") - } - } - if name.(string) == "event_updating_model_update_IsDirty_name1" { - if event.IsDirty("name") { - name := event.GetAttribute("name") - if name != "event_updating_model_update_IsDirty_name1" { - return errors.New("error") - } - event.SetAttribute("avatar", "event_updating_model_update_IsDirty_avatar") - } - } - } - - avatar := event.GetAttribute("avatar") - if avatar != nil { - if avatar.(string) == "event_updating_save_avatar" { - event.SetAttribute("avatar", "event_updating_save_avatar1") - } - if avatar.(string) == "event_updating_model_update_avatar" { - event.SetAttribute("avatar", "event_updating_model_update_avatar1") - } - } - - return nil - }, - ormcontract.EventUpdated: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil { - if name.(string) == "event_updated_create_name" { - event.SetAttribute("avatar", "event_updated_create_avatar") - } - if name.(string) == "event_updated_save_name" { - event.SetAttribute("avatar", "event_updated_save_avatar") - } - } - - avatar := event.GetAttribute("avatar") - if avatar != nil { - if avatar.(string) == "event_updated_save_avatar" { - event.SetAttribute("avatar", "event_updated_save_avatar1") - } - if avatar.(string) == "event_updated_model_update_avatar" { - event.SetAttribute("avatar", "event_updated_model_update_avatar1") - } - } - - return nil - }, - ormcontract.EventDeleting: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil && name.(string) == "event_deleting_name" { - return errors.New("deleting error") - } - - return nil - }, - ormcontract.EventDeleted: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil && name.(string) == "event_deleted_name" { - return errors.New("deleted error") - } - - return nil - }, - ormcontract.EventForceDeleting: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil && name.(string) == "event_force_deleting_name" { - return errors.New("force deleting error") - } - - return nil - }, - ormcontract.EventForceDeleted: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil && name.(string) == "event_force_deleted_name" { - return errors.New("force deleted error") - } - - return nil - }, - ormcontract.EventRetrieved: func(event ormcontract.Event) error { - name := event.GetAttribute("name") - if name != nil && name.(string) == "event_retrieved_name" { - event.SetAttribute("name", "event_retrieved_name1") - } - - return nil - }, - } -} - -type Role struct { - orm.Model - Name string - Users []*User `gorm:"many2many:role_user"` -} - -type Address struct { - orm.Model - UserID uint - Name string - Province string - User *User -} - -type Book struct { - orm.Model - UserID uint - Name string - User *User - Author *Author -} - -type Author struct { - orm.Model - BookID uint - Name string -} - -type House struct { - orm.Model - Name string - HouseableID uint - HouseableType string -} - -func (h *House) Factory() string { - return "house" -} - -type Phone struct { - orm.Model - Name string - PhoneableID uint - PhoneableType string -} - -type Product struct { - orm.Model - orm.SoftDeletes - Name string -} - -func (p *Product) Connection() string { - return "postgresql" -} - -type Review struct { - orm.Model - orm.SoftDeletes - Body string -} - -func (r *Review) Connection() string { - return "" -} - -type Person struct { - orm.Model - orm.SoftDeletes - Name string -} - -func (p *Person) Connection() string { - return "dummy" -} - type QueryTestSuite struct { suite.Suite - queries map[ormcontract.Driver]ormcontract.Query + queries map[ormcontract.Driver]ormcontract.Query + mysqlDocker *MysqlDocker + mysqlDocker1 *MysqlDocker + postgresqlDocker *PostgresqlDocker + sqliteDocker *SqliteDocker + sqlserverDocker *SqlserverDocker } func TestQueryTestSuite(t *testing.T) { @@ -349,6 +44,12 @@ func TestQueryTestSuite(t *testing.T) { log.Fatalf("Init mysql error: %s", err) } + mysqlDocker1 := NewMysqlDocker() + mysqlPool1, mysqlResource1, _, err := mysqlDocker1.New() + if err != nil { + log.Fatalf("Init mysql1 error: %s", err) + } + postgresqlDocker := NewPostgresqlDocker() postgresqlPool, postgresqlResource, postgresqlQuery, err := postgresqlDocker.New() if err != nil { @@ -374,10 +75,16 @@ func TestQueryTestSuite(t *testing.T) { ormcontract.DriverSqlite: sqliteQuery, ormcontract.DriverSqlserver: sqlserverQuery, }, + mysqlDocker: mysqlDocker, + mysqlDocker1: mysqlDocker1, + postgresqlDocker: postgresqlDocker, + sqliteDocker: sqliteDocker, + sqlserverDocker: sqlserverDocker, }) assert.Nil(t, file.Remove(dbDatabase)) assert.Nil(t, mysqlPool.Purge(mysqlResource)) + assert.Nil(t, mysqlPool1.Purge(mysqlResource1)) assert.Nil(t, postgresqlPool.Purge(postgresqlResource)) assert.Nil(t, sqlserverPool.Purge(sqlserverResource)) } @@ -393,7 +100,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "Find", setup: func() { - user := &User{ + user := User{ Name: "association_find_name", Address: &Address{ Name: "association_find_address", @@ -418,7 +125,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "hasOne Append", setup: func() { - user := &User{ + user := User{ Name: "association_has_one_append_name", Address: &Address{ Name: "association_has_one_append_address", @@ -442,7 +149,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "hasMany Append", setup: func() { - user := &User{ + user := User{ Name: "association_has_many_append_name", Books: []*Book{ {Name: "association_has_many_append_address1"}, @@ -468,7 +175,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "hasOne Replace", setup: func() { - user := &User{ + user := User{ Name: "association_has_one_append_name", Address: &Address{ Name: "association_has_one_append_address", @@ -492,7 +199,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "hasMany Replace", setup: func() { - user := &User{ + user := User{ Name: "association_has_many_replace_name", Books: []*Book{ {Name: "association_has_many_replace_address1"}, @@ -518,7 +225,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "Delete", setup: func() { - user := &User{ + user := User{ Name: "association_delete_name", Address: &Address{ Name: "association_delete_address", @@ -554,7 +261,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "Clear", setup: func() { - user := &User{ + user := User{ Name: "association_clear_name", Address: &Address{ Name: "association_clear_address", @@ -578,7 +285,7 @@ func (s *QueryTestSuite) TestAssociation() { { name: "Count", setup: func() { - user := &User{ + user := User{ Name: "association_count_name", Books: []*Book{ {Name: "association_count_address1"}, @@ -652,11 +359,25 @@ func (s *QueryTestSuite) TestCount() { } func (s *QueryTestSuite) TestCreate() { - for _, query := range s.queries { + for driver, query := range s.queries { tests := []struct { name string setup func() }{ + { + name: "success when refresh connection", + setup: func() { + s.mockDummyConnection(driver) + + people := People{Body: "create_people"} + s.Nil(query.Create(&people)) + s.True(people.ID > 0) + + people1 := People{Body: "create_people1"} + s.Nil(query.Model(&People{}).Create(&people1)) + s.True(people1.ID > 0) + }, + }, { name: "success when create with no relationships", setup: func() { @@ -769,8 +490,10 @@ func (s *QueryTestSuite) TestCreate() { func (s *QueryTestSuite) TestCursor() { for driver, query := range s.queries { s.Run(driver.String(), func() { - user := User{Name: "cursor_user", Avatar: "cursor_avatar"} - s.Nil(query.Create(&user)) + user := User{Name: "cursor_user", Avatar: "cursor_avatar", Address: &Address{Name: "cursor_address"}, Books: []*Book{ + {Name: "cursor_book"}, + }} + s.Nil(query.Select(orm.Associations).Create(&user)) s.True(user.ID > 0) user1 := User{Name: "cursor_user", Avatar: "cursor_avatar1"} @@ -784,9 +507,11 @@ func (s *QueryTestSuite) TestCursor() { s.Nil(err) s.Equal(int64(1), res.RowsAffected) - users, err := query.Model(&User{}).Where("name = ?", "cursor_user").WithTrashed().Cursor() + users, err := query.Model(&User{}).Where("name = ?", "cursor_user").WithTrashed().With("Address").With("Books").Cursor() s.Nil(err) var size int + var addressNum int + var bookNum int for row := range users { var tempUser User s.Nil(row.Scan(&tempUser)) @@ -796,14 +521,48 @@ func (s *QueryTestSuite) TestCursor() { s.NotEmpty(tempUser.UpdatedAt.String()) s.Equal(tempUser.DeletedAt.Valid, tempUser.ID == user2.ID) size++ + + if tempUser.Address != nil { + addressNum++ + } + bookNum += len(tempUser.Books) } s.Equal(3, size) + s.Equal(1, addressNum) + s.Equal(1, bookNum) + }) + } +} + +func (s *QueryTestSuite) TestDBRaw() { + userName := "db_raw" + for driver, query := range s.queries { + s.Run(driver.String(), func() { + user := User{Name: userName} + + s.Nil(query.Create(&user)) + s.True(user.ID > 0) + switch driver { + case ormcontract.DriverSqlserver, ormcontract.DriverMysql: + res, err := query.Model(&user).Update("Name", databasedb.Raw("concat(name, ?)", driver.String())) + s.Nil(err) + s.Equal(int64(1), res.RowsAffected) + default: + res, err := query.Model(&user).Update("Name", databasedb.Raw("name || ?", driver.String())) + s.Nil(err) + s.Equal(int64(1), res.RowsAffected) + } + + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.True(user1.ID > 0) + s.True(user1.Name == userName+driver.String()) }) } } func (s *QueryTestSuite) TestDelete() { - for _, query := range s.queries { + for driver, query := range s.queries { tests := []struct { name string setup func() @@ -824,6 +583,37 @@ func (s *QueryTestSuite) TestDelete() { s.Equal(uint(0), user1.ID) }, }, + { + name: "success when refresh connection", + setup: func() { + user := User{Name: "delete_user", Avatar: "delete_avatar"} + s.Nil(query.Create(&user)) + s.True(user.ID > 0) + + res, err := query.Delete(&user) + s.Equal(int64(1), res.RowsAffected) + s.Nil(err) + + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.Equal(uint(0), user1.ID) + + // refresh connection + s.mockDummyConnection(driver) + + people := People{Body: "delete_people"} + s.Nil(query.Create(&people)) + s.True(people.ID > 0) + + res, err = query.Delete(&people) + s.Equal(int64(1), res.RowsAffected) + s.Nil(err) + + var people1 People + s.Nil(query.Find(&people1, people.ID)) + s.Equal(uint(0), people1.ID) + }, + }, { name: "success by id", setup: func() { @@ -1510,6 +1300,7 @@ func (s *QueryTestSuite) TestEvent_Query() { user := User{Name: "event_query"} s.Nil(query.Create(&user)) s.True(user.ID > 0) + s.Equal("event_query", user.Name) var user1 User s.Nil(query.Where("name", "event_query1").Find(&user1)) @@ -1542,36 +1333,21 @@ func (s *QueryTestSuite) TestExec() { func (s *QueryTestSuite) TestFind() { for _, query := range s.queries { - tests := []struct { - name string - setup func() - }{ - { - name: "success", - setup: func() { - user := User{Name: "find_user"} - s.Nil(query.Create(&user)) - s.True(user.ID > 0) + user := User{Name: "find_user"} + s.Nil(query.Create(&user)) + s.True(user.ID > 0) - var user2 User - s.Nil(query.Find(&user2, user.ID)) - s.True(user2.ID > 0) + var user2 User + s.Nil(query.Find(&user2, user.ID)) + s.True(user2.ID > 0) - var user3 []User - s.Nil(query.Find(&user3, []uint{user.ID})) - s.Equal(1, len(user3)) + var user3 []User + s.Nil(query.Find(&user3, []uint{user.ID})) + s.Equal(1, len(user3)) - var user4 []User - s.Nil(query.Where("id in ?", []uint{user.ID}).Find(&user4)) - s.Equal(1, len(user4)) - }, - }, - } - for _, test := range tests { - s.Run(test.name, func() { - test.setup() - }) - } + var user4 []User + s.Nil(query.Where("id in ?", []uint{user.ID}).Find(&user4)) + s.Equal(1, len(user4)) } } @@ -1610,29 +1386,25 @@ func (s *QueryTestSuite) TestFindOrFail() { } func (s *QueryTestSuite) TestFirst() { - for _, query := range s.queries { - tests := []struct { - name string - setup func() - }{ - { - name: "success", - setup: func() { - user := User{Name: "first_user"} - s.Nil(query.Create(&user)) - s.True(user.ID > 0) + for driver, query := range s.queries { + user := User{Name: "first_user"} + s.Nil(query.Create(&user)) + s.True(user.ID > 0) - var user1 User - s.Nil(query.Where("name", "first_user").First(&user1)) - s.True(user1.ID > 0) - }, - }, - } - for _, test := range tests { - s.Run(test.name, func() { - test.setup() - }) - } + var user1 User + s.Nil(query.Where("name", "first_user").First(&user1)) + s.True(user1.ID > 0) + + // refresh connection + s.mockDummyConnection(driver) + + people := People{Body: "first_people"} + s.Nil(query.Create(&people)) + s.True(people.ID > 0) + + var people1 People + s.Nil(query.Where("id in ?", []uint{people.ID}).First(&people1)) + s.True(people1.ID > 0) } } @@ -1842,7 +1614,23 @@ func (s *QueryTestSuite) TestGet() { var user1 []User s.Nil(query.Where("id in ?", []uint{user.ID}).Get(&user1)) s.Equal(1, len(user1)) + + // refresh connection + s.mockDummyConnection(driver) + + people := People{Body: "get_people"} + s.Nil(query.Create(&people)) + s.True(people.ID > 0) + + var people1 []People + s.Nil(query.Where("id in ?", []uint{people.ID}).Get(&people1)) + s.Equal(1, len(people1)) + + var user2 []User + s.Nil(query.Where("id in ?", []uint{user.ID}).Get(&user2)) + s.Equal(1, len(user2)) }) + break } } @@ -1962,21 +1750,26 @@ func (s *QueryTestSuite) TestPaginate() { s.True(user3.ID > 0) var users []User - var total int64 s.Nil(query.Where("name = ?", "paginate_user").Paginate(1, 3, &users, nil)) s.Equal(3, len(users)) - s.Nil(query.Where("name = ?", "paginate_user").Paginate(2, 3, &users, &total)) - s.Equal(1, len(users)) - s.Equal(int64(4), total) - - s.Nil(query.Model(User{}).Where("name = ?", "paginate_user").Paginate(1, 3, &users, &total)) - s.Equal(3, len(users)) - s.Equal(int64(4), total) - - s.Nil(query.Table("users").Where("name = ?", "paginate_user").Paginate(1, 3, &users, &total)) - s.Equal(3, len(users)) - s.Equal(int64(4), total) + var users1 []User + var total1 int64 + s.Nil(query.Where("name = ?", "paginate_user").Paginate(2, 3, &users1, &total1)) + s.Equal(1, len(users1)) + s.Equal(int64(4), total1) + + var users2 []User + var total2 int64 + s.Nil(query.Model(User{}).Where("name = ?", "paginate_user").Paginate(1, 3, &users2, &total2)) + s.Equal(3, len(users2)) + s.Equal(int64(4), total2) + + var users3 []User + var total3 int64 + s.Nil(query.Table("users").Where("name = ?", "paginate_user").Paginate(1, 3, &users3, &total3)) + s.Equal(3, len(users3)) + s.Equal(int64(4), total3) }) } } @@ -2154,97 +1947,97 @@ func (s *QueryTestSuite) TestLimit() { } func (s *QueryTestSuite) TestLoad() { - for driver, query := range s.queries { - s.Run(driver.String(), func() { - user := User{Name: "load_user", Address: &Address{}, Books: []*Book{&Book{}, &Book{}}} - user.Address.Name = "load_address" - user.Books[0].Name = "load_book0" - user.Books[1].Name = "load_book1" - s.Nil(query.Select(orm.Associations).Create(&user)) - s.True(user.ID > 0) - s.True(user.Address.ID > 0) - s.True(user.Books[0].ID > 0) - s.True(user.Books[1].ID > 0) + for _, query := range s.queries { + user := User{Name: "load_user", Address: &Address{}, Books: []*Book{&Book{}, &Book{}}} + user.Address.Name = "load_address" + user.Books[0].Name = "load_book0" + user.Books[1].Name = "load_book1" + s.Nil(query.Select(orm.Associations).Create(&user)) + s.True(user.ID > 0) + s.True(user.Address.ID > 0) + s.True(user.Books[0].ID > 0) + s.True(user.Books[1].ID > 0) - tests := []struct { - description string - setup func(description string) - }{ - { - description: "simple load relationship", - setup: func(description string) { - var user1 User - s.Nil(query.Find(&user1, user.ID)) - s.True(user1.ID > 0) - s.Nil(user1.Address) - s.True(len(user1.Books) == 0) - s.Nil(query.Load(&user1, "Address")) - s.True(user1.Address.ID > 0) - s.True(len(user1.Books) == 0) - s.Nil(query.Load(&user1, "Books")) - s.True(user1.Address.ID > 0) - s.True(len(user1.Books) == 2) - }, + tests := []struct { + description string + setup func(description string) + }{ + { + description: "simple load relationship", + setup: func(description string) { + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.True(user1.ID > 0) + s.Nil(user1.Address) + s.True(len(user1.Books) == 0) + s.Nil(query.Load(&user1, "Address")) + s.True(user1.Address.ID > 0) + s.True(len(user1.Books) == 0) + s.Nil(query.Load(&user1, "Books")) + s.True(user1.Address.ID > 0) + s.True(len(user1.Books) == 2) }, - { - description: "load relationship with simple condition", - setup: func(description string) { - var user1 User - s.Nil(query.Find(&user1, user.ID)) - s.True(user1.ID > 0) - s.Nil(user1.Address) - s.Equal(0, len(user1.Books)) - s.Nil(query.Load(&user1, "Books", "name = ?", "load_book0")) - s.True(user1.ID > 0) - s.Nil(user1.Address) - s.Equal(1, len(user1.Books)) - s.Equal("load_book0", user.Books[0].Name) - }, + }, + { + description: "load relationship with simple condition", + setup: func(description string) { + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.True(user1.ID > 0) + s.Nil(user1.Address) + s.Equal(0, len(user1.Books)) + s.Nil(query.Load(&user1, "Books", "name = ?", "load_book0")) + s.True(user1.ID > 0) + s.Nil(user1.Address) + s.Equal(1, len(user1.Books)) + s.Equal("load_book0", user.Books[0].Name) }, - { - description: "load relationship with func condition", - setup: func(description string) { - var user1 User - s.Nil(query.Find(&user1, user.ID)) - s.True(user1.ID > 0) - s.Nil(user1.Address) - s.Equal(0, len(user1.Books)) - s.Nil(query.Load(&user1, "Books", func(query ormcontract.Query) ormcontract.Query { - return query.Where("name = ?", "load_book0") - })) - s.True(user1.ID > 0) - s.Nil(user1.Address) - s.Equal(1, len(user1.Books)) - s.Equal("load_book0", user.Books[0].Name) - }, + }, + { + description: "load relationship with func condition", + setup: func(description string) { + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.True(user1.ID > 0) + s.Nil(user1.Address) + s.Equal(0, len(user1.Books)) + s.Nil(query.Load(&user1, "Books", func(query ormcontract.Query) ormcontract.Query { + return query.Where("name = ?", "load_book0") + })) + s.True(user1.ID > 0) + s.Nil(user1.Address) + s.Equal(1, len(user1.Books)) + s.Equal("load_book0", user.Books[0].Name) }, - { - description: "error when relation is empty", - setup: func(description string) { - var user1 User - s.Nil(query.Find(&user1, user.ID)) - s.True(user1.ID > 0) - s.Nil(user1.Address) - s.Equal(0, len(user1.Books)) - s.EqualError(query.Load(&user1, ""), "relation cannot be empty") - }, + }, + { + description: "error when relation is empty", + setup: func(description string) { + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.True(user1.ID > 0) + s.Nil(user1.Address) + s.Equal(0, len(user1.Books)) + s.EqualError(query.Load(&user1, ""), "relation cannot be empty") }, - { - description: "error when id is nil", - setup: func(description string) { - type UserNoID struct { - Name string - Avatar string - } - var userNoID UserNoID - s.EqualError(query.Load(&userNoID, "Book"), "id cannot be empty") - }, + }, + { + description: "error when id is nil", + setup: func(description string) { + type UserNoID struct { + Name string + Avatar string + } + var userNoID UserNoID + s.EqualError(query.Load(&userNoID, "Book"), "id cannot be empty") }, - } - for _, test := range tests { + }, + } + for _, test := range tests { + s.Run(test.description, func() { test.setup(test.description) - } - }) + }) + } } } @@ -2319,6 +2112,107 @@ func (s *QueryTestSuite) TestRaw() { } } +func (s *QueryTestSuite) TestReuse() { + for _, query := range s.queries { + users := []User{{Name: "reuse_user", Avatar: "reuse_avatar"}, {Name: "reuse_user1", Avatar: "reuse_avatar1"}} + s.Nil(query.Create(&users)) + s.True(users[0].ID > 0) + s.True(users[1].ID > 0) + + q := query.Where("name", "reuse_user") + + var users1 User + s.Nil(q.Where("avatar", "reuse_avatar").Find(&users1)) + s.True(users1.ID > 0) + + var users2 User + s.Nil(q.Where("avatar", "reuse_avatar1").Find(&users2)) + s.True(users2.ID == 0) + + var users3 User + s.Nil(query.Where("avatar", "reuse_avatar1").Find(&users3)) + s.True(users3.ID > 0) + } +} + +func (s *QueryTestSuite) TestRefreshConnection() { + tests := []struct { + name string + model any + setup func() + expectConnection string + expectErr string + }{ + { + name: "invalid model", + model: func() any { + var product string + return product + }(), + setup: func() {}, + expectErr: "invalid model", + }, + { + name: "the connection of model is empty", + model: func() any { + var review Review + return review + }(), + setup: func() {}, + expectConnection: "mysql", + }, + { + name: "the connection of model is same as current connection", + model: func() any { + var box Box + return box + }(), + setup: func() {}, + expectConnection: "mysql", + }, + { + name: "connections are different, but drivers are same", + model: func() any { + var people People + return people + }(), + setup: func() { + mockDummyConnection(s.mysqlDocker.MockConfig, s.mysqlDocker1.Port) + }, + expectConnection: "dummy", + }, + { + name: "connections and drivers are different", + model: func() any { + var product Product + return product + }(), + setup: func() { + mockPostgresqlConnection(s.mysqlDocker.MockConfig, s.postgresqlDocker.Port) + }, + expectConnection: "postgresql", + }, + } + + for _, test := range tests { + s.Run(test.name, func() { + test.setup() + queryImpl := s.queries[ormcontract.DriverMysql].(*QueryImpl) + query, err := queryImpl.refreshConnection(test.model) + if test.expectErr != "" { + s.EqualError(err, test.expectErr) + } else { + s.Nil(err) + } + if test.expectConnection == "" { + s.Nil(query) + } else { + s.Equal(test.expectConnection, query.connection) + } + }) + } +} + func (s *QueryTestSuite) TestSave() { for _, query := range s.queries { tests := []struct { @@ -2363,31 +2257,16 @@ func (s *QueryTestSuite) TestSave() { func (s *QueryTestSuite) TestSaveQuietly() { for _, query := range s.queries { - tests := []struct { - name string - setup func() - }{ - { - name: "success", - setup: func() { - user := User{Name: "event_save_quietly_name", Avatar: "save_quietly_avatar"} - s.Nil(query.SaveQuietly(&user)) - s.True(user.ID > 0) - s.Equal("event_save_quietly_name", user.Name) - s.Equal("save_quietly_avatar", user.Avatar) + user := User{Name: "event_save_quietly_name", Avatar: "save_quietly_avatar"} + s.Nil(query.SaveQuietly(&user)) + s.True(user.ID > 0) + s.Equal("event_save_quietly_name", user.Name) + s.Equal("save_quietly_avatar", user.Avatar) - var user1 User - s.Nil(query.Find(&user1, user.ID)) - s.Equal("event_save_quietly_name", user1.Name) - s.Equal("save_quietly_avatar", user1.Avatar) - }, - }, - } - for _, test := range tests { - s.Run(test.name, func() { - test.setup() - }) - } + var user1 User + s.Nil(query.Find(&user1, user.ID)) + s.Equal("event_save_quietly_name", user1.Name) + s.Equal("save_quietly_avatar", user1.Avatar) } } @@ -2683,7 +2562,7 @@ func (s *QueryTestSuite) TestWhere() { var user2 []User s.Nil(query.Where("name = ?", "where_user").OrWhere("avatar = ?", "where_avatar1").Find(&user2)) - s.True(len(user2) > 0) + s.Equal(2, len(user2)) var user3 User s.Nil(query.Where("name = 'where_user'").Find(&user3)) @@ -2821,30 +2700,16 @@ func (s *QueryTestSuite) TestWithNesting() { } } -func (s *QueryTestSuite) TestDBRaw() { - userName := "db_raw" - for driver, query := range s.queries { - s.Run(driver.String(), func() { - user := User{Name: userName} - - s.Nil(query.Create(&user)) - s.True(user.ID > 0) - switch driver { - case ormcontract.DriverSqlserver, ormcontract.DriverMysql: - res, err := query.Model(&user).Update("Name", databasedb.Raw("concat(name, ?)", driver.String())) - s.Nil(err) - s.Equal(int64(1), res.RowsAffected) - default: - res, err := query.Model(&user).Update("Name", databasedb.Raw("name || ?", driver.String())) - s.Nil(err) - s.Equal(int64(1), res.RowsAffected) - } - - var user1 User - s.Nil(query.Find(&user1, user.ID)) - s.True(user1.ID > 0) - s.True(user1.Name == userName+driver.String()) - }) +func (s *QueryTestSuite) mockDummyConnection(driver ormcontract.Driver) { + switch driver { + case ormcontract.DriverMysql: + mockDummyConnection(s.mysqlDocker.MockConfig, s.mysqlDocker1.Port) + case ormcontract.DriverPostgresql: + mockDummyConnection(s.postgresqlDocker.MockConfig, s.mysqlDocker1.Port) + case ormcontract.DriverSqlite: + mockDummyConnection(s.sqliteDocker.MockConfig, s.mysqlDocker1.Port) + case ormcontract.DriverSqlserver: + mockDummyConnection(s.sqlserverDocker.MockConfig, s.mysqlDocker1.Port) } } @@ -2872,18 +2737,7 @@ func TestCustomConnection(t *testing.T) { assert.Nil(t, query.Where("body", "create_review").First(&review1)) assert.True(t, review1.ID > 0) - mysqlDocker.MockConfig.On("Get", "database.connections.postgresql.read").Return(nil) - mysqlDocker.MockConfig.On("Get", "database.connections.postgresql.write").Return(nil) - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.host").Return("localhost") - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.username").Return(DbUser) - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.password").Return(DbPassword) - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.driver").Return(ormcontract.DriverPostgresql.String()) - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.database").Return("postgres") - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.sslmode").Return("disable") - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.timezone").Return("UTC") - mysqlDocker.MockConfig.On("GetString", "database.connections.postgresql.prefix").Return("") - mysqlDocker.MockConfig.On("GetBool", "database.connections.postgresql.singular").Return(false) - mysqlDocker.MockConfig.On("GetInt", "database.connections.postgresql.port").Return(cast.ToInt(postgresqlResource.GetPort("5432/tcp"))) + mockPostgresqlConnection(mysqlDocker.MockConfig, postgresqlDocker.Port) product := Product{Name: "create_product"} assert.Nil(t, query.Create(&product)) @@ -2897,7 +2751,7 @@ func TestCustomConnection(t *testing.T) { assert.Nil(t, query.Where("name", "create_product1").First(&product2)) assert.True(t, product2.ID == 0) - mysqlDocker.MockConfig.On("GetString", "database.connections.dummy.driver").Return("") + mockDummyConnection(mysqlDocker.MockConfig, mysqlDocker.Port) person := Person{Name: "create_person"} assert.NotNil(t, query.Create(&person)) @@ -2907,6 +2761,128 @@ func TestCustomConnection(t *testing.T) { assert.Nil(t, postgresqlPool.Purge(postgresqlResource)) } +func TestFilterFindConditions(t *testing.T) { + tests := []struct { + name string + conditions []any + expectErr error + }{ + { + name: "condition is empty", + }, + { + name: "condition is empty string", + conditions: []any{""}, + expectErr: ErrorMissingWhereClause, + }, + { + name: "condition is empty slice", + conditions: []any{[]string{}}, + expectErr: ErrorMissingWhereClause, + }, + { + name: "condition has value", + conditions: []any{"name = ?", "test"}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + err := filterFindConditions(test.conditions...) + if test.expectErr != nil { + assert.Equal(t, err, test.expectErr) + } else { + assert.Nil(t, err) + } + }) + } +} + +func TestGetModelConnection(t *testing.T) { + tests := []struct { + name string + model any + expectErr string + expectConnection string + }{ + { + name: "invalid model", + model: func() any { + var product string + return product + }(), + expectErr: "invalid model", + }, + { + name: "not ConnectionModel", + model: func() any { + var phone Phone + return phone + }(), + }, + { + name: "the connection of model is empty", + model: func() any { + var review Review + return review + }(), + }, + { + name: "the connection of model is not empty", + model: func() any { + var product Product + return product + }(), + expectConnection: "postgresql", + }, + { + name: "the connection of model is not empty and model is slice", + model: func() any { + var products []Product + return products + }(), + expectConnection: "postgresql", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + connection, err := getModelConnection(test.model) + if test.expectErr != "" { + assert.EqualError(t, err, test.expectErr) + } else { + assert.Nil(t, err) + } + assert.Equal(t, test.expectConnection, connection) + }) + } +} + +func TestObserver(t *testing.T) { + orm.Observers = append(orm.Observers, orm.Observer{ + Model: User{}, + Observer: &UserObserver{}, + }) + + assert.Nil(t, observer(Product{})) + assert.Equal(t, &UserObserver{}, observer(User{})) +} + +func TestObserverEvent(t *testing.T) { + assert.EqualError(t, observerEvent(ormcontract.EventRetrieved, &UserObserver{})(nil), "retrieved") + assert.EqualError(t, observerEvent(ormcontract.EventCreating, &UserObserver{})(nil), "creating") + assert.EqualError(t, observerEvent(ormcontract.EventCreated, &UserObserver{})(nil), "created") + assert.EqualError(t, observerEvent(ormcontract.EventUpdating, &UserObserver{})(nil), "updating") + assert.EqualError(t, observerEvent(ormcontract.EventUpdated, &UserObserver{})(nil), "updated") + assert.EqualError(t, observerEvent(ormcontract.EventSaving, &UserObserver{})(nil), "saving") + assert.EqualError(t, observerEvent(ormcontract.EventSaved, &UserObserver{})(nil), "saved") + assert.EqualError(t, observerEvent(ormcontract.EventDeleting, &UserObserver{})(nil), "deleting") + assert.EqualError(t, observerEvent(ormcontract.EventDeleted, &UserObserver{})(nil), "deleted") + assert.EqualError(t, observerEvent(ormcontract.EventForceDeleting, &UserObserver{})(nil), "forceDeleting") + assert.EqualError(t, observerEvent(ormcontract.EventForceDeleted, &UserObserver{})(nil), "forceDeleted") + assert.Nil(t, observerEvent("error", &UserObserver{})) +} + func TestReadWriteSeparate(t *testing.T) { if testing.Short() { t.Skip("Skipping tests of using docker") @@ -3139,3 +3115,79 @@ func paginator(page string, limit string) func(methods ormcontract.Query) ormcon return query.Offset(offset).Limit(limit) } } + +func mockDummyConnection(mockConfig *configmocks.Config, port int) { + mockConfig.On("GetString", "database.connections.dummy.prefix").Return("") + mockConfig.On("GetBool", "database.connections.dummy.singular").Return(false) + mockConfig.On("Get", "database.connections.dummy.read").Return(nil) + mockConfig.On("Get", "database.connections.dummy.write").Return(nil) + mockConfig.On("GetString", "database.connections.dummy.host").Return("127.0.0.1") + mockConfig.On("GetString", "database.connections.dummy.username").Return(DbUser) + mockConfig.On("GetString", "database.connections.dummy.password").Return(DbPassword) + mockConfig.On("GetInt", "database.connections.dummy.port").Return(port) + mockConfig.On("GetString", "database.connections.dummy.driver").Return(ormcontract.DriverMysql.String()) + mockConfig.On("GetString", "database.connections.dummy.charset").Return("utf8mb4") + mockConfig.On("GetString", "database.connections.dummy.loc").Return("Local") + mockConfig.On("GetString", "database.connections.dummy.database").Return(dbDatabase) +} + +func mockPostgresqlConnection(mockConfig *configmocks.Config, port int) { + mockConfig.On("GetString", "database.connections.postgresql.prefix").Return("") + mockConfig.On("GetBool", "database.connections.postgresql.singular").Return(false) + mockConfig.On("Get", "database.connections.postgresql.read").Return(nil) + mockConfig.On("Get", "database.connections.postgresql.write").Return(nil) + mockConfig.On("GetString", "database.connections.postgresql.host").Return("127.0.0.1") + mockConfig.On("GetString", "database.connections.postgresql.username").Return(DbUser) + mockConfig.On("GetString", "database.connections.postgresql.password").Return(DbPassword) + mockConfig.On("GetInt", "database.connections.postgresql.port").Return(port) + mockConfig.On("GetString", "database.connections.postgresql.driver").Return(ormcontract.DriverPostgresql.String()) + mockConfig.On("GetString", "database.connections.postgresql.sslmode").Return("disable") + mockConfig.On("GetString", "database.connections.postgresql.timezone").Return("UTC") + mockConfig.On("GetString", "database.connections.postgresql.database").Return("postgres") +} + +type UserObserver struct{} + +func (u *UserObserver) Retrieved(event ormcontract.Event) error { + return errors.New("retrieved") +} + +func (u *UserObserver) Creating(event ormcontract.Event) error { + return errors.New("creating") +} + +func (u *UserObserver) Created(event ormcontract.Event) error { + return errors.New("created") +} + +func (u *UserObserver) Updating(event ormcontract.Event) error { + return errors.New("updating") +} + +func (u *UserObserver) Updated(event ormcontract.Event) error { + return errors.New("updated") +} + +func (u *UserObserver) Saving(event ormcontract.Event) error { + return errors.New("saving") +} + +func (u *UserObserver) Saved(event ormcontract.Event) error { + return errors.New("saved") +} + +func (u *UserObserver) Deleting(event ormcontract.Event) error { + return errors.New("deleting") +} + +func (u *UserObserver) Deleted(event ormcontract.Event) error { + return errors.New("deleted") +} + +func (u *UserObserver) ForceDeleting(event ormcontract.Event) error { + return errors.New("forceDeleting") +} + +func (u *UserObserver) ForceDeleted(event ormcontract.Event) error { + return errors.New("forceDeleted") +} diff --git a/database/gorm/test_models.go b/database/gorm/test_models.go new file mode 100644 index 000000000..4324737d0 --- /dev/null +++ b/database/gorm/test_models.go @@ -0,0 +1,338 @@ +package gorm + +import ( + "errors" + + ormcontract "github.com/goravel/framework/contracts/database/orm" + "github.com/goravel/framework/database/orm" +) + +type contextKey int + +const testContextKey contextKey = 0 + +type User struct { + orm.Model + orm.SoftDeletes + Name string + Avatar string + Address *Address + Books []*Book + House *House `gorm:"polymorphic:Houseable"` + Phones []*Phone `gorm:"polymorphic:Phoneable"` + Roles []*Role `gorm:"many2many:role_user"` + age int +} + +func (u *User) DispatchesEvents() map[ormcontract.EventType]func(ormcontract.Event) error { + return map[ormcontract.EventType]func(ormcontract.Event) error{ + ormcontract.EventCreating: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil { + if name.(string) == "event_creating_name" { + event.SetAttribute("avatar", "event_creating_avatar") + } + if name.(string) == "event_creating_FirstOrCreate_name" { + event.SetAttribute("avatar", "event_creating_FirstOrCreate_avatar") + } + if name.(string) == "event_creating_IsDirty_name" { + if event.IsDirty("name") { + event.SetAttribute("avatar", "event_creating_IsDirty_avatar") + } + } + if name.(string) == "event_context" { + val := event.Context().Value(testContextKey) + event.SetAttribute("avatar", val.(string)) + } + if name.(string) == "event_query" { + _ = event.Query().Create(&User{Name: "event_query1"}) + } + } + + return nil + }, + ormcontract.EventCreated: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil { + if name.(string) == "event_created_name" { + event.SetAttribute("avatar", "event_created_avatar") + } + if name.(string) == "event_created_FirstOrCreate_name" { + event.SetAttribute("avatar", "event_created_FirstOrCreate_avatar") + } + } + + return nil + }, + ormcontract.EventSaving: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil { + if name.(string) == "event_saving_create_name" { + event.SetAttribute("avatar", "event_saving_create_avatar") + } + if name.(string) == "event_saving_save_name" { + event.SetAttribute("avatar", "event_saving_save_avatar") + } + if name.(string) == "event_saving_FirstOrCreate_name" { + event.SetAttribute("avatar", "event_saving_FirstOrCreate_avatar") + } + if name.(string) == "event_save_without_name" { + event.SetAttribute("avatar", "event_save_without_avatar") + } + if name.(string) == "event_save_quietly_name" { + event.SetAttribute("avatar", "event_save_quietly_avatar") + } + if name.(string) == "event_saving_IsDirty_name" { + if event.IsDirty("name") { + event.SetAttribute("avatar", "event_saving_IsDirty_avatar") + } + } + } + + avatar := event.GetAttribute("avatar") + if avatar != nil && avatar.(string) == "event_saving_single_update_avatar" { + event.SetAttribute("avatar", "event_saving_single_update_avatar1") + } + + return nil + }, + ormcontract.EventSaved: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil { + if name.(string) == "event_saved_create_name" { + event.SetAttribute("avatar", "event_saved_create_avatar") + } + if name.(string) == "event_saved_save_name" { + event.SetAttribute("avatar", "event_saved_save_avatar") + } + if name.(string) == "event_saved_FirstOrCreate_name" { + event.SetAttribute("avatar", "event_saved_FirstOrCreate_avatar") + } + if name.(string) == "event_save_without_name" { + event.SetAttribute("avatar", "event_saved_without_avatar") + } + if name.(string) == "event_save_quietly_name" { + event.SetAttribute("avatar", "event_saved_quietly_avatar") + } + } + + avatar := event.GetAttribute("avatar") + if avatar != nil && avatar.(string) == "event_saved_map_update_avatar" { + event.SetAttribute("avatar", "event_saved_map_update_avatar1") + } + + return nil + }, + ormcontract.EventUpdating: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil { + if name.(string) == "event_updating_create_name" { + event.SetAttribute("avatar", "event_updating_create_avatar") + } + if name.(string) == "event_updating_save_name" { + event.SetAttribute("avatar", "event_updating_save_avatar") + } + if name.(string) == "event_updating_single_update_IsDirty_name1" { + if event.IsDirty("name") { + name := event.GetAttribute("name") + if name != "event_updating_single_update_IsDirty_name1" { + return errors.New("error") + } + + event.SetAttribute("avatar", "event_updating_single_update_IsDirty_avatar") + } + } + if name.(string) == "event_updating_map_update_IsDirty_name1" { + if event.IsDirty("name") { + name := event.GetAttribute("name") + if name != "event_updating_map_update_IsDirty_name1" { + return errors.New("error") + } + + event.SetAttribute("avatar", "event_updating_map_update_IsDirty_avatar") + } + } + if name.(string) == "event_updating_model_update_IsDirty_name1" { + if event.IsDirty("name") { + name := event.GetAttribute("name") + if name != "event_updating_model_update_IsDirty_name1" { + return errors.New("error") + } + event.SetAttribute("avatar", "event_updating_model_update_IsDirty_avatar") + } + } + } + + avatar := event.GetAttribute("avatar") + if avatar != nil { + if avatar.(string) == "event_updating_save_avatar" { + event.SetAttribute("avatar", "event_updating_save_avatar1") + } + if avatar.(string) == "event_updating_model_update_avatar" { + event.SetAttribute("avatar", "event_updating_model_update_avatar1") + } + } + + return nil + }, + ormcontract.EventUpdated: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil { + if name.(string) == "event_updated_create_name" { + event.SetAttribute("avatar", "event_updated_create_avatar") + } + if name.(string) == "event_updated_save_name" { + event.SetAttribute("avatar", "event_updated_save_avatar") + } + } + + avatar := event.GetAttribute("avatar") + if avatar != nil { + if avatar.(string) == "event_updated_save_avatar" { + event.SetAttribute("avatar", "event_updated_save_avatar1") + } + if avatar.(string) == "event_updated_model_update_avatar" { + event.SetAttribute("avatar", "event_updated_model_update_avatar1") + } + } + + return nil + }, + ormcontract.EventDeleting: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil && name.(string) == "event_deleting_name" { + return errors.New("deleting error") + } + + return nil + }, + ormcontract.EventDeleted: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil && name.(string) == "event_deleted_name" { + return errors.New("deleted error") + } + + return nil + }, + ormcontract.EventForceDeleting: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil && name.(string) == "event_force_deleting_name" { + return errors.New("force deleting error") + } + + return nil + }, + ormcontract.EventForceDeleted: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil && name.(string) == "event_force_deleted_name" { + return errors.New("force deleted error") + } + + return nil + }, + ormcontract.EventRetrieved: func(event ormcontract.Event) error { + name := event.GetAttribute("name") + if name != nil && name.(string) == "event_retrieved_name" { + event.SetAttribute("name", "event_retrieved_name1") + } + + return nil + }, + } +} + +type Role struct { + orm.Model + Name string + Users []*User `gorm:"many2many:role_user"` +} + +type Address struct { + orm.Model + UserID uint + Name string + Province string + User *User +} + +type Book struct { + orm.Model + UserID uint + Name string + User *User + Author *Author +} + +type Author struct { + orm.Model + BookID uint + Name string +} + +type House struct { + orm.Model + Name string + HouseableID uint + HouseableType string +} + +func (h *House) Factory() string { + return "house" +} + +type Phone struct { + orm.Model + Name string + PhoneableID uint + PhoneableType string +} + +type Product struct { + orm.Model + orm.SoftDeletes + Name string +} + +func (p *Product) Connection() string { + return "postgresql" +} + +type Review struct { + orm.Model + orm.SoftDeletes + Body string +} + +func (r *Review) Connection() string { + return "" +} + +type People struct { + orm.Model + orm.SoftDeletes + Body string +} + +func (p *People) Connection() string { + return "dummy" +} + +type Person struct { + orm.Model + orm.SoftDeletes + Name string +} + +func (p *Person) Connection() string { + return "dummy" +} + +type Box struct { + orm.Model + orm.SoftDeletes + Name string +} + +func (p *Box) Connection() string { + return "mysql" +} diff --git a/database/gorm/test_utils.go b/database/gorm/test_utils.go index 7928b292c..e6bc73466 100644 --- a/database/gorm/test_utils.go +++ b/database/gorm/test_utils.go @@ -80,7 +80,7 @@ func (r *MysqlDocker) Query(createTable bool) (orm.Query, error) { } if createTable { - err = Table{}.Create(orm.DriverMysql, db) + err = Tables{}.Create(orm.DriverMysql, db) if err != nil { return nil, err } @@ -95,7 +95,7 @@ func (r *MysqlDocker) QueryWithPrefixAndSingular() (orm.Query, error) { return nil, err } - err = Table{}.CreateWithPrefixAndSingular(orm.DriverMysql, db) + err = Tables{}.CreateWithPrefixAndSingular(orm.DriverMysql, db) if err != nil { return nil, err } @@ -227,7 +227,7 @@ func (r *PostgresqlDocker) Query(createTable bool) (orm.Query, error) { } if createTable { - err = Table{}.Create(orm.DriverPostgresql, db) + err = Tables{}.Create(orm.DriverPostgresql, db) if err != nil { return nil, err } @@ -242,7 +242,7 @@ func (r *PostgresqlDocker) QueryWithPrefixAndSingular() (orm.Query, error) { return nil, err } - err = Table{}.CreateWithPrefixAndSingular(orm.DriverPostgresql, db) + err = Tables{}.CreateWithPrefixAndSingular(orm.DriverPostgresql, db) if err != nil { return nil, err } @@ -368,7 +368,7 @@ func (r *SqliteDocker) Query(createTable bool) (orm.Query, error) { } if createTable { - err = Table{}.Create(orm.DriverSqlite, db) + err = Tables{}.Create(orm.DriverSqlite, db) if err != nil { return nil, err } @@ -383,7 +383,7 @@ func (r *SqliteDocker) QueryWithPrefixAndSingular() (orm.Query, error) { return nil, err } - err = Table{}.CreateWithPrefixAndSingular(orm.DriverSqlite, db) + err = Tables{}.CreateWithPrefixAndSingular(orm.DriverSqlite, db) if err != nil { return nil, err } @@ -503,7 +503,7 @@ func (r *SqlserverDocker) Query(createTable bool) (orm.Query, error) { } if createTable { - err = Table{}.Create(orm.DriverSqlserver, db) + err = Tables{}.Create(orm.DriverSqlserver, db) if err != nil { return nil, err } @@ -518,7 +518,7 @@ func (r *SqlserverDocker) QueryWithPrefixAndSingular() (orm.Query, error) { return nil, err } - err = Table{}.CreateWithPrefixAndSingular(orm.DriverSqlserver, db) + err = Tables{}.CreateWithPrefixAndSingular(orm.DriverSqlserver, db) if err != nil { return nil, err } @@ -589,11 +589,11 @@ func (r *SqlserverDocker) query() (orm.Query, error) { return db, nil } -type Table struct { +type Tables struct { } -func (r Table) Create(driver orm.Driver, db orm.Query) error { - _, err := db.Exec(r.createPersonTable(driver)) +func (r Tables) Create(driver orm.Driver, db orm.Query) error { + _, err := db.Exec(r.createPeopleTable(driver)) if err != nil { return err } @@ -641,7 +641,7 @@ func (r Table) Create(driver orm.Driver, db orm.Query) error { return nil } -func (r Table) CreateWithPrefixAndSingular(driver orm.Driver, db orm.Query) error { +func (r Tables) CreateWithPrefixAndSingular(driver orm.Driver, db orm.Query) error { _, err := db.Exec(r.createUserTableWithPrefixAndSingular(driver)) if err != nil { return err @@ -650,11 +650,11 @@ func (r Table) CreateWithPrefixAndSingular(driver orm.Driver, db orm.Query) erro return nil } -func (r Table) createPersonTable(driver orm.Driver) string { +func (r Tables) createPeopleTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` -CREATE TABLE people ( +CREATE TABLE peoples ( id bigint(20) unsigned NOT NULL AUTO_INCREMENT, body varchar(255) NOT NULL, created_at datetime(3) NOT NULL, @@ -667,7 +667,7 @@ CREATE TABLE people ( ` case orm.DriverPostgresql: return ` -CREATE TABLE people ( +CREATE TABLE peoples ( id SERIAL PRIMARY KEY NOT NULL, body varchar(255) NOT NULL, created_at timestamp NOT NULL, @@ -677,7 +677,7 @@ CREATE TABLE people ( ` case orm.DriverSqlite: return ` -CREATE TABLE people ( +CREATE TABLE peoples ( id integer PRIMARY KEY AUTOINCREMENT NOT NULL, body varchar(255) NOT NULL, created_at datetime NOT NULL, @@ -687,7 +687,7 @@ CREATE TABLE people ( ` case orm.DriverSqlserver: return ` -CREATE TABLE people ( +CREATE TABLE peoples ( id bigint NOT NULL IDENTITY(1,1), body varchar(255) NOT NULL, created_at datetime NOT NULL, @@ -701,7 +701,7 @@ CREATE TABLE people ( } } -func (r Table) createReviewTable(driver orm.Driver) string { +func (r Tables) createReviewTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -752,7 +752,7 @@ CREATE TABLE reviews ( } } -func (r Table) createProductTable(driver orm.Driver) string { +func (r Tables) createProductTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -803,7 +803,7 @@ CREATE TABLE products ( } } -func (r Table) createUserTable(driver orm.Driver) string { +func (r Tables) createUserTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -858,7 +858,7 @@ CREATE TABLE users ( } } -func (r Table) createUserTableWithPrefixAndSingular(driver orm.Driver) string { +func (r Tables) createUserTableWithPrefixAndSingular(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -913,7 +913,7 @@ CREATE TABLE goravel_user ( } } -func (r Table) createAddressTable(driver orm.Driver) string { +func (r Tables) createAddressTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -968,7 +968,7 @@ CREATE TABLE addresses ( } } -func (r Table) createBookTable(driver orm.Driver) string { +func (r Tables) createBookTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -1019,7 +1019,7 @@ CREATE TABLE books ( } } -func (r Table) createAuthorTable(driver orm.Driver) string { +func (r Tables) createAuthorTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -1070,7 +1070,7 @@ CREATE TABLE authors ( } } -func (r Table) createRoleTable(driver orm.Driver) string { +func (r Tables) createRoleTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -1117,7 +1117,7 @@ CREATE TABLE roles ( } } -func (r Table) createHouseTable(driver orm.Driver) string { +func (r Tables) createHouseTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -1172,7 +1172,7 @@ CREATE TABLE houses ( } } -func (r Table) createPhoneTable(driver orm.Driver) string { +func (r Tables) createPhoneTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` @@ -1227,7 +1227,7 @@ CREATE TABLE phones ( } } -func (r Table) createRoleUserTable(driver orm.Driver) string { +func (r Tables) createRoleUserTable(driver orm.Driver) string { switch driver { case orm.DriverMysql: return ` diff --git a/database/gorm/transaction.go b/database/gorm/transaction.go index e72ffe958..85ff9cf8b 100644 --- a/database/gorm/transaction.go +++ b/database/gorm/transaction.go @@ -12,8 +12,8 @@ type Transaction struct { instance *gorm.DB } -func NewTransaction(tx *gorm.DB, config config.Config) *Transaction { - return &Transaction{Query: NewQueryWithWithoutEvents(tx, false, config), instance: tx} +func NewTransaction(tx *gorm.DB, config config.Config, connection string) *Transaction { + return &Transaction{Query: NewQueryImpl(tx.Statement.Context, config, connection, tx, nil), instance: tx} } func (r *Transaction) Commit() error { diff --git a/database/gorm/utils_test.go b/database/gorm/utils_test.go index 236d8c134..e6b352b1f 100644 --- a/database/gorm/utils_test.go +++ b/database/gorm/utils_test.go @@ -4,8 +4,6 @@ import ( "testing" "github.com/stretchr/testify/assert" - - "github.com/goravel/framework/support/debug" ) func TestCopyStruct(t *testing.T) { @@ -15,7 +13,7 @@ func TestCopyStruct(t *testing.T) { } data := copyStruct(Data{Name: "name", age: 18}) - debug.Dump(data) + assert.Equal(t, "name", data.Field(0).Interface().(string)) assert.Panics(t, func() { data.Field(1).Interface() diff --git a/database/gorm/wire.go b/database/gorm/wire.go index c1688e021..c43157861 100644 --- a/database/gorm/wire.go +++ b/database/gorm/wire.go @@ -23,7 +23,7 @@ func InitializeGorm(config config.Config, connection string) *GormImpl { //go:generate wire func InitializeQuery(ctx context.Context, config config.Config, connection string) (*QueryImpl, error) { - wire.Build(NewQueryImpl, GormSet, db.ConfigSet, DialectorSet) + wire.Build(BuildQueryImpl, GormSet, db.ConfigSet, DialectorSet) return nil, nil } diff --git a/database/gorm/wire_gen.go b/database/gorm/wire_gen.go index bb5259615..7147c077a 100644 --- a/database/gorm/wire_gen.go +++ b/database/gorm/wire_gen.go @@ -27,7 +27,7 @@ func InitializeQuery(ctx context.Context, config2 config.Config, connection stri configImpl := db.NewConfigImpl(config2, connection) dialectorImpl := NewDialectorImpl(config2, connection) gormImpl := NewGormImpl(config2, connection, configImpl, dialectorImpl) - queryImpl, err := NewQueryImpl(ctx, config2, gormImpl) + queryImpl, err := BuildQueryImpl(ctx, config2, connection, gormImpl) if err != nil { return nil, err } diff --git a/database/orm.go b/database/orm.go index e81b39dc2..7cb6e96d1 100644 --- a/database/orm.go +++ b/database/orm.go @@ -40,9 +40,11 @@ func (r *OrmImpl) Connection(name string) ormcontract.Orm { } if instance, exist := r.queries[name]; exist { return &OrmImpl{ - ctx: r.ctx, - query: instance, - queries: r.queries, + ctx: r.ctx, + config: r.config, + connection: name, + query: instance, + queries: r.queries, } } @@ -56,16 +58,18 @@ func (r *OrmImpl) Connection(name string) ormcontract.Orm { r.queries[name] = queue return &OrmImpl{ - ctx: r.ctx, - query: queue, - queries: r.queries, + ctx: r.ctx, + config: r.config, + connection: name, + query: queue, + queries: r.queries, } } func (r *OrmImpl) DB() (*sql.DB, error) { - db := r.Query().(*databasegorm.QueryImpl) + query := r.Query().(*databasegorm.QueryImpl) - return db.Instance().DB() + return query.Instance().DB() } func (r *OrmImpl) Query() ormcontract.Query { @@ -101,7 +105,19 @@ func (r *OrmImpl) Transaction(txFunc func(tx ormcontract.Transaction) error) err } func (r *OrmImpl) WithContext(ctx context.Context) ormcontract.Orm { - instance, _ := NewOrmImpl(ctx, r.config, r.connection, r.query) + for _, query := range r.queries { + query := query.(*databasegorm.QueryImpl) + query.SetContext(ctx) + } + + query := r.query.(*databasegorm.QueryImpl) + query.SetContext(ctx) - return instance + return &OrmImpl{ + ctx: ctx, + config: r.config, + connection: r.connection, + query: query, + queries: r.queries, + } } diff --git a/database/orm_test.go b/database/orm_test.go index c593710ad..6068f6bc1 100644 --- a/database/orm_test.go +++ b/database/orm_test.go @@ -22,6 +22,10 @@ var connections = []contractsorm.Driver{ contractsorm.DriverSqlserver, } +type contextKey int + +const testContextKey contextKey = 0 + type User struct { orm.Model orm.SoftDeletes @@ -92,8 +96,9 @@ func TestOrmSuite(t *testing.T) { func (s *OrmSuite) SetupTest() { s.orm = &OrmImpl{ - ctx: context.Background(), - query: testMysqlQuery, + connection: contractsorm.DriverMysql.String(), + ctx: context.Background(), + query: testMysqlQuery, queries: map[string]contractsorm.Query{ contractsorm.DriverMysql.String(): testMysqlQuery, contractsorm.DriverPostgresql.String(): testPostgresqlQuery, @@ -194,6 +199,38 @@ func (s *OrmSuite) TestTransactionError() { } } +func (s *OrmSuite) TestWithContext() { + s.orm.Observe(User{}, &UserObserver{}) + ctx := context.WithValue(context.Background(), testContextKey, "with_context_goravel") + user := User{Name: "with_context_name"} + + // Call Query directly + err := s.orm.WithContext(ctx).Query().Create(&user) + s.Nil(err) + s.Equal("with_context_name", user.Name) + s.Equal("with_context_goravel", user.Avatar) + + // Call Connection, then call WithContext + for _, connection := range connections { + user.ID = 0 + user.Avatar = "" + err := s.orm.Connection(connection.String()).WithContext(ctx).Query().Create(&user) + s.Nil(err) + s.Equal("with_context_name", user.Name) + s.Equal("with_context_goravel", user.Avatar) + } + + // Call WithContext, then call Connection + for _, connection := range connections { + user.ID = 0 + user.Avatar = "" + err := s.orm.WithContext(ctx).Connection(connection.String()).Query().Create(&user) + s.Nil(err) + s.Equal("with_context_name", user.Name) + s.Equal("with_context_goravel", user.Avatar) + } +} + type UserObserver struct{} func (u *UserObserver) Retrieved(event contractsorm.Event) error { @@ -202,8 +239,15 @@ func (u *UserObserver) Retrieved(event contractsorm.Event) error { func (u *UserObserver) Creating(event contractsorm.Event) error { name := event.GetAttribute("name") - if name != nil && name.(string) == "observer_name" { - return errors.New("error") + if name != nil { + if name.(string) == "observer_name" { + return errors.New("error") + } + if name.(string) == "with_context_name" { + if avatar := event.Context().Value(testContextKey); avatar != nil { + event.SetAttribute("avatar", avatar.(string)) + } + } } return nil diff --git a/database/wire_gen.go b/database/wire_gen.go index 500c34a83..f78d32167 100644 --- a/database/wire_gen.go +++ b/database/wire_gen.go @@ -20,7 +20,7 @@ func InitializeOrm(ctx context.Context, config2 config.Config, connection string configImpl := db.NewConfigImpl(config2, connection) dialectorImpl := gorm.NewDialectorImpl(config2, connection) gormImpl := gorm.NewGormImpl(config2, connection, configImpl, dialectorImpl) - queryImpl, err := gorm.NewQueryImpl(ctx, config2, gormImpl) + queryImpl, err := gorm.BuildQueryImpl(ctx, config2, connection, gormImpl) if err != nil { return nil, err } diff --git a/event/application.go b/event/application.go index 64f117400..4a6809163 100644 --- a/event/application.go +++ b/event/application.go @@ -17,10 +17,14 @@ func NewApplication(queue queuecontract.Queue) *Application { } func (app *Application) Register(events map[event.Event][]event.Listener) { - app.events = events var jobs []queuecontract.Job - for _, listeners := range events { + if app.events == nil { + app.events = map[event.Event][]event.Listener{} + } + + for e, listeners := range events { + app.events[e] = listeners for _, listener := range listeners { jobs = append(jobs, listener) } diff --git a/filesystem/application_test.go b/filesystem/application_test.go deleted file mode 100644 index 91eeb3a1a..000000000 --- a/filesystem/application_test.go +++ /dev/null @@ -1,423 +0,0 @@ -package filesystem - -import ( - "io" - "mime" - "net/http" - "os" - "testing" - - "github.com/gookit/color" - "github.com/stretchr/testify/assert" - - configmocks "github.com/goravel/framework/contracts/config/mocks" - "github.com/goravel/framework/contracts/filesystem" - "github.com/goravel/framework/support/carbon" - "github.com/goravel/framework/support/file" -) - -type TestDisk struct { - disk string - url string -} - -func TestStorage(t *testing.T) { - if !file.Exists("../.env") && os.Getenv("AWS_ACCESS_KEY_ID") == "" { - color.Redln("No filesystem tests run, need create .env based on .env.example, then initialize it") - return - } - - assert.Nil(t, file.Create("test.txt", "Goravel")) - mockConfig := initConfig() - - var driver filesystem.Driver - - disks := []TestDisk{ - { - disk: "local", - url: "http://localhost/storage", - }, - { - disk: "custom", - url: "http://localhost/storage", - }, - } - - tests := []struct { - name string - setup func(disk TestDisk) - }{ - { - name: "AllDirectories", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("AllDirectories/1.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("AllDirectories/2.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("AllDirectories/3/3.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("AllDirectories/3/5/6/6.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.MakeDirectory("AllDirectories/3/4"), disk.disk) - assert.True(t, driver.Exists("AllDirectories/1.txt"), disk.disk) - assert.True(t, driver.Exists("AllDirectories/2.txt"), disk.disk) - assert.True(t, driver.Exists("AllDirectories/3/3.txt"), disk.disk) - assert.True(t, driver.Exists("AllDirectories/3/4/"), disk.disk) - assert.True(t, driver.Exists("AllDirectories/3/5/6/6.txt"), disk.disk) - files, err := driver.AllDirectories("AllDirectories") - assert.Nil(t, err) - assert.Equal(t, []string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) - files, err = driver.AllDirectories("./AllDirectories") - assert.Nil(t, err) - assert.Equal(t, []string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) - files, err = driver.AllDirectories("/AllDirectories") - assert.Nil(t, err) - assert.Equal(t, []string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) - files, err = driver.AllDirectories("./AllDirectories/") - assert.Nil(t, err) - assert.Equal(t, []string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) - assert.Nil(t, driver.DeleteDirectory("AllDirectories"), disk.disk) - }, - }, - { - name: "AllFiles", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("AllFiles/1.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("AllFiles/2.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("AllFiles/3/3.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("AllFiles/3/4/4.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("AllFiles/1.txt"), disk.disk) - assert.True(t, driver.Exists("AllFiles/2.txt"), disk.disk) - assert.True(t, driver.Exists("AllFiles/3/3.txt"), disk.disk) - assert.True(t, driver.Exists("AllFiles/3/4/4.txt"), disk.disk) - files, err := driver.AllFiles("AllFiles") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) - files, err = driver.AllFiles("./AllFiles") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) - files, err = driver.AllFiles("/AllFiles") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) - files, err = driver.AllFiles("./AllFiles/") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) - assert.Nil(t, driver.DeleteDirectory("AllFiles"), disk.disk) - }, - }, - { - name: "Copy", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Copy/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Copy/1.txt"), disk.disk) - assert.Nil(t, driver.Copy("Copy/1.txt", "Copy1/1.txt"), disk.disk) - assert.True(t, driver.Exists("Copy/1.txt"), disk.disk) - assert.True(t, driver.Exists("Copy1/1.txt"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("Copy"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("Copy1"), disk.disk) - }, - }, - { - name: "Delete", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Delete/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Delete/1.txt"), disk.disk) - assert.Nil(t, driver.Delete("Delete/1.txt"), disk.disk) - assert.True(t, driver.Missing("Delete/1.txt"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("Delete"), disk.disk) - }, - }, - { - name: "DeleteDirectory", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("DeleteDirectory/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("DeleteDirectory/1.txt"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("DeleteDirectory"), disk.disk) - assert.True(t, driver.Missing("DeleteDirectory/1.txt"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("DeleteDirectory"), disk.disk) - }, - }, - { - name: "Directories", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Directories/1.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("Directories/2.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("Directories/3/3.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("Directories/3/5/5.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.MakeDirectory("Directories/3/4"), disk.disk) - assert.True(t, driver.Exists("Directories/1.txt"), disk.disk) - assert.True(t, driver.Exists("Directories/2.txt"), disk.disk) - assert.True(t, driver.Exists("Directories/3/3.txt"), disk.disk) - assert.True(t, driver.Exists("Directories/3/4/"), disk.disk) - assert.True(t, driver.Exists("Directories/3/5/5.txt"), disk.disk) - files, err := driver.Directories("Directories") - assert.Nil(t, err) - assert.Equal(t, []string{"3/"}, files) - files, err = driver.Directories("./Directories") - assert.Nil(t, err) - assert.Equal(t, []string{"3/"}, files) - files, err = driver.Directories("/Directories") - assert.Nil(t, err) - assert.Equal(t, []string{"3/"}, files) - files, err = driver.Directories("./Directories/") - assert.Nil(t, err) - assert.Equal(t, []string{"3/"}, files) - assert.Nil(t, driver.DeleteDirectory("Directories"), disk.disk) - }, - }, - { - name: "Files", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Files/1.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("Files/2.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("Files/3/3.txt", "Goravel"), disk.disk) - assert.Nil(t, driver.Put("Files/3/4/4.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Files/1.txt"), disk.disk) - assert.True(t, driver.Exists("Files/2.txt"), disk.disk) - assert.True(t, driver.Exists("Files/3/3.txt"), disk.disk) - assert.True(t, driver.Exists("Files/3/4/4.txt"), disk.disk) - files, err := driver.Files("Files") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt"}, files) - files, err = driver.Files("./Files") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt"}, files) - files, err = driver.Files("/Files") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt"}, files) - files, err = driver.Files("./Files/") - assert.Nil(t, err) - assert.Equal(t, []string{"1.txt", "2.txt"}, files) - assert.Nil(t, driver.DeleteDirectory("Files"), disk.disk) - }, - }, - { - name: "Get", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Get/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Get/1.txt"), disk.disk) - data, err := driver.Get("Get/1.txt") - assert.Nil(t, err) - assert.Equal(t, "Goravel", data) - length, err := driver.Size("Get/1.txt") - assert.Nil(t, err) - assert.Equal(t, int64(7), length) - assert.Nil(t, driver.DeleteDirectory("Get"), disk.disk) - }, - }, - { - name: "LastModified", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("LastModified/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("LastModified/1.txt"), disk.disk) - date, err := driver.LastModified("LastModified/1.txt") - assert.Nil(t, err) - - assert.Nil(t, err, disk.disk) - assert.Equal(t, carbon.Now().ToDateString(), carbon.FromStdTime(date).ToDateString(), disk.disk) - assert.Nil(t, driver.DeleteDirectory("LastModified"), disk.disk) - }, - }, - { - name: "MakeDirectory", - setup: func(disk TestDisk) { - assert.Nil(t, driver.MakeDirectory("MakeDirectory1/"), disk.disk) - assert.Nil(t, driver.MakeDirectory("MakeDirectory2"), disk.disk) - assert.Nil(t, driver.MakeDirectory("MakeDirectory3/MakeDirectory4"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("MakeDirectory1"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("MakeDirectory2"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("MakeDirectory3"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("MakeDirectory4"), disk.disk) - }, - }, - { - name: "MimeType", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("MimeType/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("MimeType/1.txt"), disk.disk) - mimeType, err := driver.MimeType("MimeType/1.txt") - assert.Nil(t, err, disk.disk) - mediaType, _, err := mime.ParseMediaType(mimeType) - assert.Nil(t, err, disk.disk) - assert.Equal(t, "text/plain", mediaType, disk.disk) - - fileInfo, err := NewFile("../logo.png") - assert.Nil(t, err, disk.disk) - path, err := driver.PutFile("MimeType", fileInfo) - assert.Nil(t, err, disk.disk) - assert.True(t, driver.Exists(path), disk.disk) - mimeType, err = driver.MimeType(path) - assert.Nil(t, err, disk.disk) - assert.Equal(t, "image/png", mimeType, disk.disk) - }, - }, - { - name: "Move", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Move/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Move/1.txt"), disk.disk) - assert.Nil(t, driver.Move("Move/1.txt", "Move1/1.txt"), disk.disk) - assert.True(t, driver.Missing("Move/1.txt"), disk.disk) - assert.True(t, driver.Exists("Move1/1.txt"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("Move"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("Move1"), disk.disk) - }, - }, - { - name: "Put", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Put/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Put/1.txt"), disk.disk) - assert.True(t, driver.Missing("Put/2.txt"), disk.disk) - assert.Nil(t, driver.DeleteDirectory("Put"), disk.disk) - }, - }, - { - name: "PutFile_Image", - setup: func(disk TestDisk) { - fileInfo, err := NewFile("../logo.png") - assert.Nil(t, err) - path, err := driver.PutFile("PutFile1", fileInfo) - assert.Nil(t, err) - assert.True(t, driver.Exists(path), disk.disk) - assert.Nil(t, driver.DeleteDirectory("PutFile1"), disk.disk) - }, - }, - { - name: "PutFile_Text", - setup: func(disk TestDisk) { - fileInfo, err := NewFile("./test.txt") - assert.Nil(t, err) - path, err := driver.PutFile("PutFile", fileInfo) - assert.Nil(t, err) - assert.True(t, driver.Exists(path), disk.disk) - data, err := driver.Get(path) - assert.Nil(t, err) - assert.Equal(t, "Goravel", data) - assert.Nil(t, driver.DeleteDirectory("PutFile"), disk.disk) - }, - }, - { - name: "PutFileAs_Text", - setup: func(disk TestDisk) { - fileInfo, err := NewFile("./test.txt") - assert.Nil(t, err) - path, err := driver.PutFileAs("PutFileAs", fileInfo, "text") - assert.Nil(t, err) - assert.Equal(t, "PutFileAs/text.txt", path) - assert.True(t, driver.Exists(path), disk.disk) - data, err := driver.Get(path) - assert.Nil(t, err) - assert.Equal(t, "Goravel", data) - - path, err = driver.PutFileAs("PutFileAs", fileInfo, "text1.txt") - assert.Nil(t, err) - assert.Equal(t, "PutFileAs/text1.txt", path) - assert.True(t, driver.Exists(path), disk.disk) - data, err = driver.Get(path) - assert.Nil(t, err) - assert.Equal(t, "Goravel", data) - - assert.Nil(t, driver.DeleteDirectory("PutFileAs"), disk.disk) - }, - }, - { - name: "PutFileAs_Image", - setup: func(disk TestDisk) { - fileInfo, err := NewFile("../logo.png") - assert.Nil(t, err) - path, err := driver.PutFileAs("PutFileAs1", fileInfo, "image") - assert.Nil(t, err) - assert.Equal(t, "PutFileAs1/image.png", path) - assert.True(t, driver.Exists(path), disk.disk) - - path, err = driver.PutFileAs("PutFileAs1", fileInfo, "image1.png") - assert.Nil(t, err) - assert.Equal(t, "PutFileAs1/image1.png", path) - assert.True(t, driver.Exists(path), disk.disk) - - assert.Nil(t, driver.DeleteDirectory("PutFileAs1"), disk.disk) - }, - }, - { - name: "Size", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Size/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Size/1.txt"), disk.disk) - length, err := driver.Size("Size/1.txt") - assert.Nil(t, err) - assert.Equal(t, int64(7), length) - assert.Nil(t, driver.DeleteDirectory("Size"), disk.disk) - }, - }, - { - name: "TemporaryUrl", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("TemporaryUrl/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("TemporaryUrl/1.txt"), disk.disk) - url, err := driver.TemporaryUrl("TemporaryUrl/1.txt", carbon.Now().AddSeconds(5).ToStdTime()) - assert.Nil(t, err) - assert.NotEmpty(t, url) - if disk.disk != "local" && disk.disk != "custom" { - resp, err := http.Get(url) - assert.Nil(t, err) - content, err := io.ReadAll(resp.Body) - assert.Nil(t, resp.Body.Close()) - assert.Nil(t, err) - assert.Equal(t, "Goravel", string(content), disk.disk) - } - assert.Nil(t, driver.DeleteDirectory("TemporaryUrl"), disk.disk) - }, - }, - { - name: "Url", - setup: func(disk TestDisk) { - assert.Nil(t, driver.Put("Url/1.txt", "Goravel"), disk.disk) - assert.True(t, driver.Exists("Url/1.txt"), disk.disk) - url := disk.url + "/Url/1.txt" - assert.Equal(t, url, driver.Url("Url/1.txt"), disk.disk) - if disk.disk != "local" && disk.disk != "custom" { - resp, err := http.Get(url) - assert.Nil(t, err) - content, err := io.ReadAll(resp.Body) - assert.Nil(t, resp.Body.Close()) - assert.Nil(t, err) - assert.Equal(t, "Goravel", string(content), disk.disk) - } - assert.Nil(t, driver.DeleteDirectory("Url"), disk.disk) - }, - }, - } - - for _, disk := range disks { - var err error - driver, err = NewDriver(mockConfig, disk.disk) - assert.NotNil(t, driver) - assert.Nil(t, err) - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - test.setup(disk) - }) - } - - if disk.disk == "local" || disk.disk == "custom" { - assert.Nil(t, file.Remove("./storage")) - } - } - - assert.Nil(t, file.Remove("test.txt")) -} - -func initConfig() *configmocks.Config { - mockConfig := &configmocks.Config{} - ConfigFacade = mockConfig - mockConfig.On("GetString", "app.timezone").Return("UTC") - mockConfig.On("GetString", "filesystems.default").Return("local") - mockConfig.On("GetString", "filesystems.disks.local.driver").Return("local") - mockConfig.On("GetString", "filesystems.disks.local.root").Return("storage/app") - mockConfig.On("GetString", "filesystems.disks.local.url").Return("http://localhost/storage") - mockConfig.On("GetString", "filesystems.disks.custom.driver").Return("custom") - mockConfig.On("Get", "filesystems.disks.custom.via").Return(&Local{ - config: mockConfig, - root: "storage/app/public", - url: "http://localhost/storage", - }) - - return mockConfig -} diff --git a/filesystem/local.go b/filesystem/local.go index 583c12824..c92248882 100644 --- a/filesystem/local.go +++ b/filesystem/local.go @@ -135,12 +135,18 @@ func (r *Local) Files(path string) ([]string, error) { } func (r *Local) Get(file string) (string, error) { + data, err := r.GetBytes(file) + + return string(data), err +} + +func (r *Local) GetBytes(file string) ([]byte, error) { data, err := os.ReadFile(r.fullPath(file)) if err != nil { - return "", err + return nil, err } - return string(data), nil + return data, nil } func (r *Local) LastModified(file string) (time.Time, error) { @@ -230,7 +236,7 @@ func (r *Local) WithContext(ctx context.Context) filesystem.Driver { } func (r *Local) Url(file string) string { - return strings.TrimSuffix(r.url, "/") + "/" + strings.TrimPrefix(file, "/") + return strings.TrimSuffix(r.url, "/") + "/" + strings.TrimPrefix(filepath.ToSlash(file), "/") } func (r *Local) fullPath(path string) string { diff --git a/filesystem/local_test.go b/filesystem/local_test.go new file mode 100644 index 000000000..147aca901 --- /dev/null +++ b/filesystem/local_test.go @@ -0,0 +1,430 @@ +package filesystem + +import ( + "context" + "mime" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/suite" + + configmock "github.com/goravel/framework/contracts/config/mocks" + "github.com/goravel/framework/support/carbon" + "github.com/goravel/framework/support/env" + "github.com/goravel/framework/support/file" +) + +type LocalTestSuite struct { + suite.Suite + local *Local + file *File + mockConfig *configmock.Config +} + +func TestLocalTestSuite(t *testing.T) { + suite.Run(t, new(LocalTestSuite)) + + assert.Nil(t, file.Remove("test.txt")) +} + +func (s *LocalTestSuite) SetupTest() { + s.mockConfig = &configmock.Config{} + + dir, err := os.MkdirTemp("", "local-test") + s.Nil(err) + + err = os.WriteFile(dir+"/test.txt", []byte("goravel"), 0644) + s.Nil(err) + + err = os.Mkdir(dir+"/test", 0755) + s.Nil(err) + + s.mockConfig.On("GetString", "filesystems.default").Return("local").Once() + s.mockConfig.On("GetString", "filesystems.disks.local.root").Return(dir).Once() + s.mockConfig.On("GetString", "filesystems.disks.local.url").Return("https://goravel.dev").Once() + ConfigFacade = s.mockConfig + + s.local, err = NewLocal(s.mockConfig, "local") + s.Nil(err) + s.NotNil(s.local) + + s.file, err = NewFile("./file.go") + s.Nil(err) + s.NotNil(s.file) + + s.mockConfig.AssertExpectations(s.T()) +} + +func (s *LocalTestSuite) TestAllDirectories() { + s.Nil(s.local.Put("AllDirectories/1.txt", "Goravel")) + s.Nil(s.local.Put("AllDirectories/2.txt", "Goravel")) + s.Nil(s.local.Put("AllDirectories/3/3.txt", "Goravel")) + s.Nil(s.local.Put("AllDirectories/3/5/6/6.txt", "Goravel")) + s.Nil(s.local.MakeDirectory("AllDirectories/3/4")) + s.True(s.local.Exists("AllDirectories/1.txt")) + s.True(s.local.Exists("AllDirectories/2.txt")) + s.True(s.local.Exists("AllDirectories/3/3.txt")) + s.True(s.local.Exists("AllDirectories/3/4/")) + s.True(s.local.Exists("AllDirectories/3/5/6/6.txt")) + files, err := s.local.AllDirectories("AllDirectories") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\", "3\\4\\", "3\\5\\", "3\\5\\6\\"}, files) + } else { + s.Equal([]string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) + } + files, err = s.local.AllDirectories("./AllDirectories") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\", "3\\4\\", "3\\5\\", "3\\5\\6\\"}, files) + } else { + s.Equal([]string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) + } + files, err = s.local.AllDirectories("/AllDirectories") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\", "3\\4\\", "3\\5\\", "3\\5\\6\\"}, files) + } else { + s.Equal([]string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) + } + files, err = s.local.AllDirectories("./AllDirectories/") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\", "3\\4\\", "3\\5\\", "3\\5\\6\\"}, files) + } else { + s.Equal([]string{"3/", "3/4/", "3/5/", "3/5/6/"}, files) + } + s.Nil(s.local.DeleteDirectory("AllDirectories")) +} + +func (s *LocalTestSuite) TestAllFiles() { + s.Nil(s.local.Put("AllFiles/1.txt", "Goravel")) + s.Nil(s.local.Put("AllFiles/2.txt", "Goravel")) + s.Nil(s.local.Put("AllFiles/3/3.txt", "Goravel")) + s.Nil(s.local.Put("AllFiles/3/4/4.txt", "Goravel")) + s.True(s.local.Exists("AllFiles/1.txt")) + s.True(s.local.Exists("AllFiles/2.txt")) + s.True(s.local.Exists("AllFiles/3/3.txt")) + s.True(s.local.Exists("AllFiles/3/4/4.txt")) + files, err := s.local.AllFiles("AllFiles") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"1.txt", "2.txt", "3\\3.txt", "3\\4\\4.txt"}, files) + } else { + s.Equal([]string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) + } + files, err = s.local.AllFiles("./AllFiles") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"1.txt", "2.txt", "3\\3.txt", "3\\4\\4.txt"}, files) + } else { + s.Equal([]string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) + } + files, err = s.local.AllFiles("/AllFiles") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"1.txt", "2.txt", "3\\3.txt", "3\\4\\4.txt"}, files) + } else { + s.Equal([]string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) + } + files, err = s.local.AllFiles("./AllFiles/") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"1.txt", "2.txt", "3\\3.txt", "3\\4\\4.txt"}, files) + } else { + s.Equal([]string{"1.txt", "2.txt", "3/3.txt", "3/4/4.txt"}, files) + } + s.Nil(s.local.DeleteDirectory("AllFiles")) +} + +func (s *LocalTestSuite) TestCopy() { + s.Nil(s.local.Put("Copy/1.txt", "Goravel")) + s.True(s.local.Exists("Copy/1.txt")) + s.Nil(s.local.Copy("Copy/1.txt", "Copy1/1.txt")) + s.True(s.local.Exists("Copy/1.txt")) + s.True(s.local.Exists("Copy1/1.txt")) + s.Nil(s.local.DeleteDirectory("Copy")) + s.Nil(s.local.DeleteDirectory("Copy1")) +} + +func (s *LocalTestSuite) TestDelete() { + s.Nil(s.local.Put("Delete/1.txt", "Goravel")) + s.True(s.local.Exists("Delete/1.txt")) + s.Nil(s.local.Delete("Delete/1.txt")) + s.True(s.local.Missing("Delete/1.txt")) + s.Nil(s.local.DeleteDirectory("Delete")) +} + +func (s *LocalTestSuite) TestDeleteDirectory() { + s.Nil(s.local.Put("DeleteDirectory/1.txt", "Goravel")) + s.True(s.local.Exists("DeleteDirectory/1.txt")) + s.Nil(s.local.DeleteDirectory("DeleteDirectory")) + s.True(s.local.Missing("DeleteDirectory/1.txt")) + s.Nil(s.local.DeleteDirectory("DeleteDirectory")) +} + +func (s *LocalTestSuite) TestDirectories() { + s.Nil(s.local.Put("Directories/1.txt", "Goravel")) + s.Nil(s.local.Put("Directories/2.txt", "Goravel")) + s.Nil(s.local.Put("Directories/3/3.txt", "Goravel")) + s.Nil(s.local.Put("Directories/3/5/5.txt", "Goravel")) + s.Nil(s.local.MakeDirectory("Directories/3/4")) + s.True(s.local.Exists("Directories/1.txt")) + s.True(s.local.Exists("Directories/2.txt")) + s.True(s.local.Exists("Directories/3/3.txt")) + s.True(s.local.Exists("Directories/3/4/")) + s.True(s.local.Exists("Directories/3/5/5.txt")) + files, err := s.local.Directories("Directories") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\"}, files) + } else { + s.Equal([]string{"3/"}, files) + } + files, err = s.local.Directories("./Directories") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\"}, files) + } else { + s.Equal([]string{"3/"}, files) + } + files, err = s.local.Directories("/Directories") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\"}, files) + } else { + s.Equal([]string{"3/"}, files) + } + files, err = s.local.Directories("./Directories/") + s.Nil(err) + if env.IsWindows() { + s.Equal([]string{"3\\"}, files) + } else { + s.Equal([]string{"3/"}, files) + } + s.Nil(s.local.DeleteDirectory("Directories")) +} + +func (s *LocalTestSuite) TestExists() { + exists := s.local.Exists("test.txt") + s.True(exists) + + exists = s.local.Exists("test1.txt") + s.False(exists) +} + +func (s *LocalTestSuite) TestFiles() { + s.Nil(s.local.Put("Files/1.txt", "Goravel")) + s.Nil(s.local.Put("Files/2.txt", "Goravel")) + s.Nil(s.local.Put("Files/3/3.txt", "Goravel")) + s.Nil(s.local.Put("Files/3/4/4.txt", "Goravel")) + s.True(s.local.Exists("Files/1.txt")) + s.True(s.local.Exists("Files/2.txt")) + s.True(s.local.Exists("Files/3/3.txt")) + s.True(s.local.Exists("Files/3/4/4.txt")) + files, err := s.local.Files("Files") + s.Nil(err) + s.Equal([]string{"1.txt", "2.txt"}, files) + files, err = s.local.Files("./Files") + s.Nil(err) + s.Equal([]string{"1.txt", "2.txt"}, files) + files, err = s.local.Files("/Files") + s.Nil(err) + s.Equal([]string{"1.txt", "2.txt"}, files) + files, err = s.local.Files("./Files/") + s.Nil(err) + s.Equal([]string{"1.txt", "2.txt"}, files) + s.Nil(s.local.DeleteDirectory("Files")) +} + +func (s *LocalTestSuite) TestGet() { + s.Nil(s.local.Put("Get/1.txt", "Goravel")) + s.True(s.local.Exists("Get/1.txt")) + data, err := s.local.Get("Get/1.txt") + s.Nil(err) + s.Equal("Goravel", data) + length, err := s.local.Size("Get/1.txt") + s.Nil(err) + s.Equal(int64(7), length) + s.Nil(s.local.DeleteDirectory("Get")) +} + +func (s *LocalTestSuite) TestGetBytes() { + s.Nil(s.local.Put("Get/1.txt", "Goravel")) + s.True(s.local.Exists("Get/1.txt")) + data, err := s.local.GetBytes("Get/1.txt") + s.Nil(err) + s.Equal([]byte("Goravel"), data) + length, err := s.local.Size("Get/1.txt") + s.Nil(err) + s.Equal(int64(7), length) + s.Nil(s.local.DeleteDirectory("Get")) +} + +func (s *LocalTestSuite) TestLastModified() { + s.mockConfig.On("GetString", "app.timezone").Return("UTC").Once() + + s.Nil(s.local.Put("LastModified/1.txt", "Goravel")) + s.True(s.local.Exists("LastModified/1.txt")) + date, err := s.local.LastModified("LastModified/1.txt") + s.Nil(err) + + s.Nil(err) + s.Equal(carbon.Now().ToDateString(), carbon.FromStdTime(date).ToDateString()) + s.Nil(s.local.DeleteDirectory("LastModified")) +} + +func (s *LocalTestSuite) TestMakeDirectory() { + s.Nil(s.local.MakeDirectory("MakeDirectory1/")) + s.Nil(s.local.MakeDirectory("MakeDirectory2")) + s.Nil(s.local.MakeDirectory("MakeDirectory3/MakeDirectory4")) + s.Nil(s.local.DeleteDirectory("MakeDirectory1")) + s.Nil(s.local.DeleteDirectory("MakeDirectory2")) + s.Nil(s.local.DeleteDirectory("MakeDirectory3")) + s.Nil(s.local.DeleteDirectory("MakeDirectory4")) +} + +func (s *LocalTestSuite) TestMimeType_File() { + s.Nil(s.local.Put("MimeType/1.txt", "Goravel")) + s.True(s.local.Exists("MimeType/1.txt")) + mimeType, err := s.local.MimeType("MimeType/1.txt") + s.Nil(err) + mediaType, _, err := mime.ParseMediaType(mimeType) + s.Nil(err) + s.Equal("text/plain", mediaType) +} + +func (s *LocalTestSuite) TestMimeType_Image() { + s.mockConfig.On("GetString", "filesystems.default").Return("local").Once() + + fileInfo, err := NewFile("../logo.png") + s.Nil(err) + path, err := s.local.PutFile("MimeType", fileInfo) + s.Nil(err) + s.True(s.local.Exists(path)) + mimeType, err := s.local.MimeType(path) + s.Nil(err) + s.Equal("image/png", mimeType) +} + +func (s *LocalTestSuite) TestMissing() { + missing := s.local.Missing("test.txt") + s.False(missing) + + missing = s.local.Missing("test1.txt") + s.True(missing) +} + +func (s *LocalTestSuite) TestMove() { + s.Nil(s.local.Put("Move/1.txt", "Goravel")) + s.True(s.local.Exists("Move/1.txt")) + s.Nil(s.local.Move("Move/1.txt", "Move1/1.txt")) + s.True(s.local.Missing("Move/1.txt")) + s.True(s.local.Exists("Move1/1.txt")) + s.Nil(s.local.DeleteDirectory("Move")) + s.Nil(s.local.DeleteDirectory("Move1")) +} + +func (s *LocalTestSuite) TestPath() { + path := s.local.Path("test.txt") + s.Equal(filepath.Join(s.local.root, "test.txt"), path) +} + +func (s *LocalTestSuite) TestPut() { + s.Nil(s.local.Put("Put/1.txt", "Goravel")) + s.True(s.local.Exists("Put/1.txt")) + s.True(s.local.Missing("Put/2.txt")) + s.Nil(s.local.DeleteDirectory("Put")) +} + +func (s *LocalTestSuite) TestPutFile_Text() { + path, err := s.local.PutFile("PutFile", s.file) + s.Nil(err) + s.True(s.local.Exists(path)) + data, err := s.local.Get(path) + s.Nil(err) + s.NotEmpty(data) + s.Nil(s.local.DeleteDirectory("PutFile")) +} + +func (s *LocalTestSuite) TestPutFile_Image() { + s.mockConfig.On("GetString", "filesystems.default").Return("local").Once() + + fileInfo, err := NewFile("../logo.png") + s.Nil(err) + path, err := s.local.PutFile("PutFile1", fileInfo) + s.Nil(err) + s.True(s.local.Exists(path)) + s.Nil(s.local.DeleteDirectory("PutFile1")) +} + +func (s *LocalTestSuite) TestPutFileAs_Text() { + path, err := s.local.PutFileAs("PutFileAs", s.file, "text") + s.Nil(err) + s.Equal(filepath.Join("PutFileAs", "text.txt"), path) + s.True(s.local.Exists(path)) + data, err := s.local.Get(path) + s.Nil(err) + s.NotEmpty(data) + + path, err = s.local.PutFileAs("PutFileAs", s.file, "text1.txt") + s.Nil(err) + s.Equal(filepath.Join("PutFileAs", "text1.txt"), path) + s.True(s.local.Exists(path)) + data, err = s.local.Get(path) + s.Nil(err) + s.NotEmpty(data) + + s.Nil(s.local.DeleteDirectory("PutFileAs")) +} + +func (s *LocalTestSuite) TestPutFileAs_Image() { + s.mockConfig.On("GetString", "filesystems.default").Return("local").Once() + + fileInfo, err := NewFile("../logo.png") + s.Nil(err) + path, err := s.local.PutFileAs("PutFileAs1", fileInfo, "image") + s.Nil(err) + s.Equal(filepath.Join("PutFileAs1", "image.png"), path) + s.True(s.local.Exists(path)) + + path, err = s.local.PutFileAs("PutFileAs1", fileInfo, "image1.png") + s.Nil(err) + s.Equal(filepath.Join("PutFileAs1", "image1.png"), path) + s.True(s.local.Exists(path)) + + s.Nil(s.local.DeleteDirectory("PutFileAs1")) +} + +func (s *LocalTestSuite) TestSize() { + s.Nil(s.local.Put("Size/1.txt", "Goravel")) + s.True(s.local.Exists("Size/1.txt")) + length, err := s.local.Size("Size/1.txt") + s.Nil(err) + s.Equal(int64(7), length) + s.Nil(s.local.DeleteDirectory("Size")) +} + +func (s *LocalTestSuite) TestTemporaryUrl() { + s.Nil(s.local.Put("TemporaryUrl/1.txt", "Goravel")) + s.True(s.local.Exists("TemporaryUrl/1.txt")) + url, err := s.local.TemporaryUrl("TemporaryUrl/1.txt", carbon.Now().AddSeconds(5).ToStdTime()) + s.Nil(err) + s.NotEmpty(url) + s.Nil(s.local.DeleteDirectory("TemporaryUrl")) +} + +func (s *LocalTestSuite) TestWithContext() { + driver := s.local.WithContext(context.Background()) + s.Equal(s.local, driver) +} + +func (s *LocalTestSuite) TestUrl() { + s.Equal("https://goravel.dev/Url/1.txt", s.local.Url("Url/1.txt")) + + if env.IsWindows() { + s.Equal("https://goravel.dev/Url/2.txt", s.local.Url(`Url\2.txt`)) + } +} diff --git a/foundation/application.go b/foundation/application.go index d97487401..4348c796e 100644 --- a/foundation/application.go +++ b/foundation/application.go @@ -181,7 +181,9 @@ func setEnv() { for _, arg := range args[1:] { if arg == "artisan" { support.Env = support.EnvArtisan - break + } + if arg == "key:generate" { + support.IsKeyGenerateCommand = true } } } diff --git a/go.mod b/go.mod index b6a89bc69..c516c9591 100644 --- a/go.mod +++ b/go.mod @@ -4,7 +4,7 @@ go 1.20 require ( github.com/RichardKnop/machinery/v2 v2.0.11 - github.com/bytedance/sonic v1.10.0 + github.com/bytedance/sonic v1.10.1 github.com/davecgh/go-spew v1.1.1 github.com/gabriel-vasile/mimetype v1.4.2 github.com/glebarez/go-sqlite v1.21.2 @@ -13,7 +13,7 @@ require ( github.com/go-sql-driver/mysql v1.7.1 github.com/golang-jwt/jwt/v5 v5.0.0 github.com/golang-migrate/migrate/v4 v4.16.2 - github.com/golang-module/carbon/v2 v2.2.6 + github.com/golang-module/carbon/v2 v2.2.8 github.com/golang/protobuf v1.5.3 github.com/google/uuid v1.3.1 github.com/google/wire v0.5.0 @@ -36,7 +36,8 @@ require ( github.com/urfave/cli/v2 v2.25.7 go.uber.org/atomic v1.11.0 golang.org/x/crypto v0.13.0 - google.golang.org/grpc v1.58.0 + golang.org/x/exp v0.0.0-20230315142452-642cacee5cc0 + google.golang.org/grpc v1.58.2 gorm.io/driver/mysql v1.5.1 gorm.io/driver/postgres v1.5.2 gorm.io/driver/sqlserver v1.5.1 @@ -142,7 +143,7 @@ require ( golang.org/x/oauth2 v0.10.0 // indirect golang.org/x/sync v0.3.0 // indirect golang.org/x/sys v0.12.0 // indirect - golang.org/x/text v0.13.0 // indirect + golang.org/x/text v0.13.0 golang.org/x/tools v0.9.1 // indirect google.golang.org/api v0.126.0 // indirect google.golang.org/appengine v1.6.7 // indirect diff --git a/go.sum b/go.sum index 94a2d4f08..089973e21 100644 --- a/go.sum +++ b/go.sum @@ -86,8 +86,8 @@ github.com/brianvoe/gofakeit/v6 v6.23.2 h1:lVde18uhad5wII/f5RMVFLtdQNE0HaGFuBUXm github.com/brianvoe/gofakeit/v6 v6.23.2/go.mod h1:Ow6qC71xtwm79anlwKRlWZW6zVq9D2XHE4QSSMP/rU8= github.com/bytedance/sonic v1.5.0/go.mod h1:ED5hyg4y6t3/9Ku1R6dU/4KyJ48DZ4jPhfY1O2AihPM= github.com/bytedance/sonic v1.10.0-rc/go.mod h1:ElCzW+ufi8qKqNW0FY314xriJhyJhuoJ3gFZdAHF7NM= -github.com/bytedance/sonic v1.10.0 h1:qtNZduETEIWJVIyDl01BeNxur2rW9OwTQ/yBqFRkKEk= -github.com/bytedance/sonic v1.10.0/go.mod h1:iZcSUejdk5aukTND/Eu/ivjQuEL0Cu9/rf50Hi0u/g4= +github.com/bytedance/sonic v1.10.1 h1:7a1wuFXL1cMy7a3f7/VFcEtriuXQnUBhtoVfOZiaysc= +github.com/bytedance/sonic v1.10.1/go.mod h1:iZcSUejdk5aukTND/Eu/ivjQuEL0Cu9/rf50Hi0u/g4= github.com/cenkalti/backoff/v4 v4.2.0 h1:HN5dHm3WBOgndBH6E8V0q2jIYIR3s9yglV8k/+MN3u4= github.com/cenkalti/backoff/v4 v4.2.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE= github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU= @@ -221,8 +221,8 @@ github.com/golang-jwt/jwt/v5 v5.0.0 h1:1n1XNM9hk7O9mnQoNBGolZvzebBQ7p93ULHRc28XJ github.com/golang-jwt/jwt/v5 v5.0.0/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk= github.com/golang-migrate/migrate/v4 v4.16.2 h1:8coYbMKUyInrFk1lfGfRovTLAW7PhWp8qQDT2iKfuoA= github.com/golang-migrate/migrate/v4 v4.16.2/go.mod h1:pfcJX4nPHaVdc5nmdCikFBWtm+UBpiZjRNNsyBbp0/o= -github.com/golang-module/carbon/v2 v2.2.6 h1:4nHkIx7A7d+aoRDAGLwxZEJWAAZpmH5Bm/pVUl58FXI= -github.com/golang-module/carbon/v2 v2.2.6/go.mod h1:XDALX7KgqmHk95xyLeaqX9/LJGbfLATyruTziq68SZ8= +github.com/golang-module/carbon/v2 v2.2.8 h1:a1VxHHKAR7fc1ho7sYXhS1s5S4x7+oqAf2EY5p8C46A= +github.com/golang-module/carbon/v2 v2.2.8/go.mod h1:XDALX7KgqmHk95xyLeaqX9/LJGbfLATyruTziq68SZ8= github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA= github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0= github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A= @@ -623,6 +623,7 @@ golang.org/x/exp v0.0.0-20200119233911-0405dc783f0a/go.mod h1:2RIsYlXP63K8oxa1u0 golang.org/x/exp v0.0.0-20200207192155-f17229e696bd/go.mod h1:J/WKrq2StrnmMY6+EHIKF9dgMWnmCNThgcyBT1FY9mM= golang.org/x/exp v0.0.0-20200224162631-6cc2880d07d6/go.mod h1:3jZMyOhIsHpP37uCMkUooju7aAi5cS1Q23tOzKc+0MU= golang.org/x/exp v0.0.0-20230315142452-642cacee5cc0 h1:pVgRXcIictcr+lBQIFeiwuwtDIs4eL21OuM9nyAADmo= +golang.org/x/exp v0.0.0-20230315142452-642cacee5cc0/go.mod h1:CxIveKay+FTh1D0yPZemJVgC/95VzuuOLq5Qi4xnoYc= golang.org/x/image v0.0.0-20190227222117-0694c2d4d067/go.mod h1:kZ7UVZpmo3dzQBMxlp+ypCbDeSB+sBbTgSJuh5dn5js= golang.org/x/image v0.0.0-20190802002840-cff245a6509b/go.mod h1:FeLwcggjj3mMvU+oOTbSwawSJRM1uh48EjtB4UJZlP0= golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= @@ -970,8 +971,8 @@ google.golang.org/grpc v1.34.0/go.mod h1:WotjhfgOW/POjDeRt8vscBtXq+2VjORFy659qA5 google.golang.org/grpc v1.35.0/go.mod h1:qjiiYl8FncCW8feJPdyg3v6XW24KsRHe+dy9BAGRRjU= google.golang.org/grpc v1.36.0/go.mod h1:qjiiYl8FncCW8feJPdyg3v6XW24KsRHe+dy9BAGRRjU= google.golang.org/grpc v1.45.0/go.mod h1:lN7owxKUQEqMfSyQikvvk5tf/6zMPsrK+ONuO11+0rQ= -google.golang.org/grpc v1.58.0 h1:32JY8YpPMSR45K+c3o6b8VL73V+rR8k+DeMIr4vRH8o= -google.golang.org/grpc v1.58.0/go.mod h1:tgX3ZQDlNJGU96V6yHh1T/JeoBQ2TXdr43YbYSsCJk0= +google.golang.org/grpc v1.58.2 h1:SXUpjxeVF3FKrTYQI4f4KvbGD5u2xccdYdurwowix5I= +google.golang.org/grpc v1.58.2/go.mod h1:tgX3ZQDlNJGU96V6yHh1T/JeoBQ2TXdr43YbYSsCJk0= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM= diff --git a/log/formatter/general.go b/log/formatter/general.go index 7336824da..e5af5d5a7 100644 --- a/log/formatter/general.go +++ b/log/formatter/general.go @@ -6,10 +6,11 @@ import ( "strings" "time" - "github.com/bytedance/sonic" "github.com/sirupsen/logrus" + "github.com/spf13/cast" "github.com/goravel/framework/contracts/config" + "github.com/goravel/framework/support/json" ) type General struct { @@ -53,44 +54,38 @@ func formatData(data logrus.Fields) (string, error) { var builder strings.Builder if len(data) > 0 { - dataBytes, err := sonic.Marshal(data) - if err != nil { - return "", err - } - removedData := deleteKey(data, "root") if len(removedData) > 0 { - removedDataBytes, err := sonic.Marshal(removedData) + removedDataBytes, err := json.Marshal(removedData) if err != nil { return "", err } + builder.WriteString(fmt.Sprintf("fields: %s\n", string(removedDataBytes))) } - root, err := sonic.Get(dataBytes, "root") + root, err := cast.ToStringMapE(data["root"]) if err != nil { return "", err } for _, key := range []string{"code", "context", "domain", "hint", "owner", "request", "response", "tags", "user"} { - if value := root.Get(key); value.Valid() { - info, err := value.Raw() + if value, exists := root[key]; exists && value != nil { + v, err := json.Marshal(value) if err != nil { return "", err } - builder.WriteString(fmt.Sprintf("%s: %s\n", key, info)) + + builder.WriteString(fmt.Sprintf(`%s: %v"\n`, key, string(v))) } } - if stackTraceValue := root.Get("stacktrace"); stackTraceValue.Valid() { - stackTraces, err := stackTraceValue.Interface() - if err != nil { - return "", err - } - traces, err := formatStackTraces(stackTraces) + if stackTraceValue, exists := root["stacktrace"]; exists && stackTraceValue != nil { + traces, err := formatStackTraces(stackTraceValue) if err != nil { return "", err } + builder.WriteString(traces) } } @@ -105,6 +100,7 @@ func deleteKey(data logrus.Fields, keyToDelete string) logrus.Fields { dataCopy[key] = value } } + return dataCopy } @@ -121,12 +117,13 @@ type StackTrace struct { func formatStackTraces(stackTraces any) (string, error) { var formattedTraces strings.Builder - data, err := sonic.Marshal(stackTraces) + data, err := json.Marshal(stackTraces) + if err != nil { return "", err } var traces StackTrace - err = sonic.Unmarshal(data, &traces) + err = json.Unmarshal(data, &traces) if err != nil { return "", err } diff --git a/log/formatter/general_test.go b/log/formatter/general_test.go index b36b70f61..15d188126 100644 --- a/log/formatter/general_test.go +++ b/log/formatter/general_test.go @@ -56,7 +56,7 @@ func (s *GeneralTestSuite) TestFormat() { name: "Data is not empty", setup: func() { s.entry.Data = logrus.Fields{ - "root": map[string]interface{}{ + "root": map[string]any{ "code": "200", "domain": "example.com", "owner": "owner", @@ -119,14 +119,14 @@ func TestFormatData(t *testing.T) { name: "Invalid data type", setup: func() { data = logrus.Fields{ - "root": map[string]interface{}{ + "root": map[string]any{ "code": "123", "context": "sample", "domain": "example.com", "hint": make(chan int), // Invalid data type that will cause an error during value extraction "owner": "owner", - "request": map[string]interface{}{"method": "GET", "uri": "http://localhost"}, - "response": map[string]interface{}{"status": 200}, + "request": map[string]any{"method": "GET", "uri": "http://localhost"}, + "response": map[string]any{"status": 200}, "tags": []string{"tag1", "tag2"}, "user": "user1", }, @@ -142,7 +142,7 @@ func TestFormatData(t *testing.T) { name: "Data is not empty", setup: func() { data = logrus.Fields{ - "root": map[string]interface{}{ + "root": map[string]any{ "code": "200", "domain": "example.com", "owner": "owner", @@ -234,8 +234,8 @@ func TestFormatStackTraces(t *testing.T) { { name: "StackTraces is not nil", setup: func() { - stackTraces = map[string]interface{}{ - "root": map[string]interface{}{ + stackTraces = map[string]any{ + "root": map[string]any{ "message": "error bad request", // root cause "stack": []string{ "main.main:/dummy/examples/logging/example.go:143", // original calling method @@ -244,7 +244,7 @@ func TestFormatStackTraces(t *testing.T) { "main.(*Request).Validate:/dummy/examples/logging/example.go:28", // location of the root }, }, - "wrap": []map[string]interface{}{ + "wrap": []map[string]any{ { "message": "received a request with no ID", // additional context "stack": "main.(*Request).Validate:/dummy/examples/logging/example.go:29", // location of Wrap call diff --git a/mail/application_test.go b/mail/application_test.go index b9f1e9714..474f5268a 100644 --- a/mail/application_test.go +++ b/mail/application_test.go @@ -13,6 +13,7 @@ import ( "github.com/stretchr/testify/suite" configmock "github.com/goravel/framework/contracts/config/mocks" + logmock "github.com/goravel/framework/contracts/log/mocks" "github.com/goravel/framework/contracts/mail" queuecontract "github.com/goravel/framework/contracts/queue" "github.com/goravel/framework/queue" @@ -85,8 +86,9 @@ func (s *ApplicationTestSuite) TestSendMailWithFrom() { func (s *ApplicationTestSuite) TestQueueMail() { mockConfig := mockConfig(587, s.redisPort) + mockLog := &logmock.Log{} - queueFacade := queue.NewApplication(mockConfig) + queueFacade := queue.NewApplication(mockConfig, mockLog) queueFacade.Register([]queuecontract.Job{ NewSendMailJob(mockConfig), }) @@ -115,6 +117,7 @@ func (s *ApplicationTestSuite) TestQueueMail() { func mockConfig(mailPort, redisPort int) *configmock.Config { mockConfig := &configmock.Config{} mockConfig.On("GetString", "app.name").Return("goravel") + mockConfig.On("GetBool", "app.debug").Return(false) mockConfig.On("GetString", "queue.default").Return("redis") mockConfig.On("GetString", "queue.connections.sync.driver").Return("sync") mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis") diff --git a/queue/application.go b/queue/application.go index f7aaf081a..1a20943ed 100644 --- a/queue/application.go +++ b/queue/application.go @@ -2,17 +2,20 @@ package queue import ( configcontract "github.com/goravel/framework/contracts/config" + "github.com/goravel/framework/contracts/log" "github.com/goravel/framework/contracts/queue" ) type Application struct { config *Config jobs []queue.Job + log log.Log } -func NewApplication(config configcontract.Config) *Application { +func NewApplication(config configcontract.Config, log log.Log) *Application { return &Application{ config: NewConfig(config), + log: log, } } @@ -20,13 +23,13 @@ func (app *Application) Worker(args *queue.Args) queue.Worker { defaultConnection := app.config.DefaultConnection() if args == nil { - return NewWorker(app.config, 1, defaultConnection, app.jobs, app.config.Queue(defaultConnection, "")) + return NewWorker(app.config, app.log, 1, defaultConnection, app.jobs, app.config.Queue(defaultConnection, "")) } if args.Connection == "" { args.Connection = defaultConnection } - return NewWorker(app.config, args.Concurrent, args.Connection, app.jobs, app.config.Queue(args.Connection, args.Queue)) + return NewWorker(app.config, app.log, args.Concurrent, args.Connection, app.jobs, app.config.Queue(args.Connection, args.Queue)) } func (app *Application) Register(jobs []queue.Job) { @@ -38,9 +41,9 @@ func (app *Application) GetJobs() []queue.Job { } func (app *Application) Job(job queue.Job, args []queue.Arg) queue.Task { - return NewTask(app.config, job, args) + return NewTask(app.config, app.log, job, args) } func (app *Application) Chain(jobs []queue.Jobs) queue.Task { - return NewChainTask(app.config, jobs) + return NewChainTask(app.config, app.log, jobs) } diff --git a/queue/application_test.go b/queue/application_test.go index 628bf8562..8da5fb994 100644 --- a/queue/application_test.go +++ b/queue/application_test.go @@ -2,29 +2,34 @@ package queue import ( "context" + "errors" "log" "testing" "time" "github.com/ory/dockertest/v3" "github.com/spf13/cast" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/suite" configmock "github.com/goravel/framework/contracts/config/mocks" + logmock "github.com/goravel/framework/contracts/log/mocks" "github.com/goravel/framework/contracts/queue" - queuemock "github.com/goravel/framework/contracts/queue/mocks" "github.com/goravel/framework/support/carbon" testingdocker "github.com/goravel/framework/support/docker" ) var ( - testSyncJob = 0 - testAsyncJob = 0 - testDelayAsyncJob = 0 - testCustomAsyncJob = 0 - testErrorAsyncJob = 0 - testChainAsyncJob = 0 - testChainSyncJob = 0 + testSyncJob = 0 + testAsyncJob = 0 + testAsyncJobOfDisableDebug = 0 + testDelayAsyncJob = 0 + testCustomAsyncJob = 0 + testErrorAsyncJob = 0 + testChainAsyncJob = 0 + testChainSyncJob = 0 + testChainAsyncJobError = 0 + testChainSyncJobError = 0 ) type QueueTestSuite struct { @@ -32,7 +37,7 @@ type QueueTestSuite struct { app *Application redisResource *dockertest.Resource mockConfig *configmock.Config - mockQueue *queuemock.Queue + mockLog *logmock.Log } func TestQueueTestSuite(t *testing.T) { @@ -56,8 +61,8 @@ func TestQueueTestSuite(t *testing.T) { func (s *QueueTestSuite) SetupTest() { s.mockConfig = &configmock.Config{} - s.mockQueue = &queuemock.Queue{} - s.app = NewApplication(s.mockConfig) + s.mockLog = &logmock.Log{} + s.app = NewApplication(s.mockConfig, s.mockLog) } func (s *QueueTestSuite) TestSyncQueue() { @@ -69,9 +74,53 @@ func (s *QueueTestSuite) TestSyncQueue() { s.Equal(1, testSyncJob) } -func (s *QueueTestSuite) TestDefaultAsyncQueue() { +func (s *QueueTestSuite) TestDefaultAsyncQueue_EnableDebug() { + s.mockConfig.On("GetString", "queue.default").Return("redis").Twice() + s.mockConfig.On("GetString", "app.name").Return("goravel").Times(4) + s.mockConfig.On("GetBool", "app.debug").Return(true).Times(2) + s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Times(2) + s.mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis").Times(3) + s.mockConfig.On("GetString", "queue.connections.redis.connection").Return("default").Twice() + s.mockConfig.On("GetString", "database.redis.default.host").Return("localhost").Twice() + s.mockConfig.On("GetString", "database.redis.default.password").Return("").Twice() + s.mockConfig.On("GetInt", "database.redis.default.port").Return(cast.ToInt(s.redisResource.GetPort("6379/tcp"))).Twice() + s.mockConfig.On("GetInt", "database.redis.default.database").Return(0).Twice() + s.mockLog.On("Infof", "Launching a worker with the following settings:").Once() + s.mockLog.On("Infof", "- Broker: %s", "://").Once() + s.mockLog.On("Infof", "- DefaultQueue: %s", "goravel_queues:debug").Once() + s.mockLog.On("Infof", "- ResultBackend: %s", "://").Once() + s.mockLog.On("Info", "[*] Waiting for messages. To exit press CTRL+C").Once() + s.mockLog.On("Debugf", "Received new message: %s", mock.Anything).Once() + s.mockLog.On("Debugf", "Processed task %s. Results = %s", mock.Anything, mock.Anything).Once() + s.app.jobs = []queue.Job{&TestAsyncJob{}} + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + go func(ctx context.Context) { + s.Nil(s.app.Worker(&queue.Args{ + Queue: "debug", + }).Run()) + + for range ctx.Done() { + return + } + }(ctx) + time.Sleep(2 * time.Second) + s.Nil(s.app.Job(&TestAsyncJob{}, []queue.Arg{ + {Type: "string", Value: "TestDefaultAsyncQueue_EnableDebug"}, + {Type: "int", Value: 1}, + }).OnQueue("debug").Dispatch()) + time.Sleep(2 * time.Second) + s.Equal(1, testAsyncJob) + + s.mockConfig.AssertExpectations(s.T()) + s.mockLog.AssertExpectations(s.T()) +} + +func (s *QueueTestSuite) TestDefaultAsyncQueue_DisableDebug() { s.mockConfig.On("GetString", "queue.default").Return("redis").Twice() s.mockConfig.On("GetString", "app.name").Return("goravel").Times(3) + s.mockConfig.On("GetBool", "app.debug").Return(false).Times(2) s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Times(3) s.mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis").Times(3) s.mockConfig.On("GetString", "queue.connections.redis.connection").Return("default").Twice() @@ -79,7 +128,7 @@ func (s *QueueTestSuite) TestDefaultAsyncQueue() { s.mockConfig.On("GetString", "database.redis.default.password").Return("").Twice() s.mockConfig.On("GetInt", "database.redis.default.port").Return(cast.ToInt(s.redisResource.GetPort("6379/tcp"))).Twice() s.mockConfig.On("GetInt", "database.redis.default.database").Return(0).Twice() - s.app.jobs = []queue.Job{&TestAsyncJob{}} + s.app.jobs = []queue.Job{&TestAsyncJobOfDisableDebug{}} ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -91,20 +140,21 @@ func (s *QueueTestSuite) TestDefaultAsyncQueue() { } }(ctx) time.Sleep(2 * time.Second) - s.Nil(s.app.Job(&TestAsyncJob{}, []queue.Arg{ - {Type: "string", Value: "TestDefaultAsyncQueue"}, + s.Nil(s.app.Job(&TestAsyncJobOfDisableDebug{}, []queue.Arg{ + {Type: "string", Value: "TestDefaultAsyncQueue_DisableDebug"}, {Type: "int", Value: 1}, }).Dispatch()) time.Sleep(2 * time.Second) - s.Equal(1, testAsyncJob) + s.Equal(1, testAsyncJobOfDisableDebug) s.mockConfig.AssertExpectations(s.T()) - s.mockQueue.AssertExpectations(s.T()) + s.mockLog.AssertExpectations(s.T()) } func (s *QueueTestSuite) TestDelayAsyncQueue() { s.mockConfig.On("GetString", "queue.default").Return("redis").Times(2) s.mockConfig.On("GetString", "app.name").Return("goravel").Times(4) + s.mockConfig.On("GetBool", "app.debug").Return(false).Times(2) s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Twice() s.mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis").Times(3) s.mockConfig.On("GetString", "queue.connections.redis.connection").Return("default").Twice() @@ -136,12 +186,12 @@ func (s *QueueTestSuite) TestDelayAsyncQueue() { s.Equal(1, testDelayAsyncJob) s.mockConfig.AssertExpectations(s.T()) - s.mockQueue.AssertExpectations(s.T()) } func (s *QueueTestSuite) TestCustomAsyncQueue() { s.mockConfig.On("GetString", "queue.default").Return("redis").Twice() s.mockConfig.On("GetString", "app.name").Return("goravel").Times(4) + s.mockConfig.On("GetBool", "app.debug").Return(false).Times(2) s.mockConfig.On("GetString", "queue.connections.custom.queue", "default").Return("default").Twice() s.mockConfig.On("GetString", "queue.connections.custom.driver").Return("redis").Times(3) s.mockConfig.On("GetString", "queue.connections.custom.connection").Return("default").Twice() @@ -173,12 +223,12 @@ func (s *QueueTestSuite) TestCustomAsyncQueue() { s.Equal(1, testCustomAsyncJob) s.mockConfig.AssertExpectations(s.T()) - s.mockQueue.AssertExpectations(s.T()) } func (s *QueueTestSuite) TestErrorAsyncQueue() { s.mockConfig.On("GetString", "queue.default").Return("redis").Twice() s.mockConfig.On("GetString", "app.name").Return("goravel").Times(4) + s.mockConfig.On("GetBool", "app.debug").Return(false).Times(2) s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Twice() s.mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis").Times(3) s.mockConfig.On("GetString", "queue.connections.redis.connection").Return("default").Twice() @@ -208,12 +258,12 @@ func (s *QueueTestSuite) TestErrorAsyncQueue() { s.Equal(0, testErrorAsyncJob) s.mockConfig.AssertExpectations(s.T()) - s.mockQueue.AssertExpectations(s.T()) } func (s *QueueTestSuite) TestChainAsyncQueue() { s.mockConfig.On("GetString", "queue.default").Return("redis").Times(2) s.mockConfig.On("GetString", "app.name").Return("goravel").Times(4) + s.mockConfig.On("GetBool", "app.debug").Return(false).Times(2) s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Twice() s.mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis").Times(3) s.mockConfig.On("GetString", "queue.connections.redis.connection").Return("default").Twice() @@ -260,6 +310,54 @@ func (s *QueueTestSuite) TestChainAsyncQueue() { s.mockConfig.AssertExpectations(s.T()) } +func (s *QueueTestSuite) TestChainAsyncQueue_Error() { + s.mockConfig.On("GetString", "queue.default").Return("redis").Times(2) + s.mockConfig.On("GetString", "app.name").Return("goravel").Times(4) + s.mockConfig.On("GetBool", "app.debug").Return(false).Times(2) + s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Twice() + s.mockConfig.On("GetString", "queue.connections.redis.driver").Return("redis").Times(3) + s.mockConfig.On("GetString", "queue.connections.redis.connection").Return("default").Twice() + s.mockConfig.On("GetString", "database.redis.default.host").Return("localhost").Twice() + s.mockConfig.On("GetString", "database.redis.default.password").Return("").Twice() + s.mockConfig.On("GetInt", "database.redis.default.port").Return(cast.ToInt(s.redisResource.GetPort("6379/tcp"))).Twice() + s.mockConfig.On("GetInt", "database.redis.default.database").Return(0).Twice() + s.mockLog.On("Errorf", "Failed processing task %s. Error = %v", mock.Anything, errors.New("error")).Once() + s.app.jobs = []queue.Job{&TestChainAsyncJob{}, &TestChainSyncJob{}} + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + go func(ctx context.Context) { + s.Nil(s.app.Worker(&queue.Args{ + Queue: "chain", + }).Run()) + + for range ctx.Done() { + return + } + }(ctx) + + time.Sleep(2 * time.Second) + s.Nil(s.app.Chain([]queue.Jobs{ + { + Job: &TestChainAsyncJob{}, + Args: []queue.Arg{ + {Type: "bool", Value: true}, + }, + }, + { + Job: &TestChainSyncJob{}, + Args: []queue.Arg{}, + }, + }).OnQueue("chain").Dispatch()) + + time.Sleep(2 * time.Second) + s.Equal(1, testChainAsyncJobError) + s.Equal(0, testChainSyncJobError) + + s.mockConfig.AssertExpectations(s.T()) + s.mockLog.AssertExpectations(s.T()) +} + type TestAsyncJob struct { } @@ -275,6 +373,21 @@ func (receiver *TestAsyncJob) Handle(args ...any) error { return nil } +type TestAsyncJobOfDisableDebug struct { +} + +// Signature The name and signature of the job. +func (receiver *TestAsyncJobOfDisableDebug) Signature() string { + return "test_async_job_of_disable_debug" +} + +// Handle Execute the job. +func (receiver *TestAsyncJobOfDisableDebug) Handle(args ...any) error { + testAsyncJobOfDisableDebug++ + + return nil +} + type TestDelayAsyncJob struct { } @@ -345,6 +458,12 @@ func (receiver *TestChainAsyncJob) Signature() string { // Handle Execute the job. func (receiver *TestChainAsyncJob) Handle(args ...any) error { + if len(args) > 0 && cast.ToBool(args[0]) { + testChainAsyncJobError++ + + return errors.New("error") + } + testChainAsyncJob++ return nil diff --git a/queue/log.go b/queue/log.go new file mode 100644 index 000000000..898b5bb54 --- /dev/null +++ b/queue/log.go @@ -0,0 +1,257 @@ +package queue + +import ( + "github.com/goravel/framework/contracts/log" +) + +type Debug struct { + debug bool + log log.Log +} + +func NewDebug(debug bool, log log.Log) *Debug { + return &Debug{ + debug: debug, + log: log, + } +} + +func (r *Debug) Print(args ...any) { + if r.debug { + r.log.Debug(args...) + } +} + +func (r *Debug) Printf(format string, args ...any) { + if r.debug { + r.log.Debugf(format, args...) + } +} + +func (r *Debug) Println(args ...any) { + if r.debug { + r.log.Debug(args...) + } +} + +func (r *Debug) Fatal(args ...any) { + r.log.Error(args...) +} + +func (r *Debug) Fatalf(format string, args ...any) { + r.log.Errorf(format, args...) +} + +func (r *Debug) Fatalln(args ...any) { + r.log.Error(args...) +} + +func (r *Debug) Panic(args ...any) { + r.log.Panic(args...) +} + +func (r *Debug) Panicf(format string, args ...any) { + r.log.Panicf(format, args...) +} + +func (r *Debug) Panicln(args ...any) { + r.log.Panic(args...) +} + +type Info struct { + debug bool + log log.Log +} + +func NewInfo(debug bool, log log.Log) *Info { + return &Info{ + debug: debug, + log: log, + } +} + +func (r *Info) Print(args ...any) { + if r.debug { + r.log.Info(args...) + } +} + +func (r *Info) Printf(format string, args ...any) { + if r.debug { + r.log.Infof(format, args...) + } +} + +func (r *Info) Println(args ...any) { + if r.debug { + r.log.Info(args...) + } +} + +func (r *Info) Fatal(args ...any) { + r.log.Error(args...) +} + +func (r *Info) Fatalf(format string, args ...any) { + r.log.Errorf(format, args...) +} + +func (r *Info) Fatalln(args ...any) { + r.log.Error(args...) +} + +func (r *Info) Panic(args ...any) { + r.log.Panic(args...) +} + +func (r *Info) Panicf(format string, args ...any) { + r.log.Panicf(format, args...) +} + +func (r *Info) Panicln(args ...any) { + r.log.Panic(args...) +} + +type Warning struct { + debug bool + log log.Log +} + +func NewWarning(debug bool, log log.Log) *Warning { + return &Warning{ + debug: debug, + log: log, + } +} + +func (r *Warning) Print(args ...any) { + r.log.Warning(args...) +} + +func (r *Warning) Printf(format string, args ...any) { + r.log.Warningf(format, args...) +} + +func (r *Warning) Println(args ...any) { + r.log.Warning(args...) +} + +func (r *Warning) Fatal(args ...any) { + r.log.Error(args...) +} + +func (r *Warning) Fatalf(format string, args ...any) { + r.log.Errorf(format, args...) +} + +func (r *Warning) Fatalln(args ...any) { + r.log.Error(args...) +} + +func (r *Warning) Panic(args ...any) { + r.log.Panic(args...) +} + +func (r *Warning) Panicf(format string, args ...any) { + r.log.Panicf(format, args...) +} + +func (r *Warning) Panicln(args ...any) { + r.log.Panic(args...) +} + +type Error struct { + debug bool + log log.Log +} + +func NewError(debug bool, log log.Log) *Error { + return &Error{ + debug: debug, + log: log, + } +} + +func (r *Error) Print(args ...any) { + r.log.Error(args...) +} + +func (r *Error) Printf(format string, args ...any) { + r.log.Errorf(format, args...) +} + +func (r *Error) Println(args ...any) { + r.log.Error(args...) +} + +func (r *Error) Fatal(args ...any) { + r.log.Error(args...) +} + +func (r *Error) Fatalf(format string, args ...any) { + r.log.Errorf(format, args...) +} + +func (r *Error) Fatalln(args ...any) { + r.log.Error(args...) +} + +func (r *Error) Panic(args ...any) { + r.log.Panic(args...) +} + +func (r *Error) Panicf(format string, args ...any) { + r.log.Panicf(format, args...) +} + +func (r *Error) Panicln(args ...any) { + r.log.Panic(args...) +} + +type Fatal struct { + debug bool + log log.Log +} + +func NewFatal(debug bool, log log.Log) *Fatal { + return &Fatal{ + debug: debug, + log: log, + } +} + +func (r *Fatal) Print(args ...any) { + r.log.Fatal(args...) +} + +func (r *Fatal) Printf(format string, args ...any) { + r.log.Fatalf(format, args...) +} + +func (r *Fatal) Println(args ...any) { + r.log.Fatal(args...) +} + +func (r *Fatal) Fatal(args ...any) { + r.log.Fatal(args...) +} + +func (r *Fatal) Fatalf(format string, args ...any) { + r.log.Fatalf(format, args...) +} + +func (r *Fatal) Fatalln(args ...any) { + r.log.Fatal(args...) +} + +func (r *Fatal) Panic(args ...any) { + r.log.Panic(args...) +} + +func (r *Fatal) Panicf(format string, args ...any) { + r.log.Panicf(format, args...) +} + +func (r *Fatal) Panicln(args ...any) { + r.log.Panic(args...) +} diff --git a/queue/machinery.go b/queue/machinery.go index f77001c1f..1b6c5d90d 100644 --- a/queue/machinery.go +++ b/queue/machinery.go @@ -8,15 +8,19 @@ import ( redisbroker "github.com/RichardKnop/machinery/v2/brokers/redis" "github.com/RichardKnop/machinery/v2/config" "github.com/RichardKnop/machinery/v2/locks/eager" + "github.com/RichardKnop/machinery/v2/log" "github.com/gookit/color" + + logcontract "github.com/goravel/framework/contracts/log" ) type Machinery struct { config *Config + log logcontract.Log } -func NewMachinery(config *Config) *Machinery { - return &Machinery{config: config} +func NewMachinery(config *Config, log logcontract.Log) *Machinery { + return &Machinery{config: config, log: log} } func (m *Machinery) Server(connection string, queue string) (*machinery.Server, error) { @@ -49,5 +53,12 @@ func (m *Machinery) redisServer(connection string, queue string) *machinery.Serv backend := redisbackend.NewGR(cnf, []string{redisConfig}, database) lock := eager.New() + debug := m.config.config.GetBool("app.debug") + log.DEBUG = NewDebug(debug, m.log) + log.INFO = NewInfo(debug, m.log) + log.WARNING = NewWarning(debug, m.log) + log.ERROR = NewError(debug, m.log) + log.FATAL = NewFatal(debug, m.log) + return machinery.NewServer(cnf, broker, backend, lock) } diff --git a/queue/machinery_test.go b/queue/machinery_test.go index cfcddb8f1..9721622a4 100644 --- a/queue/machinery_test.go +++ b/queue/machinery_test.go @@ -6,11 +6,13 @@ import ( "github.com/stretchr/testify/suite" configmock "github.com/goravel/framework/contracts/config/mocks" + logmock "github.com/goravel/framework/contracts/log/mocks" ) type MachineryTestSuite struct { suite.Suite mockConfig *configmock.Config + mockLog *logmock.Log machinery *Machinery } @@ -20,7 +22,8 @@ func TestMachineryTestSuite(t *testing.T) { func (s *MachineryTestSuite) SetupTest() { s.mockConfig = &configmock.Config{} - s.machinery = NewMachinery(NewConfig(s.mockConfig)) + s.mockLog = &logmock.Log{} + s.machinery = NewMachinery(NewConfig(s.mockConfig), s.mockLog) } func (s *MachineryTestSuite) TestServer() { @@ -51,6 +54,7 @@ func (s *MachineryTestSuite) TestServer() { s.mockConfig.On("GetInt", "database.redis.default.database").Return(0).Once() s.mockConfig.On("GetString", "queue.connections.redis.queue", "default").Return("default").Once() s.mockConfig.On("GetString", "app.name").Return("goravel").Once() + s.mockConfig.On("GetBool", "app.debug").Return(true).Once() }, expectServer: true, }, diff --git a/queue/service_provider.go b/queue/service_provider.go index 083fe2d5d..dfebea0a1 100644 --- a/queue/service_provider.go +++ b/queue/service_provider.go @@ -13,7 +13,7 @@ type ServiceProvider struct { func (receiver *ServiceProvider) Register(app foundation.Application) { app.Singleton(Binding, func(app foundation.Application) (any, error) { - return NewApplication(app.MakeConfig()), nil + return NewApplication(app.MakeConfig(), app.MakeLog()), nil }) } diff --git a/queue/task.go b/queue/task.go index d98c864b0..b9eedee81 100644 --- a/queue/task.go +++ b/queue/task.go @@ -7,6 +7,7 @@ import ( "github.com/RichardKnop/machinery/v2" "github.com/RichardKnop/machinery/v2/tasks" + "github.com/goravel/framework/contracts/log" "github.com/goravel/framework/contracts/queue" ) @@ -21,11 +22,11 @@ type Task struct { server *machinery.Server } -func NewTask(config *Config, job queue.Job, args []queue.Arg) *Task { +func NewTask(config *Config, log log.Log, job queue.Job, args []queue.Arg) *Task { return &Task{ config: config, connection: config.DefaultConnection(), - machinery: NewMachinery(config), + machinery: NewMachinery(config, log), jobs: []queue.Jobs{ { Job: job, @@ -35,12 +36,12 @@ func NewTask(config *Config, job queue.Job, args []queue.Arg) *Task { } } -func NewChainTask(config *Config, jobs []queue.Jobs) *Task { +func NewChainTask(config *Config, log log.Log, jobs []queue.Jobs) *Task { return &Task{ config: config, connection: config.DefaultConnection(), chain: true, - machinery: NewMachinery(config), + machinery: NewMachinery(config, log), jobs: jobs, } } @@ -68,13 +69,7 @@ func (receiver *Task) Dispatch() error { receiver.server = server if receiver.chain { - for _, job := range receiver.jobs { - if err := receiver.handleAsync(job.Job, job.Args); err != nil { - return err - } - } - - return nil + return receiver.handleChain(receiver.jobs) } else { job := receiver.jobs[0] @@ -110,6 +105,34 @@ func (receiver *Task) OnQueue(queue string) queue.Task { return receiver } +func (receiver *Task) handleChain(jobs []queue.Jobs) error { + var signatures []*tasks.Signature + for _, job := range jobs { + var realArgs []tasks.Arg + for _, arg := range job.Args { + realArgs = append(realArgs, tasks.Arg{ + Type: arg.Type, + Value: arg.Value, + }) + } + + signatures = append(signatures, &tasks.Signature{ + Name: job.Job.Signature(), + Args: realArgs, + ETA: receiver.delay, + }) + } + + chain, err := tasks.NewChain(signatures...) + if err != nil { + return err + } + + _, err = receiver.server.SendChain(chain) + + return err +} + func (receiver *Task) handleAsync(job queue.Job, args []queue.Arg) error { var realArgs []tasks.Arg for _, arg := range args { diff --git a/queue/worker.go b/queue/worker.go index fe6e2cd2f..38c6a3f9c 100644 --- a/queue/worker.go +++ b/queue/worker.go @@ -1,6 +1,7 @@ package queue import ( + "github.com/goravel/framework/contracts/log" "github.com/goravel/framework/contracts/queue" ) @@ -15,11 +16,11 @@ type Worker struct { queue string } -func NewWorker(config *Config, concurrent int, connection string, jobs []queue.Job, queue string) *Worker { +func NewWorker(config *Config, log log.Log, concurrent int, connection string, jobs []queue.Job, queue string) *Worker { return &Worker{ concurrent: concurrent, connection: connection, - machinery: NewMachinery(config), + machinery: NewMachinery(config, log), jobs: jobs, queue: queue, } diff --git a/support/carbon/json_test.go b/support/carbon/json_test.go index 77b0b793d..8809366ce 100644 --- a/support/carbon/json_test.go +++ b/support/carbon/json_test.go @@ -1,10 +1,10 @@ package carbon import ( + "encoding/json" "fmt" "testing" - "github.com/bytedance/sonic" "github.com/stretchr/testify/assert" ) @@ -44,7 +44,7 @@ func TestMarshalJSON(t *testing.T) { CreatedAt3: TimestampMicro{Parse("2025-08-05 13:14:15.999999")}, CreatedAt4: TimestampNano{Parse("2025-08-05 13:14:15.999999999")}, } - data, err := sonic.Marshal(&person) + data, err := json.Marshal(&person) assert.Nil(t, err) fmt.Printf("Person output by json:\n%s\n", data) } @@ -67,7 +67,7 @@ func TestUnmarshalJSON(t *testing.T) { "created_at4": 1596604455999999999 }` - err := sonic.Unmarshal([]byte(str), &person) + err := json.Unmarshal([]byte(str), &person) assert.Nil(t, err) fmt.Printf("Json string parse to person:\n%+v\n", person) } @@ -89,7 +89,7 @@ func TestErrorJson(t *testing.T) { "created_at3": 0, "created_at4": 0 }` - err := sonic.Unmarshal([]byte(str), &person) + err := json.Unmarshal([]byte(str), &person) assert.NotNil(t, err) fmt.Printf("Json string parse to person:\n%+v\n", person) } diff --git a/support/constant.go b/support/constant.go index 43fa04fe4..1a86c3225 100644 --- a/support/constant.go +++ b/support/constant.go @@ -1,6 +1,6 @@ package support -const Version string = "v1.13.0" +const Version string = "v1.13.8" const ( EnvRuntime = "runtime" @@ -9,8 +9,9 @@ const ( ) var ( - Env = EnvRuntime - EnvPath = ".env" - RelativePath string - RootPath string + Env = EnvRuntime + EnvPath = ".env" + IsKeyGenerateCommand = false + RelativePath string + RootPath string ) diff --git a/support/database/database.go b/support/database/database.go index f29786c74..7fd79cd66 100644 --- a/support/database/database.go +++ b/support/database/database.go @@ -30,7 +30,7 @@ func GetIDByReflect(t reflect.Type, v reflect.Value) any { if t.Field(i).Name == "Model" && v.Field(i).Type().Kind() == reflect.Struct { structField := v.Field(i).Type() for j := 0; j < structField.NumField(); j++ { - if !structField.Field(i).IsExported() { + if !structField.Field(j).IsExported() { continue } if strings.Contains(structField.Field(j).Tag.Get("gorm"), "primaryKey") { diff --git a/support/env/env.go b/support/env/env.go new file mode 100644 index 000000000..076f547c2 --- /dev/null +++ b/support/env/env.go @@ -0,0 +1,39 @@ +package env + +import "runtime" + +// IsWindows returns whether the current operating system is Windows. +// IsWindows 返回当前操作系统是否为 Windows。 +func IsWindows() bool { + return runtime.GOOS == "windows" +} + +// IsLinux returns whether the current operating system is Linux. +// IsLinux 返回当前操作系统是否为 Linux。 +func IsLinux() bool { + return runtime.GOOS == "linux" +} + +// IsDarwin returns whether the current operating system is Darwin. +// IsDarwin 返回当前操作系统是否为 Darwin。 +func IsDarwin() bool { + return runtime.GOOS == "darwin" +} + +// IsArm returns whether the current CPU architecture is ARM. +// IsArm 返回当前 CPU 架构是否为 ARM。 +func IsArm() bool { + return runtime.GOARCH == "arm" || runtime.GOARCH == "arm64" +} + +// IsX86 returns whether the current CPU architecture is X86. +// IsX86 返回当前 CPU 架构是否为 X86。 +func IsX86() bool { + return runtime.GOARCH == "386" || runtime.GOARCH == "amd64" +} + +// Is64Bit returns whether the current CPU architecture is 64-bit. +// Is64Bit 返回当前 CPU 架构是否为 64 位。 +func Is64Bit() bool { + return runtime.GOARCH == "amd64" || runtime.GOARCH == "arm64" +} diff --git a/support/file/file.go b/support/file/file.go index 69bda1096..34e548c55 100644 --- a/support/file/file.go +++ b/support/file/file.go @@ -35,11 +35,7 @@ func Create(file string, content string) error { if err != nil { return err } - defer func() { - if closeErr := f.Close(); closeErr != nil && err == nil { - err = closeErr - } - }() + defer f.Close() if _, err = f.WriteString(content); err != nil { return err @@ -114,6 +110,7 @@ func Size(file string) (int64, error) { if err != nil { return 0, err } + defer fileInfo.Close() fi, err := fileInfo.Stat() if err != nil { diff --git a/support/json/json.go b/support/json/json.go new file mode 100644 index 000000000..0f57768fc --- /dev/null +++ b/support/json/json.go @@ -0,0 +1,32 @@ +//go:build !amd64 + +package json + +import ( + "encoding/json" +) + +// Marshal is a wrapper of json.Marshal. +// Marshal 是 json.Marshal 的包装器。 +func Marshal(v any) ([]byte, error) { + return json.Marshal(v) +} + +// Unmarshal is a wrapper of json.Unmarshal. +// Unmarshal 是 json.Unmarshal 的包装器。 +func Unmarshal(data []byte, v any) error { + return json.Unmarshal(data, v) +} + +// MarshalString is a wrapper of json.Marshal. +// MarshalString 是 json.Marshal 的包装器。 +func MarshalString(v any) (string, error) { + s, err := json.Marshal(v) + return string(s), err +} + +// UnmarshalString is a wrapper of json.Unmarshal. +// UnmarshalString 是 json.Unmarshal 的包装器。 +func UnmarshalString(data string, v any) error { + return json.Unmarshal([]byte(data), v) +} diff --git a/support/json/json_test.go b/support/json/json_test.go new file mode 100644 index 000000000..a0cc218af --- /dev/null +++ b/support/json/json_test.go @@ -0,0 +1,33 @@ +package json + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestMarshal(t *testing.T) { + json, err := Marshal(map[string]int{"a": 1}) + assert.Equal(t, []byte(`{"a":1}`), json) + assert.Nil(t, err) +} + +func TestUnmarshal(t *testing.T) { + var m map[string]int + err := Unmarshal([]byte(`{"a":1}`), &m) + assert.Equal(t, map[string]int{"a": 1}, m) + assert.Nil(t, err) +} + +func TestMarshalString(t *testing.T) { + json, err := MarshalString(map[string]int{"a": 1}) + assert.Equal(t, `{"a":1}`, json) + assert.Nil(t, err) +} + +func TestUnmarshalString(t *testing.T) { + var m map[string]int + err := UnmarshalString(`{"a":1}`, &m) + assert.Equal(t, map[string]int{"a": 1}, m) + assert.Nil(t, err) +} diff --git a/support/json/json_x86.go b/support/json/json_x86.go new file mode 100644 index 000000000..818a4578f --- /dev/null +++ b/support/json/json_x86.go @@ -0,0 +1,31 @@ +//go:build amd64 + +package json + +import ( + "github.com/bytedance/sonic" +) + +// Marshal is a wrapper of sonic.Marshal. +// Marshal 是 sonic.Marshal 的包装器。 +func Marshal(v any) ([]byte, error) { + return sonic.Marshal(v) +} + +// Unmarshal is a wrapper of sonic.Unmarshal. +// Unmarshal 是 sonic.Unmarshal 的包装器。 +func Unmarshal(data []byte, v any) error { + return sonic.Unmarshal(data, v) +} + +// MarshalString is a wrapper of sonic.MarshalString. +// MarshalString 是 sonic.MarshalString 的包装器。 +func MarshalString(v any) (string, error) { + return sonic.MarshalString(v) +} + +// UnmarshalString is a wrapper of sonic.UnmarshalString. +// UnmarshalString 是 sonic.UnmarshalString 的包装器。 +func UnmarshalString(data string, v any) error { + return sonic.UnmarshalString(data, v) +} diff --git a/support/str/str.go b/support/str/str.go index 39b4d3ae1..37f3b6243 100644 --- a/support/str/str.go +++ b/support/str/str.go @@ -3,11 +3,927 @@ package str import ( "bytes" "crypto/rand" + "encoding/json" + "path/filepath" + "regexp" "strconv" "strings" "unicode" + "unicode/utf8" + + "golang.org/x/exp/constraints" + "golang.org/x/text/cases" + "golang.org/x/text/language" ) +type String struct { + value string +} + +// ExcerptOption is the option for Excerpt method +type ExcerptOption struct { + Radius int + Omission string +} + +// Of creates a new String instance with the given value. +func Of(value string) *String { + return &String{value: value} +} + +// After returns a new String instance with the substring after the first occurrence of the specified search string. +func (s *String) After(search string) *String { + if search == "" { + return s + } + index := strings.Index(s.value, search) + if index != -1 { + s.value = s.value[index+len(search):] + } + return s +} + +// AfterLast returns the String instance with the substring after the last occurrence of the specified search string. +func (s *String) AfterLast(search string) *String { + index := strings.LastIndex(s.value, search) + if index != -1 { + s.value = s.value[index+len(search):] + } + + return s +} + +// Append appends one or more strings to the current string. +func (s *String) Append(values ...string) *String { + s.value += strings.Join(values, "") + return s +} + +// Basename returns the String instance with the basename of the current file path string, +// and trims the suffix based on the parameter(optional). +func (s *String) Basename(suffix ...string) *String { + s.value = filepath.Base(s.value) + if len(suffix) > 0 && suffix[0] != "" { + s.value = strings.TrimSuffix(s.value, suffix[0]) + } + return s +} + +// Before returns the String instance with the substring before the first occurrence of the specified search string. +func (s *String) Before(search string) *String { + index := strings.Index(s.value, search) + if index != -1 { + s.value = s.value[:index] + } + + return s +} + +// BeforeLast returns the String instance with the substring before the last occurrence of the specified search string. +func (s *String) BeforeLast(search string) *String { + index := strings.LastIndex(s.value, search) + if index != -1 { + s.value = s.value[:index] + } + + return s +} + +// Between returns the String instance with the substring between the given start and end strings. +func (s *String) Between(start, end string) *String { + if start == "" || end == "" { + return s + } + return s.After(start).BeforeLast(end) +} + +// BetweenFirst returns the String instance with the substring between the first occurrence of the given start string and the given end string. +func (s *String) BetweenFirst(start, end string) *String { + if start == "" || end == "" { + return s + } + return s.Before(end).After(start) +} + +// Camel returns the String instance in camel case. +func (s *String) Camel() *String { + return s.Studly().LcFirst() +} + +// CharAt returns the character at the specified index. +func (s *String) CharAt(index int) string { + length := len(s.value) + // return zero string when char doesn't exists + if index < 0 && index < -length || index > length-1 { + return "" + } + return Substr(s.value, index, 1) +} + +// Contains returns true if the string contains the given value or any of the values. +func (s *String) Contains(values ...string) bool { + for _, value := range values { + if value != "" && strings.Contains(s.value, value) { + return true + } + } + + return false +} + +// ContainsAll returns true if the string contains all of the given values. +func (s *String) ContainsAll(values ...string) bool { + for _, value := range values { + if !strings.Contains(s.value, value) { + return false + } + } + + return true +} + +// Dirname returns the String instance with the directory name of the current file path string. +func (s *String) Dirname(levels ...int) *String { + defaultLevels := 1 + if len(levels) > 0 { + defaultLevels = levels[0] + } + + dir := s.value + for i := 0; i < defaultLevels; i++ { + dir = filepath.Dir(dir) + } + + s.value = dir + return s +} + +// EndsWith returns true if the string ends with the given value or any of the values. +func (s *String) EndsWith(values ...string) bool { + for _, value := range values { + if value != "" && strings.HasSuffix(s.value, value) { + return true + } + } + + return false +} + +// Exactly returns true if the string is exactly the given value. +func (s *String) Exactly(value string) bool { + return s.value == value +} + +// Excerpt returns the String instance truncated to the given length. +func (s *String) Excerpt(phrase string, options ...ExcerptOption) *String { + defaultOptions := ExcerptOption{ + Radius: 100, + Omission: "...", + } + + if len(options) > 0 { + if options[0].Radius != 0 { + defaultOptions.Radius = options[0].Radius + } + if options[0].Omission != "" { + defaultOptions.Omission = options[0].Omission + } + } + + radius := maximum(0, defaultOptions.Radius) + omission := defaultOptions.Omission + + regex := regexp.MustCompile(`(.*?)(` + regexp.QuoteMeta(phrase) + `)(.*)`) + matches := regex.FindStringSubmatch(s.value) + + if len(matches) == 0 { + return s + } + + start := strings.TrimRight(matches[1], "") + end := strings.TrimLeft(matches[3], "") + + end = Of(Substr(end, 0, radius)).LTrim(""). + Unless(func(s *String) bool { + return s.Exactly(end) + }, func(s *String) *String { + return s.Append(omission) + }).String() + + s.value = Of(Substr(start, maximum(len(start)-radius, 0), radius)).LTrim(""). + Unless(func(s *String) bool { + return s.Exactly(start) + }, func(s *String) *String { + return s.Prepend(omission) + }).Append(matches[2], end).String() + + return s +} + +// Explode splits the string by given delimiter string. +func (s *String) Explode(delimiter string, limit ...int) []string { + defaultLimit := 1 + isLimitSet := false + if len(limit) > 0 && limit[0] != 0 { + defaultLimit = limit[0] + isLimitSet = true + } + tempExplode := strings.Split(s.value, delimiter) + if !isLimitSet || len(tempExplode) <= defaultLimit { + return tempExplode + } + + if defaultLimit > 0 { + return append(tempExplode[:defaultLimit-1], strings.Join(tempExplode[defaultLimit-1:], delimiter)) + } + + if defaultLimit < 0 && len(tempExplode) <= -defaultLimit { + return []string{} + } + + return tempExplode[:len(tempExplode)+defaultLimit] +} + +// Finish returns the String instance with the given value appended. +// If the given value already ends with the suffix, it will not be added twice. +func (s *String) Finish(value string) *String { + quoted := regexp.QuoteMeta(value) + reg := regexp.MustCompile("(?:" + quoted + ")+$") + s.value = reg.ReplaceAllString(s.value, "") + value + return s +} + +// Headline returns the String instance in headline case. +func (s *String) Headline() *String { + parts := s.Explode(" ") + + if len(parts) > 1 { + return s.Title() + } + + parts = Of(strings.Join(parts, "_")).Studly().UcSplit() + collapsed := Of(strings.Join(parts, "_")). + Replace("-", "_"). + Replace(" ", "_"). + Replace("_", "_").Explode("_") + + s.value = strings.Join(collapsed, " ") + return s +} + +// Is returns true if the string matches any of the given patterns. +func (s *String) Is(patterns ...string) bool { + for _, pattern := range patterns { + if pattern == s.value { + return true + } + + // Escape special characters in the pattern + pattern = regexp.QuoteMeta(pattern) + + // Replace asterisks with regular expression wildcards + pattern = strings.ReplaceAll(pattern, `\*`, ".*") + + // Create a regular expression pattern for matching + regexPattern := "^" + pattern + "$" + + // Compile the regular expression + regex := regexp.MustCompile(regexPattern) + + // Check if the value matches the pattern + if regex.MatchString(s.value) { + return true + } + } + + return false +} + +// IsEmpty returns true if the string is empty. +func (s *String) IsEmpty() bool { + return s.value == "" +} + +// IsNotEmpty returns true if the string is not empty. +func (s *String) IsNotEmpty() bool { + return !s.IsEmpty() +} + +// IsAscii returns true if the string contains only ASCII characters. +func (s *String) IsAscii() bool { + return s.IsMatch(`^[\x00-\x7F]+$`) +} + +// IsMap returns true if the string is a valid Map. +func (s *String) IsMap() bool { + var obj map[string]interface{} + return json.Unmarshal([]byte(s.value), &obj) == nil +} + +// IsSlice returns true if the string is a valid Slice. +func (s *String) IsSlice() bool { + var arr []interface{} + return json.Unmarshal([]byte(s.value), &arr) == nil +} + +// IsUlid returns true if the string is a valid ULID. +func (s *String) IsUlid() bool { + return s.IsMatch(`^[0-9A-Z]{26}$`) +} + +// IsUuid returns true if the string is a valid UUID. +func (s *String) IsUuid() bool { + return s.IsMatch(`(?i)^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`) +} + +// Kebab returns the String instance in kebab case. +func (s *String) Kebab() *String { + return s.Snake("-") +} + +// LcFirst returns the String instance with the first character lowercased. +func (s *String) LcFirst() *String { + if s.Length() == 0 { + return s + } + s.value = strings.ToLower(Substr(s.value, 0, 1)) + Substr(s.value, 1) + return s +} + +// Length returns the length of the string. +func (s *String) Length() int { + return utf8.RuneCountInString(s.value) +} + +// Limit returns the String instance truncated to the given length. +func (s *String) Limit(limit int, end ...string) *String { + defaultEnd := "..." + if len(end) > 0 { + defaultEnd = end[0] + } + + if s.Length() <= limit { + return s + } + s.value = Substr(s.value, 0, limit) + defaultEnd + return s +} + +// Lower returns the String instance in lower case. +func (s *String) Lower() *String { + s.value = strings.ToLower(s.value) + return s +} + +// Ltrim returns the String instance with the leftmost occurrence of the given value removed. +func (s *String) LTrim(characters ...string) *String { + if len(characters) == 0 { + s.value = strings.TrimLeft(s.value, " ") + return s + } + + s.value = strings.TrimLeft(s.value, characters[0]) + return s +} + +// Mask returns the String instance with the given character masking the specified number of characters. +func (s *String) Mask(character string, index int, length ...int) *String { + // Check if the character is empty, if so, return the original string. + if character == "" { + return s + } + + segment := Substr(s.value, index, length...) + + // Check if the segment is empty, if so, return the original string. + if segment == "" { + return s + } + + strLen := utf8.RuneCountInString(s.value) + startIndex := index + + // Check if the start index is out of bounds. + if index < 0 { + if index < -strLen { + startIndex = 0 + } else { + startIndex = strLen + index + } + } + + start := Substr(s.value, 0, startIndex) + segmentLen := utf8.RuneCountInString(segment) + end := Substr(s.value, startIndex+segmentLen) + + s.value = start + strings.Repeat(Substr(character, 0, 1), segmentLen) + end + return s +} + +// Match returns the String instance with the first occurrence of the given pattern. +func (s *String) Match(pattern string) *String { + if pattern == "" { + return s + } + reg := regexp.MustCompile(pattern) + s.value = reg.FindString(s.value) + return s +} + +// MatchAll returns all matches for the given regular expression. +func (s *String) MatchAll(pattern string) []string { + if pattern == "" { + return []string{s.value} + } + reg := regexp.MustCompile(pattern) + return reg.FindAllString(s.value, -1) +} + +// IsMatch returns true if the string matches any of the given patterns. +func (s *String) IsMatch(patterns ...string) bool { + for _, pattern := range patterns { + reg := regexp.MustCompile(pattern) + if reg.MatchString(s.value) { + return true + } + } + + return false +} + +// NewLine appends one or more new lines to the current string. +func (s *String) NewLine(count ...int) *String { + if len(count) == 0 { + s.value += "\n" + return s + } + + s.value += strings.Repeat("\n", count[0]) + return s +} + +// PadBoth returns the String instance padded to the left and right sides of the given length. +func (s *String) PadBoth(length int, pad ...string) *String { + defaultPad := " " + if len(pad) > 0 { + defaultPad = pad[0] + } + short := maximum(0, length-s.Length()) + left := short / 2 + right := short/2 + short%2 + + s.value = Substr(strings.Repeat(defaultPad, left), 0, left) + s.value + Substr(strings.Repeat(defaultPad, right), 0, right) + + return s +} + +// PadLeft returns the String instance padded to the left side of the given length. +func (s *String) PadLeft(length int, pad ...string) *String { + defaultPad := " " + if len(pad) > 0 { + defaultPad = pad[0] + } + short := maximum(0, length-s.Length()) + + s.value = Substr(strings.Repeat(defaultPad, short), 0, short) + s.value + return s +} + +// PadRight returns the String instance padded to the right side of the given length. +func (s *String) PadRight(length int, pad ...string) *String { + defaultPad := " " + if len(pad) > 0 { + defaultPad = pad[0] + } + short := maximum(0, length-s.Length()) + + s.value = s.value + Substr(strings.Repeat(defaultPad, short), 0, short) + return s +} + +// Pipe passes the string to the given callback and returns the result. +func (s *String) Pipe(callback func(s string) string) *String { + s.value = callback(s.value) + return s +} + +// Prepend one or more strings to the current string. +func (s *String) Prepend(values ...string) *String { + s.value = strings.Join(values, "") + s.value + return s +} + +// Remove returns the String instance with the first occurrence of the given value removed. +func (s *String) Remove(values ...string) *String { + for _, value := range values { + s.value = strings.ReplaceAll(s.value, value, "") + } + + return s +} + +// Repeat returns the String instance repeated the given number of times. +func (s *String) Repeat(times int) *String { + s.value = strings.Repeat(s.value, times) + return s +} + +// Replace returns the String instance with all occurrences of the search string replaced by the given replacement string. +func (s *String) Replace(search string, replace string, caseSensitive ...bool) *String { + caseSensitive = append(caseSensitive, true) + if len(caseSensitive) > 0 && !caseSensitive[0] { + s.value = regexp.MustCompile("(?i)"+search).ReplaceAllString(s.value, replace) + return s + } + s.value = strings.ReplaceAll(s.value, search, replace) + return s +} + +// ReplaceEnd returns the String instance with the last occurrence of the given value replaced. +func (s *String) ReplaceEnd(search string, replace string) *String { + if search == "" { + return s + } + + if s.EndsWith(search) { + return s.ReplaceLast(search, replace) + } + + return s +} + +// ReplaceFirst returns the String instance with the first occurrence of the given value replaced. +func (s *String) ReplaceFirst(search string, replace string) *String { + if search == "" { + return s + } + s.value = strings.Replace(s.value, search, replace, 1) + return s +} + +// ReplaceLast returns the String instance with the last occurrence of the given value replaced. +func (s *String) ReplaceLast(search string, replace string) *String { + if search == "" { + return s + } + index := strings.LastIndex(s.value, search) + if index != -1 { + s.value = s.value[:index] + replace + s.value[index+len(search):] + return s + } + + return s +} + +// ReplaceMatches returns the String instance with all occurrences of the given pattern +// replaced by the given replacement string. +func (s *String) ReplaceMatches(pattern string, replace string) *String { + s.value = regexp.MustCompile(pattern).ReplaceAllString(s.value, replace) + return s +} + +// ReplaceStart returns the String instance with the first occurrence of the given value replaced. +func (s *String) ReplaceStart(search string, replace string) *String { + if search == "" { + return s + } + + if s.StartsWith(search) { + return s.ReplaceFirst(search, replace) + } + + return s +} + +// RTrim returns the String instance with the right occurrences of the given value removed. +func (s *String) RTrim(characters ...string) *String { + if len(characters) == 0 { + s.value = strings.TrimRight(s.value, " ") + return s + } + + s.value = strings.TrimRight(s.value, characters[0]) + return s +} + +// Snake returns the String instance in snake case. +func (s *String) Snake(delimiter ...string) *String { + defaultDelimiter := "_" + if len(delimiter) > 0 { + defaultDelimiter = delimiter[0] + } + words := fieldsFunc(s.value, func(r rune) bool { + return r == ' ' || r == ',' || r == '.' || r == '-' || r == '_' + }, func(r rune) bool { + return unicode.IsUpper(r) + }) + + casesLower := cases.Lower(language.Und) + var studlyWords []string + for _, word := range words { + studlyWords = append(studlyWords, casesLower.String(word)) + } + + s.value = strings.Join(studlyWords, defaultDelimiter) + return s +} + +// Split splits the string by given pattern string. +func (s *String) Split(pattern string, limit ...int) []string { + r := regexp.MustCompile(pattern) + defaultLimit := -1 + if len(limit) != 0 { + defaultLimit = limit[0] + } + + return r.Split(s.value, defaultLimit) +} + +// Squish returns the String instance with consecutive whitespace characters collapsed into a single space. +func (s *String) Squish() *String { + leadWhitespace := regexp.MustCompile(`^[\s\p{Zs}]+|[\s\p{Zs}]+$`) + insideWhitespace := regexp.MustCompile(`[\s\p{Zs}]{2,}`) + s.value = leadWhitespace.ReplaceAllString(s.value, "") + s.value = insideWhitespace.ReplaceAllString(s.value, " ") + return s +} + +// Start returns the String instance with the given value prepended. +func (s *String) Start(prefix string) *String { + quoted := regexp.QuoteMeta(prefix) + re := regexp.MustCompile(`^(` + quoted + `)+`) + s.value = prefix + re.ReplaceAllString(s.value, "") + return s +} + +// StartsWith returns true if the string starts with the given value or any of the values. +func (s *String) StartsWith(values ...string) bool { + for _, value := range values { + if strings.HasPrefix(s.value, value) { + return true + } + } + + return false +} + +// String returns the string value. +func (s *String) String() string { + return s.value +} + +// Studly returns the String instance in studly case. +func (s *String) Studly() *String { + words := fieldsFunc(s.value, func(r rune) bool { + return r == '_' || r == ' ' || r == '-' || r == ',' || r == '.' + }, func(r rune) bool { + return unicode.IsUpper(r) + }) + + casesTitle := cases.Title(language.Und) + var studlyWords []string + for _, word := range words { + studlyWords = append(studlyWords, casesTitle.String(word)) + } + + s.value = strings.Join(studlyWords, "") + return s +} + +// Substr returns the String instance starting at the given index with the specified length. +func (s *String) Substr(start int, length ...int) *String { + s.value = Substr(s.value, start, length...) + return s +} + +// Swap replaces all occurrences of the search string with the given replacement string. +func (s *String) Swap(replacements map[string]string) *String { + if len(replacements) == 0 { + return s + } + + oldNewPairs := make([]string, 0, len(replacements)*2) + for k, v := range replacements { + if k == "" { + return s + } + oldNewPairs = append(oldNewPairs, k, v) + } + + s.value = strings.NewReplacer(oldNewPairs...).Replace(s.value) + return s +} + +// Tap passes the string to the given callback and returns the string. +func (s *String) Tap(callback func(String)) *String { + callback(*s) + return s +} + +// Test returns true if the string matches the given pattern. +func (s *String) Test(pattern string) bool { + return s.IsMatch(pattern) +} + +// Title returns the String instance in title case. +func (s *String) Title() *String { + casesTitle := cases.Title(language.Und) + s.value = casesTitle.String(s.value) + return s +} + +// Trim returns the String instance with trimmed characters from the left and right sides. +func (s *String) Trim(characters ...string) *String { + if len(characters) == 0 { + s.value = strings.TrimSpace(s.value) + return s + } + + s.value = strings.Trim(s.value, characters[0]) + return s +} + +// UcFirst returns the String instance with the first character uppercased. +func (s *String) UcFirst() *String { + if s.Length() == 0 { + return s + } + s.value = strings.ToUpper(Substr(s.value, 0, 1)) + Substr(s.value, 1) + return s +} + +// UcSplit splits the string into words using uppercase characters as the delimiter. +func (s *String) UcSplit() []string { + words := fieldsFunc(s.value, func(r rune) bool { + return false + }, func(r rune) bool { + return unicode.IsUpper(r) + }) + return words +} + +// Unless returns the String instance with the given fallback applied if the given condition is false. +func (s *String) Unless(callback func(*String) bool, fallback func(*String) *String) *String { + if !callback(s) { + return fallback(s) + } + + return s +} + +// Upper returns the String instance in upper case. +func (s *String) Upper() *String { + s.value = strings.ToUpper(s.value) + return s +} + +// When returns the String instance with the given callback applied if the given condition is true. +// If the condition is false, the fallback callback is applied.(if provided) +func (s *String) When(condition bool, callback ...func(*String) *String) *String { + if condition { + return callback[0](s) + } else { + if len(callback) > 1 { + return callback[1](s) + } + } + + return s +} + +// WhenContains returns the String instance with the given callback applied if the string contains the given value. +func (s *String) WhenContains(value string, callback ...func(*String) *String) *String { + return s.When(s.Contains(value), callback...) +} + +// WhenContainsAll returns the String instance with the given callback applied if the string contains all the given values. +func (s *String) WhenContainsAll(values []string, callback ...func(*String) *String) *String { + return s.When(s.ContainsAll(values...), callback...) +} + +// WhenEmpty returns the String instance with the given callback applied if the string is empty. +func (s *String) WhenEmpty(callback ...func(*String) *String) *String { + return s.When(s.IsEmpty(), callback...) +} + +// WhenIsAscii returns the String instance with the given callback applied if the string contains only ASCII characters. +func (s *String) WhenIsAscii(callback ...func(*String) *String) *String { + return s.When(s.IsAscii(), callback...) +} + +// WhenNotEmpty returns the String instance with the given callback applied if the string is not empty. +func (s *String) WhenNotEmpty(callback ...func(*String) *String) *String { + return s.When(s.IsNotEmpty(), callback...) +} + +// WhenStartsWith returns the String instance with the given callback applied if the string starts with the given value. +func (s *String) WhenStartsWith(value []string, callback ...func(*String) *String) *String { + return s.When(s.StartsWith(value...), callback...) +} + +// WhenEndsWith returns the String instance with the given callback applied if the string ends with the given value. +func (s *String) WhenEndsWith(value []string, callback ...func(*String) *String) *String { + return s.When(s.EndsWith(value...), callback...) +} + +// WhenExactly returns the String instance with the given callback applied if the string is exactly the given value. +func (s *String) WhenExactly(value string, callback ...func(*String) *String) *String { + return s.When(s.Exactly(value), callback...) +} + +// WhenNotExactly returns the String instance with the given callback applied if the string is not exactly the given value. +func (s *String) WhenNotExactly(value string, callback ...func(*String) *String) *String { + return s.When(!s.Exactly(value), callback...) +} + +// WhenIs returns the String instance with the given callback applied if the string matches any of the given patterns. +func (s *String) WhenIs(value string, callback ...func(*String) *String) *String { + return s.When(s.Is(value), callback...) +} + +// WhenIsUlid returns the String instance with the given callback applied if the string is a valid ULID. +func (s *String) WhenIsUlid(callback ...func(*String) *String) *String { + return s.When(s.IsUlid(), callback...) +} + +// WhenIsUuid returns the String instance with the given callback applied if the string is a valid UUID. +func (s *String) WhenIsUuid(callback ...func(*String) *String) *String { + return s.When(s.IsUuid(), callback...) +} + +// WhenTest returns the String instance with the given callback applied if the string matches the given pattern. +func (s *String) WhenTest(pattern string, callback ...func(*String) *String) *String { + return s.When(s.Test(pattern), callback...) +} + +// WordCount returns the number of words in the string. +func (s *String) WordCount() int { + return len(strings.Fields(s.value)) +} + +// Words return the String instance truncated to the given number of words. +func (s *String) Words(limit int, end ...string) *String { + defaultEnd := "..." + if len(end) > 0 { + defaultEnd = end[0] + } + + words := strings.Fields(s.value) + if len(words) <= limit { + return s + } + + s.value = strings.Join(words[:limit], " ") + defaultEnd + return s +} + +// Substr returns a substring of a given string, starting at the specified index +// and with a specified length. +// It handles UTF-8 encoded strings. +func Substr(str string, start int, length ...int) string { + // Convert the string to a rune slice for proper handling of UTF-8 encoding. + runes := []rune(str) + strLen := utf8.RuneCountInString(str) + end := strLen + // Check if the start index is out of bounds. + if start >= strLen { + return "" + } + + // If the start index is negative, count backwards from the end of the string. + if start < 0 { + start = strLen + start + if start < 0 { + start = 0 + } + } + + if len(length) > 0 { + if length[0] >= 0 { + end = start + length[0] + } else { + end = strLen + length[0] + } + } + + // If the length is 0, return the substring from start to the end of the string. + if len(length) == 0 { + return string(runes[start:]) + } + + // Handle the case where lenArg is negative and less than start + if end < start { + return "" + } + + if end > strLen { + end = strLen + } + + // Return the substring. + return string(runes[start:end]) +} + func Random(length int) string { b := make([]byte, length) _, err := rand.Read(b) @@ -91,3 +1007,50 @@ func (b *Buffer) append(s string) *Buffer { return b } + +// fieldsFunc splits the input string into words with preservation, following the rules defined by +// the provided functions f and preserveFunc. +func fieldsFunc(s string, f func(rune) bool, preserveFunc ...func(rune) bool) []string { + var fields []string + var currentField strings.Builder + + shouldPreserve := func(r rune) bool { + for _, preserveFn := range preserveFunc { + if preserveFn(r) { + return true + } + } + return false + } + + for _, r := range s { + if f(r) { + if currentField.Len() > 0 { + fields = append(fields, currentField.String()) + currentField.Reset() + } + } else if shouldPreserve(r) { + if currentField.Len() > 0 { + fields = append(fields, currentField.String()) + currentField.Reset() + } + currentField.WriteRune(r) + } else { + currentField.WriteRune(r) + } + } + + if currentField.Len() > 0 { + fields = append(fields, currentField.String()) + } + + return fields +} + +// maximum returns the largest of x or y. +func maximum[T constraints.Ordered](x T, y T) T { + if x > y { + return x + } + return y +} diff --git a/support/str/str_test.go b/support/str/str_test.go index 6c4a83cc6..12ba19663 100644 --- a/support/str/str_test.go +++ b/support/str/str_test.go @@ -4,11 +4,1097 @@ import ( "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/suite" + + "github.com/goravel/framework/support/env" ) +type StringTestSuite struct { + suite.Suite +} + +func TestStringTestSuite(t *testing.T) { + suite.Run(t, &StringTestSuite{}) +} + +func (s *StringTestSuite) SetupTest() { +} + +func (s *StringTestSuite) TestAfter() { + s.Equal("Framework", Of("GoravelFramework").After("Goravel").String()) + s.Equal("lel", Of("parallel").After("l").String()) + s.Equal("3def", Of("abc123def").After("2").String()) + s.Equal("abc123def", Of("abc123def").After("4").String()) + s.Equal("GoravelFramework", Of("GoravelFramework").After("").String()) +} + +func (s *StringTestSuite) TestAfterLast() { + s.Equal("Framework", Of("GoravelFramework").AfterLast("Goravel").String()) + s.Equal("", Of("parallel").AfterLast("l").String()) + s.Equal("3def", Of("abc123def").AfterLast("2").String()) + s.Equal("abc123def", Of("abc123def").AfterLast("4").String()) +} + +func (s *StringTestSuite) TestAppend() { + s.Equal("foobar", Of("foo").Append("bar").String()) + s.Equal("foobar", Of("foo").Append("bar").Append("").String()) + s.Equal("foobar", Of("foo").Append("bar").Append().String()) +} + +func (s *StringTestSuite) TestBasename() { + s.Equal("str", Of("/framework/support/str").Basename().String()) + s.Equal("str", Of("/framework/support/str/").Basename().String()) + s.Equal("str", Of("str").Basename().String()) + s.Equal("str", Of("/str").Basename().String()) + s.Equal("str", Of("/str/").Basename().String()) + s.Equal("str", Of("str/").Basename().String()) + + str := Of("/").Basename().String() + if env.IsWindows() { + s.Equal("\\", str) + } else { + s.Equal("/", str) + } + + s.Equal(".", Of("").Basename().String()) + s.Equal("str", Of("/framework/support/str/str.go").Basename(".go").String()) +} + +func (s *StringTestSuite) TestBefore() { + s.Equal("Goravel", Of("GoravelFramework").Before("Framework").String()) + s.Equal("para", Of("parallel").Before("l").String()) + s.Equal("abc123", Of("abc123def").Before("def").String()) + s.Equal("abc", Of("abc123def").Before("123").String()) +} + +func (s *StringTestSuite) TestBeforeLast() { + s.Equal("Goravel", Of("GoravelFramework").BeforeLast("Framework").String()) + s.Equal("paralle", Of("parallel").BeforeLast("l").String()) + s.Equal("abc123", Of("abc123def").BeforeLast("def").String()) + s.Equal("abc", Of("abc123def").BeforeLast("123").String()) +} + +func (s *StringTestSuite) TestBetween() { + s.Equal("foobarbaz", Of("foobarbaz").Between("", "b").String()) + s.Equal("foobarbaz", Of("foobarbaz").Between("f", "").String()) + s.Equal("foobarbaz", Of("foobarbaz").Between("", "").String()) + s.Equal("obar", Of("foobarbaz").Between("o", "b").String()) + s.Equal("bar", Of("foobarbaz").Between("foo", "baz").String()) + s.Equal("foo][bar][baz", Of("[foo][bar][baz]").Between("[", "]").String()) +} + +func (s *StringTestSuite) TestBetweenFirst() { + s.Equal("foobarbaz", Of("foobarbaz").BetweenFirst("", "b").String()) + s.Equal("foobarbaz", Of("foobarbaz").BetweenFirst("f", "").String()) + s.Equal("foobarbaz", Of("foobarbaz").BetweenFirst("", "").String()) + s.Equal("o", Of("foobarbaz").BetweenFirst("o", "b").String()) + s.Equal("foo", Of("[foo][bar][baz]").BetweenFirst("[", "]").String()) + s.Equal("foobar", Of("foofoobarbaz").BetweenFirst("foo", "baz").String()) +} + +func (s *StringTestSuite) TestCamel() { + s.Equal("goravelGOFramework", Of("Goravel_g_o_framework").Camel().String()) + s.Equal("goravelGOFramework", Of("Goravel_gO_framework").Camel().String()) + s.Equal("goravelGoFramework", Of("Goravel -_- go -_- framework ").Camel().String()) + + s.Equal("fooBar", Of("FooBar").Camel().String()) + s.Equal("fooBar", Of("foo_bar").Camel().String()) + s.Equal("fooBar", Of("foo-Bar").Camel().String()) + s.Equal("fooBar", Of("foo bar").Camel().String()) + s.Equal("fooBar", Of("foo.bar").Camel().String()) +} + +func (s *StringTestSuite) TestCharAt() { + s.Equal("好", Of("你好,世界!").CharAt(1)) + s.Equal("त", Of("नमस्ते, दुनिया!").CharAt(4)) + s.Equal("w", Of("Привет, world!").CharAt(8)) + s.Equal("계", Of("안녕하세요, 세계!").CharAt(-2)) + s.Equal("", Of("こんにちは、世界!").CharAt(-200)) +} + +func (s *StringTestSuite) TestContains() { + s.True(Of("kkumar").Contains("uma")) + s.True(Of("kkumar").Contains("kumar")) + s.True(Of("kkumar").Contains("uma", "xyz")) + s.False(Of("kkumar").Contains("xyz")) + s.False(Of("kkumar").Contains("")) +} + +func (s *StringTestSuite) TestContainsAll() { + s.True(Of("krishan kumar").ContainsAll("krishan", "kumar")) + s.True(Of("krishan kumar").ContainsAll("kumar")) + s.False(Of("krishan kumar").ContainsAll("kumar", "xyz")) +} + +func (s *StringTestSuite) TestDirname() { + str := Of("/framework/support/str").Dirname().String() + if env.IsWindows() { + s.Equal("\\framework\\support", str) + } else { + s.Equal("/framework/support", str) + } + + str = Of("/framework/support/str").Dirname(2).String() + if env.IsWindows() { + s.Equal("\\framework", str) + } else { + s.Equal("/framework", str) + } + + s.Equal(".", Of("framework").Dirname().String()) + s.Equal(".", Of(".").Dirname().String()) + + str = Of("/").Dirname().String() + if env.IsWindows() { + s.Equal("\\", str) + } else { + s.Equal("/", str) + } + + str = Of("/framework/").Dirname(2).String() + if env.IsWindows() { + s.Equal("\\", str) + } else { + s.Equal("/", str) + } +} + +func (s *StringTestSuite) TestEndsWith() { + s.True(Of("bowen").EndsWith("wen")) + s.True(Of("bowen").EndsWith("bowen")) + s.True(Of("bowen").EndsWith("wen", "xyz")) + s.False(Of("bowen").EndsWith("xyz")) + s.False(Of("bowen").EndsWith("")) + s.False(Of("bowen").EndsWith()) + s.False(Of("bowen").EndsWith("N")) + s.True(Of("a7.12").EndsWith("7.12")) + // Test for muti-byte string + s.True(Of("你好").EndsWith("好")) + s.True(Of("你好").EndsWith("你好")) + s.True(Of("你好").EndsWith("好", "xyz")) + s.False(Of("你好").EndsWith("xyz")) + s.False(Of("你好").EndsWith("")) +} + +func (s *StringTestSuite) TestExactly() { + s.True(Of("foo").Exactly("foo")) + s.False(Of("foo").Exactly("Foo")) +} + +func (s *StringTestSuite) TestExcerpt() { + s.Equal("...is a beautiful morn...", Of("This is a beautiful morning").Excerpt("beautiful", ExcerptOption{ + Radius: 5, + }).String()) + s.Equal("This is a beautiful morning", Of("This is a beautiful morning").Excerpt("foo", ExcerptOption{ + Radius: 5, + }).String()) + s.Equal("(...)is a beautiful morn(...)", Of("This is a beautiful morning").Excerpt("beautiful", ExcerptOption{ + Omission: "(...)", + Radius: 5, + }).String()) +} + +func (s *StringTestSuite) TestExplode() { + s.Equal([]string{"Foo", "Bar", "Baz"}, Of("Foo Bar Baz").Explode(" ")) + // with limit + s.Equal([]string{"Foo", "Bar Baz"}, Of("Foo Bar Baz").Explode(" ", 2)) + s.Equal([]string{"Foo", "Bar"}, Of("Foo Bar Baz").Explode(" ", -1)) + s.Equal([]string{}, Of("Foo Bar Baz").Explode(" ", -10)) +} + +func (s *StringTestSuite) TestFinish() { + s.Equal("abbc", Of("ab").Finish("bc").String()) + s.Equal("abbc", Of("abbcbc").Finish("bc").String()) + s.Equal("abcbbc", Of("abcbbcbc").Finish("bc").String()) +} + +func (s *StringTestSuite) TestHeadline() { + s.Equal("Hello", Of("hello").Headline().String()) + s.Equal("This Is A Headline", Of("this is a headline").Headline().String()) + s.Equal("Camelcase Is A Headline", Of("CamelCase is a headline").Headline().String()) + s.Equal("Kebab-Case Is A Headline", Of("kebab-case is a headline").Headline().String()) +} + +func (s *StringTestSuite) TestIs() { + s.True(Of("foo").Is("foo", "bar", "baz")) + s.True(Of("foo123").Is("bar*", "baz*", "foo*")) + s.False(Of("foo").Is("bar", "baz")) + s.True(Of("a.b").Is("a.b", "c.*")) + s.False(Of("abc*").Is("abc\\*", "xyz*")) + s.False(Of("").Is("foo")) + s.True(Of("foo/bar/baz").Is("foo/*", "bar/*", "baz*")) + // Is case-sensitive + s.False(Of("foo/bar/baz").Is("*BAZ*")) +} + +func (s *StringTestSuite) TestIsEmpty() { + s.True(Of("").IsEmpty()) + s.False(Of("F").IsEmpty()) +} + +func (s *StringTestSuite) TestIsNotEmpty() { + s.False(Of("").IsNotEmpty()) + s.True(Of("F").IsNotEmpty()) +} + +func (s *StringTestSuite) TestIsAscii() { + s.True(Of("abc").IsAscii()) + s.False(Of("你好").IsAscii()) +} + +func (s *StringTestSuite) TestIsSlice() { + // Test when the string represents a valid JSON array + s.True(Of(`["apple", "banana", "cherry"]`).IsSlice()) + + // Test when the string represents a valid JSON array with objects + s.True(Of(`[{"name": "John"}, {"name": "Alice"}]`).IsSlice()) + + // Test when the string represents an empty JSON array + s.True(Of(`[]`).IsSlice()) + + // Test when the string represents an invalid JSON object + s.False(Of(`{"name": "John"}`).IsSlice()) + + // Test when the string is not valid JSON + s.False(Of(`Not a JSON array`).IsSlice()) + + // Test when the string is empty + s.False(Of("").IsSlice()) +} + +func (s *StringTestSuite) TestIsMap() { + // Test when the string represents a valid JSON object + s.True(Of(`{"name": "John", "age": 30}`).IsMap()) + + // Test when the string represents a valid JSON object with nested objects + s.True(Of(`{"person": {"name": "Alice", "age": 25}}`).IsMap()) + + // Test when the string represents an empty JSON object + s.True(Of(`{}`).IsMap()) + + // Test when the string represents an invalid JSON array + s.False(Of(`["apple", "banana", "cherry"]`).IsMap()) + + // Test when the string is not valid JSON + s.False(Of(`Not a JSON object`).IsMap()) + + // Test when the string is empty + s.False(Of("").IsMap()) +} + +func (s *StringTestSuite) TestIsUlid() { + s.True(Of("01E65Z7XCHCR7X1P2MKF78ENRP").IsUlid()) + // lowercase characters are not allowed + s.False(Of("01e65z7xchcr7x1p2mkf78enrp").IsUlid()) + // too short (ULIDS must be 26 characters long) + s.False(Of("01E65Z7XCHCR7X1P2MKF78E").IsUlid()) + // contains invalid characters + s.False(Of("01E65Z7XCHCR7X1P2MKF78ENR!").IsUlid()) +} + +func (s *StringTestSuite) TestIsUuid() { + s.True(Of("3f2504e0-4f89-41d3-9a0c-0305e82c3301").IsUuid()) + s.False(Of("3f2504e0-4f89-41d3-9a0c-0305e82c3301-extra").IsUuid()) +} + +func (s *StringTestSuite) TestKebab() { + s.Equal("goravel-framework", Of("GoravelFramework").Kebab().String()) +} + +func (s *StringTestSuite) TestLcFirst() { + s.Equal("framework", Of("Framework").LcFirst().String()) + s.Equal("framework", Of("framework").LcFirst().String()) +} + +func (s *StringTestSuite) TestLength() { + s.Equal(11, Of("foo bar baz").Length()) + s.Equal(0, Of("").Length()) +} + +func (s *StringTestSuite) TestLimit() { + s.Equal("This is...", Of("This is a beautiful morning").Limit(7).String()) + s.Equal("This is****", Of("This is a beautiful morning").Limit(7, "****").String()) + s.Equal("这是一...", Of("这是一段中文").Limit(3).String()) + s.Equal("这是一段中文", Of("这是一段中文").Limit(9).String()) +} + +func (s *StringTestSuite) TestLower() { + s.Equal("foo bar baz", Of("FOO BAR BAZ").Lower().String()) + s.Equal("foo bar baz", Of("fOo Bar bAz").Lower().String()) +} + +func (s *StringTestSuite) TestLTrim() { + s.Equal("foo ", Of(" foo ").LTrim().String()) +} + +func (s *StringTestSuite) TestMask() { + s.Equal("kri**************", Of("krishan@email.com").Mask("*", 3).String()) + s.Equal("*******@email.com", Of("krishan@email.com").Mask("*", 0, 7).String()) + s.Equal("kris*************", Of("krishan@email.com").Mask("*", -13).String()) + s.Equal("kris***@email.com", Of("krishan@email.com").Mask("*", -13, 3).String()) + + s.Equal("*****************", Of("krishan@email.com").Mask("*", -17).String()) + s.Equal("*****an@email.com", Of("krishan@email.com").Mask("*", -99, 5).String()) + + s.Equal("krishan@email.com", Of("krishan@email.com").Mask("*", 17).String()) + s.Equal("krishan@email.com", Of("krishan@email.com").Mask("*", 17, 99).String()) + + s.Equal("krishan@email.com", Of("krishan@email.com").Mask("", 3).String()) + + s.Equal("krissssssssssssss", Of("krishan@email.com").Mask("something", 3).String()) + + s.Equal("这是一***", Of("这是一段中文").Mask("*", 3).String()) + s.Equal("**一段中文", Of("这是一段中文").Mask("*", 0, 2).String()) +} + +func (s *StringTestSuite) TestMatch() { + s.Equal("World", Of("Hello, World!").Match("World").String()) + s.Equal("(test)", Of("This is a (test) string").Match(`\([^)]+\)`).String()) + s.Equal("123", Of("abc123def456def").Match(`\d+`).String()) + s.Equal("", Of("No match here").Match(`\d+`).String()) + s.Equal("Hello, World!", Of("Hello, World!").Match("").String()) + s.Equal("[456]", Of("123 [456]").Match(`\[456\]`).String()) +} + +func (s *StringTestSuite) TestMatchAll() { + s.Equal([]string{"World"}, Of("Hello, World!").MatchAll("World")) + s.Equal([]string{"(test)"}, Of("This is a (test) string").MatchAll(`\([^)]+\)`)) + s.Equal([]string{"123", "456"}, Of("abc123def456def").MatchAll(`\d+`)) + s.Equal([]string(nil), Of("No match here").MatchAll(`\d+`)) + s.Equal([]string{"Hello, World!"}, Of("Hello, World!").MatchAll("")) + s.Equal([]string{"[456]"}, Of("123 [456]").MatchAll(`\[456\]`)) +} + +func (s *StringTestSuite) TestIsMatch() { + // Test matching with a single pattern + s.True(Of("Hello, Goravel!").IsMatch(`.*,.*!`)) + s.True(Of("Hello, Goravel!").IsMatch(`^.*$(.*)`)) + s.True(Of("Hello, Goravel!").IsMatch(`(?i)goravel`)) + s.True(Of("Hello, GOravel!").IsMatch(`^(.*(.*(.*)))`)) + + // Test non-matching with a single pattern + s.False(Of("Hello, Goravel!").IsMatch(`H.o`)) + s.False(Of("Hello, Goravel!").IsMatch(`^goravel!`)) + s.False(Of("Hello, Goravel!").IsMatch(`goravel!(.*)`)) + s.False(Of("Hello, Goravel!").IsMatch(`^[a-zA-Z,!]+$`)) + + // Test with multiple patterns + s.True(Of("Hello, Goravel!").IsMatch(`.*,.*!`, `H.o`)) + s.True(Of("Hello, Goravel!").IsMatch(`(?i)goravel`, `^.*$(.*)`)) + s.True(Of("Hello, Goravel!").IsMatch(`(?i)goravel`, `goravel!(.*)`)) + s.True(Of("Hello, Goravel!").IsMatch(`^[a-zA-Z,!]+$`, `^(.*(.*(.*)))`)) +} + +func (s *StringTestSuite) TestNewLine() { + s.Equal("Goravel\n", Of("Goravel").NewLine().String()) + s.Equal("Goravel\n\nbar", Of("Goravel").NewLine(2).Append("bar").String()) +} + +func (s *StringTestSuite) TestPadBoth() { + // Test padding with spaces + s.Equal(" Hello ", Of("Hello").PadBoth(11, " ").String()) + s.Equal(" World! ", Of("World!").PadBoth(10, " ").String()) + s.Equal("==Hello===", Of("Hello").PadBoth(10, "=").String()) + s.Equal("Hello", Of("Hello").PadBoth(3, " ").String()) + s.Equal(" ", Of("").PadBoth(6, " ").String()) +} + +func (s *StringTestSuite) TestPadLeft() { + s.Equal(" Goravel", Of("Goravel").PadLeft(10, " ").String()) + s.Equal("==Goravel", Of("Goravel").PadLeft(9, "=").String()) + s.Equal("Goravel", Of("Goravel").PadLeft(3, " ").String()) +} + +func (s *StringTestSuite) TestPadRight() { + s.Equal("Goravel ", Of("Goravel").PadRight(10, " ").String()) + s.Equal("Goravel==", Of("Goravel").PadRight(9, "=").String()) + s.Equal("Goravel", Of("Goravel").PadRight(3, " ").String()) +} + +func (s *StringTestSuite) TestPipe() { + callback := func(str string) string { + return Of(str).Append("bar").String() + } + s.Equal("foobar", Of("foo").Pipe(callback).String()) +} + +func (s *StringTestSuite) TestPrepend() { + s.Equal("foobar", Of("bar").Prepend("foo").String()) + s.Equal("foobar", Of("bar").Prepend("foo").Prepend("").String()) + s.Equal("foobar", Of("bar").Prepend("foo").Prepend().String()) +} + +func (s *StringTestSuite) TestRemove() { + s.Equal("Fbar", Of("Foobar").Remove("o").String()) + s.Equal("Foo", Of("Foobar").Remove("bar").String()) + s.Equal("oobar", Of("Foobar").Remove("F").String()) + s.Equal("Foobar", Of("Foobar").Remove("f").String()) + + s.Equal("Fbr", Of("Foobar").Remove("o", "a").String()) + s.Equal("Fooar", Of("Foobar").Remove("f", "b").String()) + s.Equal("Foobar", Of("Foo|bar").Remove("f", "|").String()) +} + +func (s *StringTestSuite) TestRepeat() { + s.Equal("aaaaa", Of("a").Repeat(5).String()) + s.Equal("", Of("").Repeat(5).String()) +} + +func (s *StringTestSuite) TestReplace() { + s.Equal("foo/foo/foo", Of("?/?/?").Replace("?", "foo").String()) + s.Equal("foo/foo/foo", Of("x/x/x").Replace("X", "foo", false).String()) + s.Equal("bar/bar", Of("?/?").Replace("?", "bar").String()) + s.Equal("?/?/?", Of("? ? ?").Replace(" ", "/").String()) +} + +func (s *StringTestSuite) TestReplaceEnd() { + s.Equal("Golang is great!", Of("Golang is good!").ReplaceEnd("good!", "great!").String()) + s.Equal("Hello, World!", Of("Hello, Earth!").ReplaceEnd("Earth!", "World!").String()) + s.Equal("München Berlin", Of("München Frankfurt").ReplaceEnd("Frankfurt", "Berlin").String()) + s.Equal("Café Latte", Of("Café Americano").ReplaceEnd("Americano", "Latte").String()) + s.Equal("Golang is good!", Of("Golang is good!").ReplaceEnd("", "great!").String()) + s.Equal("Golang is good!", Of("Golang is good!").ReplaceEnd("excellent!", "great!").String()) +} + +func (s *StringTestSuite) TestReplaceFirst() { + s.Equal("fooqux foobar", Of("foobar foobar").ReplaceFirst("bar", "qux").String()) + s.Equal("foo/qux? foo/bar?", Of("foo/bar? foo/bar?").ReplaceFirst("bar?", "qux?").String()) + s.Equal("foo foobar", Of("foobar foobar").ReplaceFirst("bar", "").String()) + s.Equal("foobar foobar", Of("foobar foobar").ReplaceFirst("xxx", "yyy").String()) + s.Equal("foobar foobar", Of("foobar foobar").ReplaceFirst("", "yyy").String()) + // Test for multibyte string support + s.Equal("Jxxxnköping Malmö", Of("Jönköping Malmö").ReplaceFirst("ö", "xxx").String()) + s.Equal("Jönköping Malmö", Of("Jönköping Malmö").ReplaceFirst("", "yyy").String()) +} + +func (s *StringTestSuite) TestReplaceLast() { + s.Equal("foobar fooqux", Of("foobar foobar").ReplaceLast("bar", "qux").String()) + s.Equal("foo/bar? foo/qux?", Of("foo/bar? foo/bar?").ReplaceLast("bar?", "qux?").String()) + s.Equal("foobar foo", Of("foobar foobar").ReplaceLast("bar", "").String()) + s.Equal("foobar foobar", Of("foobar foobar").ReplaceLast("xxx", "yyy").String()) + s.Equal("foobar foobar", Of("foobar foobar").ReplaceLast("", "yyy").String()) + // Test for multibyte string support + s.Equal("Malmö Jönkxxxping", Of("Malmö Jönköping").ReplaceLast("ö", "xxx").String()) + s.Equal("Malmö Jönköping", Of("Malmö Jönköping").ReplaceLast("", "yyy").String()) +} + +func (s *StringTestSuite) TestReplaceMatches() { + s.Equal("Golang is great!", Of("Golang is good!").ReplaceMatches("good", "great").String()) + s.Equal("Hello, World!", Of("Hello, Earth!").ReplaceMatches("Earth", "World").String()) + s.Equal("Apples, Apples, Apples", Of("Oranges, Oranges, Oranges").ReplaceMatches("Oranges", "Apples").String()) + s.Equal("1, 2, 3, 4, 5", Of("10, 20, 30, 40, 50").ReplaceMatches("0", "").String()) + s.Equal("München Berlin", Of("München Frankfurt").ReplaceMatches("Frankfurt", "Berlin").String()) + s.Equal("Café Latte", Of("Café Americano").ReplaceMatches("Americano", "Latte").String()) + s.Equal("The quick brown fox", Of("The quick brown fox").ReplaceMatches(`\b([a-z])`, `$1`).String()) + s.Equal("One, One, One", Of("1, 2, 3").ReplaceMatches(`\d`, "One").String()) + s.Equal("Hello, World!", Of("Hello, World!").ReplaceMatches("Earth", "").String()) + s.Equal("Hello, World!", Of("Hello, World!").ReplaceMatches("Golang", "Great").String()) +} + +func (s *StringTestSuite) TestReplaceStart() { + s.Equal("foobar foobar", Of("foobar foobar").ReplaceStart("bar", "qux").String()) + s.Equal("foo/bar? foo/bar?", Of("foo/bar? foo/bar?").ReplaceStart("bar?", "qux?").String()) + s.Equal("quxbar foobar", Of("foobar foobar").ReplaceStart("foo", "qux").String()) + s.Equal("qux? foo/bar?", Of("foo/bar? foo/bar?").ReplaceStart("foo/bar?", "qux?").String()) + s.Equal("bar foobar", Of("foobar foobar").ReplaceStart("foo", "").String()) + s.Equal("1", Of("0").ReplaceStart("0", "1").String()) + // Test for multibyte string support + s.Equal("xxxnköping Malmö", Of("Jönköping Malmö").ReplaceStart("Jö", "xxx").String()) + s.Equal("Jönköping Malmö", Of("Jönköping Malmö").ReplaceStart("", "yyy").String()) +} + +func (s *StringTestSuite) TestRTrim() { + s.Equal(" foo", Of(" foo ").RTrim().String()) + s.Equal(" foo", Of(" foo__").RTrim("_").String()) +} + +func (s *StringTestSuite) TestSnake() { + s.Equal("goravel_g_o_framework", Of("GoravelGOFramework").Snake().String()) + s.Equal("goravel_go_framework", Of("GoravelGoFramework").Snake().String()) + s.Equal("goravel go framework", Of("GoravelGoFramework").Snake(" ").String()) + s.Equal("goravel_go_framework", Of("Goravel Go Framework").Snake().String()) + s.Equal("goravel_go_framework", Of("Goravel Go Framework ").Snake().String()) + s.Equal("goravel__go__framework", Of("GoravelGoFramework").Snake("__").String()) + s.Equal("żółta_łódka", Of("ŻółtaŁódka").Snake().String()) +} + +func (s *StringTestSuite) TestSplit() { + s.Equal([]string{"one", "two", "three", "four"}, Of("one-two-three-four").Split("-")) + s.Equal([]string{"", "", "D", "E", "", ""}, Of(",,D,E,,").Split(",")) + s.Equal([]string{"one", "two", "three,four"}, Of("one,two,three,four").Split(",", 3)) +} + +func (s *StringTestSuite) TestSquish() { + s.Equal("Hello World", Of(" Hello World ").Squish().String()) + s.Equal("A B C", Of("A B C").Squish().String()) + s.Equal("Lorem ipsum dolor sit amet", Of(" Lorem ipsum \n dolor sit \t amet ").Squish().String()) + s.Equal("Leading and trailing spaces", Of(" Leading "+ + "and trailing "+ + " spaces ").Squish().String()) + s.Equal("", Of("").Squish().String()) +} + +func (s *StringTestSuite) TestStart() { + s.Equal("/test/string", Of("test/string").Start("/").String()) + s.Equal("/test/string", Of("/test/string").Start("/").String()) + s.Equal("/test/string", Of("//test/string").Start("/").String()) +} + +func (s *StringTestSuite) TestStartsWith() { + s.True(Of("Wenbo Han").StartsWith("Wen")) + s.True(Of("Wenbo Han").StartsWith("Wenbo")) + s.True(Of("Wenbo Han").StartsWith("Han", "Wen")) + s.False(Of("Wenbo Han").StartsWith()) + s.False(Of("Wenbo Han").StartsWith("we")) + s.True(Of("Jönköping").StartsWith("Jö")) + s.False(Of("Jönköping").StartsWith("Jonko")) +} + +func (s *StringTestSuite) TestStudly() { + s.Equal("GoravelGOFramework", Of("Goravel_g_o_framework").Studly().String()) + s.Equal("GoravelGOFramework", Of("Goravel_gO_framework").Studly().String()) + s.Equal("GoravelGoFramework", Of("Goravel -_- go -_- framework ").Studly().String()) + + s.Equal("FooBar", Of("FooBar").Studly().String()) + s.Equal("FooBar", Of("foo_bar").Studly().String()) + s.Equal("FooBar", Of("foo-Bar").Studly().String()) + s.Equal("FooBar", Of("foo bar").Studly().String()) + s.Equal("FooBar", Of("foo.bar").Studly().String()) +} + +func (s *StringTestSuite) TestSubstr() { + s.Equal("Ё", Of("БГДЖИЛЁ").Substr(-1).String()) + s.Equal("ЛЁ", Of("БГДЖИЛЁ").Substr(-2).String()) + s.Equal("И", Of("БГДЖИЛЁ").Substr(-3, 1).String()) + s.Equal("ДЖИЛ", Of("БГДЖИЛЁ").Substr(2, -1).String()) + s.Equal("", Of("БГДЖИЛЁ").Substr(4, -4).String()) + s.Equal("ИЛ", Of("БГДЖИЛЁ").Substr(-3, -1).String()) + s.Equal("ГДЖИЛЁ", Of("БГДЖИЛЁ").Substr(1).String()) + s.Equal("ГДЖ", Of("БГДЖИЛЁ").Substr(1, 3).String()) + s.Equal("БГДЖ", Of("БГДЖИЛЁ").Substr(0, 4).String()) + s.Equal("Ё", Of("БГДЖИЛЁ").Substr(-1, 1).String()) + s.Equal("", Of("Б").Substr(2).String()) +} + +func (s *StringTestSuite) TestSwap() { + s.Equal("Go is excellent", Of("Golang is awesome").Swap(map[string]string{ + "Golang": "Go", + "awesome": "excellent", + }).String()) + s.Equal("Golang is awesome", Of("Golang is awesome").Swap(map[string]string{}).String()) + s.Equal("Golang is awesome", Of("Golang is awesome").Swap(map[string]string{ + "": "Go", + "awesome": "excellent", + }).String()) +} + +func (s *StringTestSuite) TestTap() { + tap := Of("foobarbaz") + fromTehTap := "" + tap = tap.Tap(func(s String) { + fromTehTap = s.Substr(0, 3).String() + }) + s.Equal("foo", fromTehTap) + s.Equal("foobarbaz", tap.String()) +} + +func (s *StringTestSuite) TestTitle() { + s.Equal("Krishan Kumar", Of("krishan kumar").Title().String()) + s.Equal("Krishan Kumar", Of("kriSHan kuMAr").Title().String()) +} + +func (s *StringTestSuite) TestTrim() { + s.Equal("foo", Of(" foo ").Trim().String()) + s.Equal("foo", Of("_foo_").Trim("_").String()) +} + +func (s *StringTestSuite) TestUcFirst() { + s.Equal("", Of("").UcFirst().String()) + s.Equal("Framework", Of("framework").UcFirst().String()) + s.Equal("Framework", Of("Framework").UcFirst().String()) + s.Equal(" framework", Of(" framework").UcFirst().String()) + s.Equal("Goravel framework", Of("goravel framework").UcFirst().String()) +} + +func (s *StringTestSuite) TestUcSplit() { + s.Equal([]string{"Krishan", "Kumar"}, Of("KrishanKumar").UcSplit()) + s.Equal([]string{"Hello", "From", "Goravel"}, Of("HelloFromGoravel").UcSplit()) + s.Equal([]string{"He_llo_", "World"}, Of("He_llo_World").UcSplit()) +} + +func (s *StringTestSuite) TestUnless() { + str := Of("Hello, World!") + + // Test case 1: The callback returns true, so the fallback should not be applied + s.Equal("Hello, World!", str.Unless(func(s *String) bool { + return true + }, func(s *String) *String { + return Of("This should not be applied") + }).String()) + + // Test case 2: The callback returns false, so the fallback should be applied + s.Equal("Fallback Applied", str.Unless(func(s *String) bool { + return false + }, func(s *String) *String { + return Of("Fallback Applied") + }).String()) + + // Test case 3: Testing with an empty string + s.Equal("Fallback Applied", Of("").Unless(func(s *String) bool { + return false + }, func(s *String) *String { + return Of("Fallback Applied") + }).String()) +} + +func (s *StringTestSuite) TestUpper() { + s.Equal("FOO BAR BAZ", Of("foo bar baz").Upper().String()) + s.Equal("FOO BAR BAZ", Of("foO bAr BaZ").Upper().String()) +} + +func (s *StringTestSuite) TestWhen() { + // true + s.Equal("when true", Of("when ").When(true, func(s *String) *String { + return s.Append("true") + }).String()) + s.Equal("gets a value from if", Of("gets a value ").When(true, func(s *String) *String { + return s.Append("from if") + }).String()) + + // false + s.Equal("when", Of("when").When(false, func(s *String) *String { + return s.Append("true") + }).String()) + + s.Equal("when false fallbacks to default", Of("when false ").When(false, func(s *String) *String { + return s.Append("true") + }, func(s *String) *String { + return s.Append("fallbacks to default") + }).String()) +} + +func (s *StringTestSuite) TestWhenContains() { + s.Equal("Tony Stark", Of("stark").WhenContains("tar", func(s *String) *String { + return s.Prepend("Tony ").Title() + }, func(s *String) *String { + return s.Prepend("Arno ").Title() + }).String()) + + s.Equal("stark", Of("stark").WhenContains("xxx", func(s *String) *String { + return s.Prepend("Tony ").Title() + }).String()) + + s.Equal("Arno Stark", Of("stark").WhenContains("xxx", func(s *String) *String { + return s.Prepend("Tony ").Title() + }, func(s *String) *String { + return s.Prepend("Arno ").Title() + }).String()) +} + +func (s *StringTestSuite) TestWhenContainsAll() { + // Test when all values are present + s.Equal("Tony Stark", Of("tony stark").WhenContainsAll([]string{"tony", "stark"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) + + // Test when not all values are present + s.Equal("tony stark", Of("tony stark").WhenContainsAll([]string{"xxx"}, + func(s *String) *String { + return s.Title() + }, + ).String()) + + // Test when some values are present and some are not + s.Equal("TonyStark", Of("tony stark").WhenContainsAll([]string{"tony", "xxx"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenEmpty() { + // Test when the string is empty + s.Equal("DEFAULT", Of("").WhenEmpty( + func(s *String) *String { + return s.Append("default").Upper() + }).String()) + + // Test when the string is not empty + s.Equal("non-empty", Of("non-empty").WhenEmpty( + func(s *String) *String { + return s.Append("default") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenIsAscii() { + s.Equal("Ascii: A", Of("A").WhenIsAscii( + func(s *String) *String { + return s.Prepend("Ascii: ") + }).String()) + s.Equal("ù", Of("ù").WhenIsAscii( + func(s *String) *String { + return s.Prepend("Ascii: ") + }).String()) + s.Equal("Not Ascii: ù", Of("ù").WhenIsAscii( + func(s *String) *String { + return s.Prepend("Ascii: ") + }, + func(s *String) *String { + return s.Prepend("Not Ascii: ") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenNotEmpty() { + // Test when the string is not empty + s.Equal("UPPERCASE", Of("uppercase").WhenNotEmpty( + func(s *String) *String { + return s.Upper() + }, + ).String()) + + // Test when the string is empty + s.Equal("", Of("").WhenNotEmpty( + func(s *String) *String { + return s.Append("not empty") + }, + func(s *String) *String { + return s.Upper() + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenStartsWith() { + // Test when the string starts with a specific prefix + s.Equal("Tony Stark", Of("tony stark").WhenStartsWith([]string{"ton"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) + + // Test when the string starts with any of the specified prefixes + s.Equal("Tony Stark", Of("tony stark").WhenStartsWith([]string{"ton", "not"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) + + // Test when the string does not start with the specified prefix + s.Equal("tony stark", Of("tony stark").WhenStartsWith([]string{"xxx"}, + func(s *String) *String { + return s.Title() + }, + ).String()) + + // Test when the string starts with one of the specified prefixes and not the other + s.Equal("Tony Stark", Of("tony stark").WhenStartsWith([]string{"tony", "xxx"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenEndsWith() { + // Test when the string ends with a specific suffix + s.Equal("Tony Stark", Of("tony stark").WhenEndsWith([]string{"ark"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) + + // Test when the string ends with any of the specified suffixes + s.Equal("Tony Stark", Of("tony stark").WhenEndsWith([]string{"kra", "ark"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) + + // Test when the string does not end with the specified suffix + s.Equal("tony stark", Of("tony stark").WhenEndsWith([]string{"xxx"}, + func(s *String) *String { + return s.Title() + }, + ).String()) + + // Test when the string ends with one of the specified suffixes and not the other + s.Equal("TonyStark", Of("tony stark").WhenEndsWith([]string{"tony", "xxx"}, + func(s *String) *String { + return s.Title() + }, + func(s *String) *String { + return s.Studly() + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenExactly() { + // Test when the string exactly matches the expected value + s.Equal("Nailed it...!", Of("Tony Stark").WhenExactly("Tony Stark", + func(s *String) *String { + return Of("Nailed it...!") + }, + func(s *String) *String { + return Of("Swing and a miss...!") + }, + ).String()) + + // Test when the string does not exactly match the expected value + s.Equal("Swing and a miss...!", Of("Tony Stark").WhenExactly("Iron Man", + func(s *String) *String { + return Of("Nailed it...!") + }, + func(s *String) *String { + return Of("Swing and a miss...!") + }, + ).String()) + + // Test when the string exactly matches the expected value with no "else" callback + s.Equal("Tony Stark", Of("Tony Stark").WhenExactly("Iron Man", + func(s *String) *String { + return Of("Nailed it...!") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenNotExactly() { + // Test when the string does not exactly match the expected value with an "else" callback + s.Equal("Iron Man", Of("Tony").WhenNotExactly("Tony Stark", + func(s *String) *String { + return Of("Iron Man") + }, + ).String()) + + // Test when the string does not exactly match the expected value with both "if" and "else" callbacks + s.Equal("Swing and a miss...!", Of("Tony Stark").WhenNotExactly("Tony Stark", + func(s *String) *String { + return Of("Iron Man") + }, + func(s *String) *String { + return Of("Swing and a miss...!") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenIs() { + // Test when the string exactly matches the expected value with an "if" callback + s.Equal("Winner: /", Of("/").WhenIs("/", + func(s *String) *String { + return s.Prepend("Winner: ") + }, + func(s *String) *String { + return Of("Try again") + }, + ).String()) + + // Test when the string does not exactly match the expected value with an "if" callback + s.Equal("/", Of("/").WhenIs(" /", + func(s *String) *String { + return s.Prepend("Winner: ") + }, + ).String()) + + // Test when the string does not exactly match the expected value with both "if" and "else" callbacks + s.Equal("Try again", Of("/").WhenIs(" /", + func(s *String) *String { + return s.Prepend("Winner: ") + }, + func(s *String) *String { + return Of("Try again") + }, + ).String()) + + // Test when the string matches a pattern using wildcard and "if" callback + s.Equal("Winner: foo/bar/baz", Of("foo/bar/baz").WhenIs("foo/*", + func(s *String) *String { + return s.Prepend("Winner: ") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenIsUlid() { + // Test when the string is a valid ULID with an "if" callback + s.Equal("Ulid: 01GJSNW9MAF792C0XYY8RX6QFT", Of("01GJSNW9MAF792C0XYY8RX6QFT").WhenIsUlid( + func(s *String) *String { + return s.Prepend("Ulid: ") + }, + func(s *String) *String { + return s.Prepend("Not Ulid: ") + }, + ).String()) + + // Test when the string is not a valid ULID with an "if" callback + s.Equal("2cdc7039-65a6-4ac7-8e5d-d554a98", Of("2cdc7039-65a6-4ac7-8e5d-d554a98").WhenIsUlid( + func(s *String) *String { + return s.Prepend("Ulid: ") + }, + ).String()) + + // Test when the string is not a valid ULID with both "if" and "else" callbacks + s.Equal("Not Ulid: ss-01GJSNW9MAF792C0XYY8RX6QFT", Of("ss-01GJSNW9MAF792C0XYY8RX6QFT").WhenIsUlid( + func(s *String) *String { + return s.Prepend("Ulid: ") + }, + func(s *String) *String { + return s.Prepend("Not Ulid: ") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenIsUuid() { + // Test when the string is a valid UUID with an "if" callback + s.Equal("Uuid: 2cdc7039-65a6-4ac7-8e5d-d554a98e7b15", Of("2cdc7039-65a6-4ac7-8e5d-d554a98e7b15").WhenIsUuid( + func(s *String) *String { + return s.Prepend("Uuid: ") + }, + func(s *String) *String { + return s.Prepend("Not Uuid: ") + }, + ).String()) + + s.Equal("2cdc7039-65a6-4ac7-8e5d-d554a98", Of("2cdc7039-65a6-4ac7-8e5d-d554a98").WhenIsUuid( + func(s *String) *String { + return s.Prepend("Uuid: ") + }, + ).String()) + + s.Equal("Not Uuid: 2cdc7039-65a6-4ac7-8e5d-d554a98", Of("2cdc7039-65a6-4ac7-8e5d-d554a98").WhenIsUuid( + func(s *String) *String { + return s.Prepend("Uuid: ") + }, + func(s *String) *String { + return s.Prepend("Not Uuid: ") + }, + ).String()) +} + +func (s *StringTestSuite) TestWhenTest() { + // Test when the regular expression matches with an "if" callback + s.Equal("Winner: foo bar", Of("foo bar").WhenTest(`bar*`, + func(s *String) *String { + return s.Prepend("Winner: ") + }, + func(s *String) *String { + return Of("Try again") + }, + ).String()) + + // Test when the regular expression does not match with an "if" callback + s.Equal("Try again", Of("foo bar").WhenTest(`/link/`, + func(s *String) *String { + return s.Prepend("Winner: ") + }, + func(s *String) *String { + return Of("Try again") + }, + ).String()) + + // Test when the regular expression does not match with both "if" and "else" callbacks + s.Equal("foo bar", Of("foo bar").WhenTest(`/link/`, + func(s *String) *String { + return s.Prepend("Winner: ") + }, + ).String()) +} + +func (s *StringTestSuite) TestWordCount() { + s.Equal(2, Of("Hello, world!").WordCount()) + s.Equal(10, Of("Hi, this is my first contribution to the Goravel framework.").WordCount()) +} + +func (s *StringTestSuite) TestWords() { + s.Equal("Perfectly balanced, as >>>", Of("Perfectly balanced, as all things should be.").Words(3, " >>>").String()) + s.Equal("Perfectly balanced, as all things should be.", Of("Perfectly balanced, as all things should be.").Words(100).String()) +} + +func TestFieldsFunc(t *testing.T) { + tests := []struct { + input string + shouldPreserve []func(rune) bool + expected []string + }{ + // Test case 1: Basic word splitting with space separator. + { + input: "Hello World", + expected: []string{"Hello", "World"}, + }, + // Test case 2: Splitting with space and preserving hyphen. + { + input: "Hello-World", + shouldPreserve: []func(rune) bool{func(r rune) bool { return r == '-' }}, + expected: []string{"Hello", "-World"}, + }, + // Test case 3: Splitting with space and preserving multiple characters. + { + input: "Hello-World,This,Is,a,Test", + shouldPreserve: []func(rune) bool{ + func(r rune) bool { return r == '-' }, + func(r rune) bool { return r == ',' }, + }, + expected: []string{"Hello", "-World", ",This", ",Is", ",a", ",Test"}, + }, + // Test case 4: No splitting when no separator is found. + { + input: "HelloWorld", + expected: []string{"HelloWorld"}, + }, + } + + for _, test := range tests { + t.Run(test.input, func(t *testing.T) { + result := fieldsFunc(test.input, func(r rune) bool { return r == ' ' }, test.shouldPreserve...) + assert.Equal(t, test.expected, result) + }) + } +} + +func TestSubstr(t *testing.T) { + assert.Equal(t, "world", Substr("Hello, world!", 7, 5)) + assert.Equal(t, "", Substr("Golang", 10)) + assert.Equal(t, "tine", Substr("Goroutines", -5, 4)) + assert.Equal(t, "ic", Substr("Unicode", 2, -3)) + assert.Equal(t, "esting", Substr("Testing", 1, 10)) + assert.Equal(t, "", Substr("", 0, 5)) + assert.Equal(t, "世界!", Substr("你好,世界!", 3, 3)) +} + +func TestMaximum(t *testing.T) { + assert.Equal(t, 10, maximum(5, 10)) + assert.Equal(t, 3.14, maximum(3.14, 2.71)) + assert.Equal(t, "banana", maximum("apple", "banana")) + assert.Equal(t, -5, maximum(-5, -10)) + assert.Equal(t, 42, maximum(42, 42)) +} + func TestRandom(t *testing.T) { assert.Len(t, Random(10), 10) assert.Empty(t, Random(0)) + assert.Panics(t, func() { + Random(-1) + }) } func TestCase2Camel(t *testing.T) { diff --git a/testing/file/file.go b/testing/file/file.go index 7ffc0c25c..c4a68de31 100644 --- a/testing/file/file.go +++ b/testing/file/file.go @@ -24,9 +24,7 @@ func GetLineNum(file string) int { } } - defer func() { - f.Close() - }() + defer f.Close() return total } diff --git a/testing/file/file_test.go b/testing/file/file_test.go index a2419d367..857f4fb4c 100644 --- a/testing/file/file_test.go +++ b/testing/file/file_test.go @@ -7,5 +7,5 @@ import ( ) func TestGetLineNum(t *testing.T) { - assert.Equal(t, 33, GetLineNum("file.go")) + assert.Equal(t, 31, GetLineNum("file.go")) } diff --git a/testing/mock/log.go b/testing/mock/log.go index 243b178d2..1d11455ff 100644 --- a/testing/mock/log.go +++ b/testing/mock/log.go @@ -1,6 +1,7 @@ package mock import ( + "context" "fmt" "github.com/goravel/framework/contracts/http" @@ -9,108 +10,161 @@ import ( ) type TestLog struct { + *TestLogWriter } -func NewTestLog() log.Writer { - return &TestLog{} +func NewTestLog() log.Log { + return &TestLog{ + TestLogWriter: NewTestLogWriter(), + } } -func (r *TestLog) Debug(args ...any) { +func (r *TestLog) WithContext(ctx context.Context) log.Writer { + return NewTestLogWriter() +} + +type TestLogWriter struct { + data map[string]any +} + +func NewTestLogWriter() *TestLogWriter { + return &TestLogWriter{ + data: make(map[string]any), + } +} + +func (r *TestLogWriter) Debug(args ...any) { fmt.Print(prefix("debug")) fmt.Println(args...) + r.printData() } -func (r *TestLog) Debugf(format string, args ...any) { +func (r *TestLogWriter) Debugf(format string, args ...any) { fmt.Print(prefix("debug")) fmt.Printf(format+"\n", args...) + r.printData() } -func (r *TestLog) Info(args ...any) { +func (r *TestLogWriter) Info(args ...any) { fmt.Print(prefix("info")) fmt.Println(args...) + r.printData() } -func (r *TestLog) Infof(format string, args ...any) { +func (r *TestLogWriter) Infof(format string, args ...any) { fmt.Print(prefix("info")) fmt.Printf(format+"\n", args...) + r.printData() } -func (r *TestLog) Warning(args ...any) { +func (r *TestLogWriter) Warning(args ...any) { fmt.Print(prefix("warning")) fmt.Println(args...) + r.printData() } -func (r *TestLog) Warningf(format string, args ...any) { +func (r *TestLogWriter) Warningf(format string, args ...any) { fmt.Print(prefix("warning")) fmt.Printf(format+"\n", args...) + r.printData() } -func (r *TestLog) Error(args ...any) { +func (r *TestLogWriter) Error(args ...any) { fmt.Print(prefix("error")) fmt.Println(args...) + r.printData() } -func (r *TestLog) Errorf(format string, args ...any) { +func (r *TestLogWriter) Errorf(format string, args ...any) { fmt.Print(prefix("error")) fmt.Printf(format+"\n", args...) + r.printData() } -func (r *TestLog) Fatal(args ...any) { +func (r *TestLogWriter) Fatal(args ...any) { fmt.Print(prefix("fatal")) fmt.Println(args...) + r.printData() } -func (r *TestLog) Fatalf(format string, args ...any) { +func (r *TestLogWriter) Fatalf(format string, args ...any) { fmt.Print(prefix("fatal")) fmt.Printf(format+"\n", args...) + r.printData() } -func (r *TestLog) Panic(args ...any) { +func (r *TestLogWriter) Panic(args ...any) { fmt.Print(prefix("panic")) fmt.Println(args...) + r.printData() } -func (r *TestLog) Panicf(format string, args ...any) { +func (r *TestLogWriter) Panicf(format string, args ...any) { fmt.Print(prefix("panic")) fmt.Printf(format+"\n", args...) + r.printData() } -func (r *TestLog) User(user any) log.Writer { +func (r *TestLogWriter) User(user any) log.Writer { + r.data["user"] = user + return r } -func (r *TestLog) Owner(owner any) log.Writer { +func (r *TestLogWriter) Owner(owner any) log.Writer { + r.data["owner"] = owner + return r } -func (r *TestLog) Hint(hint string) log.Writer { +func (r *TestLogWriter) Hint(hint string) log.Writer { + r.data["hint"] = hint + return r } -func (r *TestLog) Code(code string) log.Writer { +func (r *TestLogWriter) Code(code string) log.Writer { + r.data["code"] = code + return r } -func (r *TestLog) With(data map[string]any) log.Writer { +func (r *TestLogWriter) With(data map[string]any) log.Writer { + r.data["with"] = data + return r } -func (r *TestLog) Tags(tags ...string) log.Writer { +func (r *TestLogWriter) Tags(tags ...string) log.Writer { + r.data["tags"] = tags + return r } -func (r *TestLog) Request(req http.ContextRequest) log.Writer { +func (r *TestLogWriter) Request(req http.ContextRequest) log.Writer { + r.data["request"] = req + return r } -func (r *TestLog) Response(res http.ContextResponse) log.Writer { +func (r *TestLogWriter) Response(res http.ContextResponse) log.Writer { + r.data["response"] = res + return r } -func (r *TestLog) In(domain string) log.Writer { +func (r *TestLogWriter) In(domain string) log.Writer { + r.data["in"] = domain + return r } +func (r *TestLogWriter) printData() { + if len(r.data) > 0 { + fmt.Println(r.data) + } +} + func prefix(model string) string { timestamp := carbon.Now().ToDateTimeString()