From 1820186611b68d7a12d3c2d040ac596c0d33aaf3 Mon Sep 17 00:00:00 2001 From: JhonatanEstabile Date: Mon, 7 Sep 2026 12:37:21 -0300 Subject: [PATCH 1/2] propagating request context to the database layer --- app/controller/base_controller.go | 43 ++++--- app/helper/context_helper.go | 10 ++ app/repository/base_repository.go | 105 ++++++++++++++---- app/repository/repository_actions.go | 35 +++--- tests/unit/controller/base_controller_test.go | 73 ++++++++++-- 5 files changed, 203 insertions(+), 63 deletions(-) create mode 100644 app/helper/context_helper.go diff --git a/app/controller/base_controller.go b/app/controller/base_controller.go index ce71026..9525206 100644 --- a/app/controller/base_controller.go +++ b/app/controller/base_controller.go @@ -63,7 +63,9 @@ func (bc *BaseController[T]) Add(w http.ResponseWriter, r *http.Request) { u.SetUpdatedAt(now) } - if err := bc.Repo.Add(m); err != nil { + ctx := helper.GetContextWithoutCancel(r) + + if err := bc.Repo.AddContext(ctx, m); err != nil { helper.JSONError(w, http.StatusInternalServerError, "Insert error", err) return } @@ -93,7 +95,9 @@ func (bc *BaseController[T]) Bulk(w http.ResponseWriter, r *http.Request) { } fields := helper.GetFieldsParamList(r, bc.Repo.New().Columns(), orderBy) - list, err := bc.Repo.Bulk(input.IDs, limit, pageCursor, orderBy, order, fields) + ctx := helper.GetContextWithoutCancel(r) + + list, err := bc.Repo.BulkContext(ctx, input.IDs, limit, pageCursor, orderBy, order, fields) if err != nil { helper.JSONError(w, http.StatusInternalServerError, "Bulk error", err) return @@ -149,7 +153,9 @@ func (bc *BaseController[T]) BulkAdd(w http.ResponseWriter, r *http.Request) { } } - if err := bc.Repo.BulkAdd(items); err != nil { + ctx := helper.GetContextWithoutCancel(r) + + if err := bc.Repo.BulkAddContext(ctx, items); err != nil { helper.JSONError(w, http.StatusInternalServerError, "Bulk insert failed", err) return } @@ -171,8 +177,9 @@ func (bc *BaseController[T]) DeadDetail(w http.ResponseWriter, r *http.Request) return } + ctx := helper.GetContextWithoutCancel(r) fields := helper.GetFieldsParamOne(r, bc.Repo.New().Columns()) - m, err := bc.Repo.DeadDetail(id, fields) + m, err := bc.Repo.DeadDetailContext(ctx, id, fields) if err != nil { helper.JSONError(w, http.StatusNotFound, "Detail error", err) return @@ -196,7 +203,9 @@ func (bc *BaseController[T]) DeadList(w http.ResponseWriter, r *http.Request) { fields := helper.GetFieldsParamList(r, bc.Repo.New().Columns(), orderBy) filters := helper.GetFilters(r, bc.Repo.New().Columns()) - list, err := bc.Repo.DeadList(limit, pageCursor, orderBy, order, fields, filters) + ctx := helper.GetContextWithoutCancel(r) + + list, err := bc.Repo.DeadListContext(ctx, limit, pageCursor, orderBy, order, fields, filters) if err != nil { helper.JSONError(w, http.StatusInternalServerError, "List error", err) return @@ -217,10 +226,11 @@ func (bc *BaseController[T]) Delete(w http.ResponseWriter, r *http.Request) { return } + ctx := helper.GetContextWithoutCancel(r) m := bc.Repo.New() bc.SetPK(m, id) - if err := bc.Repo.Delete(m); err != nil { + if err := bc.Repo.DeleteContext(ctx, m); err != nil { helper.JSONError(w, http.StatusInternalServerError, "Delete error", err) return } @@ -240,8 +250,9 @@ func (bc *BaseController[T]) Detail(w http.ResponseWriter, r *http.Request) { return } + ctx := helper.GetContextWithoutCancel(r) fields := helper.GetFieldsParamOne(r, bc.Repo.New().Columns()) - m, err := bc.Repo.Detail(id, fields) + m, err := bc.Repo.DetailContext(ctx, id, fields) if err != nil { helper.JSONError(w, http.StatusNotFound, "Detail error", err) return @@ -268,7 +279,9 @@ func (bc *BaseController[T]) Edit(w http.ResponseWriter, r *http.Request) { return } - fetched, err := bc.Repo.Detail(id, bc.Repo.New().Columns()) + ctx := helper.GetContextWithoutCancel(r) + + fetched, err := bc.Repo.DetailContext(ctx, id, bc.Repo.New().Columns()) if err != nil { helper.JSONError(w, http.StatusNotFound, "Not found", err) return @@ -298,7 +311,7 @@ func (bc *BaseController[T]) Edit(w http.ResponseWriter, r *http.Request) { m := bc.Repo.New() bc.SetPK(m, id) - if err := bc.Repo.Edit(m.TableName(), m.PrimaryKey(), m.PrimaryKeyValue(), updateCols, updateVals); err != nil { + if err := bc.Repo.EditContext(ctx, m.TableName(), m.PrimaryKey(), m.PrimaryKeyValue(), updateCols, updateVals); err != nil { helper.JSONError(w, http.StatusInternalServerError, "Edit error", err) return } @@ -312,6 +325,7 @@ func (bc *BaseController[T]) List(w http.ResponseWriter, r *http.Request) { return } + ctx := helper.GetContextWithoutCancel(r) orderBy, order := helper.GetOrderParams(r, "id") limit, pageCursor, err := helper.GetPaginationParams(r) if err != nil { @@ -321,7 +335,7 @@ func (bc *BaseController[T]) List(w http.ResponseWriter, r *http.Request) { fields := helper.GetFieldsParamList(r, bc.Repo.New().Columns(), orderBy) filters := helper.GetFilters(r, bc.Repo.New().Columns()) - list, err := bc.Repo.List(limit, pageCursor, orderBy, order, fields, filters) + list, err := bc.Repo.ListContext(ctx, limit, pageCursor, orderBy, order, fields, filters) if err != nil { helper.JSONError(w, http.StatusInternalServerError, "List error", err) return @@ -336,11 +350,12 @@ func (bc *BaseController[T]) ListOne(w http.ResponseWriter, r *http.Request) { return } + ctx := helper.GetContextWithoutCancel(r) orderBy, order := helper.GetOrderParams(r, "id") fields := helper.GetFieldsParamOne(r, bc.Repo.New().Columns()) filters := helper.GetFilters(r, bc.Repo.New().Columns()) - result, err := bc.Repo.ListOne(orderBy, order, fields, filters) + result, err := bc.Repo.ListOneContext(ctx, orderBy, order, fields, filters) if err != nil { helper.JSONError(w, http.StatusInternalServerError, "List one error", err) return @@ -387,7 +402,8 @@ func (bc *BaseController[T]) Raw(w http.ResponseWriter, r *http.Request) { return } - results, err := bc.Repo.Raw(sqlText, input.Params) + ctx := helper.GetContextWithoutCancel(r) + results, err := bc.Repo.RawContext(ctx, sqlText, input.Params) if err != nil { helper.JSONError(w, http.StatusInternalServerError, "Raw execution failed", err) return @@ -411,7 +427,8 @@ func (bc *BaseController[T]) Undelete(w http.ResponseWriter, r *http.Request) { m := bc.Repo.New() bc.SetPK(m, id) - if err := bc.Repo.Undelete(m); err != nil { + ctx := helper.GetContextWithoutCancel(r) + if err := bc.Repo.UndeleteContext(ctx, m); err != nil { helper.JSONError(w, http.StatusInternalServerError, "Undelete error", err) return } diff --git a/app/helper/context_helper.go b/app/helper/context_helper.go new file mode 100644 index 0000000..a9d6577 --- /dev/null +++ b/app/helper/context_helper.go @@ -0,0 +1,10 @@ +package helper + +import ( + "context" + "net/http" +) + +func GetContextWithoutCancel(r *http.Request) context.Context { + return context.WithoutCancel(r.Context()) +} diff --git a/app/repository/base_repository.go b/app/repository/base_repository.go index 861c767..73032fd 100644 --- a/app/repository/base_repository.go +++ b/app/repository/base_repository.go @@ -1,6 +1,7 @@ package repository import ( + "context" "database/sql" "time" @@ -48,6 +49,18 @@ type RepositoryInterface[T BaseModel] interface { ListOne(orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) Raw(query string, params map[string]any) ([]map[string]any, error) Undelete(m T) error + AddContext(ctx context.Context, m T) error + BulkContext(ctx context.Context, ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) + BulkAddContext(ctx context.Context, models []T) error + DeadDetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) + DeadListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) + DeleteContext(ctx context.Context, m T) error + DetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) + EditContext(ctx context.Context, table, pk string, pkVal interface{}, cols []string, vals []interface{}) error + ListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) + ListOneContext(ctx context.Context, orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) + RawContext(ctx context.Context, query string, params map[string]any) ([]map[string]any, error) + UndeleteContext(ctx context.Context, m T) error } type Repository[T BaseModel] struct { @@ -67,64 +80,112 @@ func (r *Repository[T]) New() T { } func (r *Repository[T]) Add(m T) error { - return addRecord(r.DB, m) + return r.AddContext(context.Background(), m) } func (r *Repository[T]) BulkAdd(m []T) error { + return r.BulkAddContext(context.Background(), m) +} + +func (r *Repository[T]) Bulk(ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { + return r.BulkContext(context.Background(), ids, limit, pageCursor, orderBy, order, fields) +} + +func (r *Repository[T]) DeadDetail(id interface{}, fields []string) (map[string]any, error) { + return r.DeadDetailContext(context.Background(), id, fields) +} + +func (r *Repository[T]) DeadList(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { + return r.DeadListContext(context.Background(), limit, pageCursor, orderBy, order, fields, filters) +} + +func (r *Repository[T]) Delete(m T) error { + return r.DeleteContext(context.Background(), m) +} + +func (r *Repository[T]) Detail(id interface{}, fields []string) (map[string]any, error) { + return r.DetailContext(context.Background(), id, fields) +} + +func (r *Repository[T]) Edit(table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { + return r.EditContext(context.Background(), table, pk, pkVal, cols, vals) +} + +func (r *Repository[T]) List(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { + return r.ListContext(context.Background(), limit, pageCursor, orderBy, order, fields, filters) +} + +func (r *Repository[T]) ListOne(orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { + return r.ListOneContext(context.Background(), orderBy, order, fields, filters) +} + +func (r *Repository[T]) Raw(query string, params map[string]any) ([]map[string]any, error) { + return r.RawContext(context.Background(), query, params) +} + +func (r *Repository[T]) Undelete(m T) error { + return r.UndeleteContext(context.Background(), m) +} + +func (r *Repository[T]) AddContext(ctx context.Context, m T) error { + return addRecord(ctx, r.DB, m) +} + +func (r *Repository[T]) BulkAddContext(ctx context.Context, m []T) error { baseModels := make([]BaseModel, len(m)) for i, model := range m { baseModels[i] = model } - return bulkAddRecords(r.DB, baseModels) + return bulkAddRecords(ctx, r.DB, baseModels) } -func (r *Repository[T]) Bulk(ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { +func (r *Repository[T]) BulkContext(ctx context.Context, ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { m := r.New() - return bulkRecords(r.DB, m.Schema(), m.TableName(), m.PrimaryKey(), fields, ids, limit, pageCursor, orderBy, order) + return bulkRecords(ctx, r.DB, m.Schema(), m.TableName(), m.PrimaryKey(), fields, ids, limit, pageCursor, orderBy, order) } -func (r *Repository[T]) DeadDetail(id interface{}, fields []string) (map[string]any, error) { +func (r *Repository[T]) DeadDetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) { m := r.New() - return getRecord(r.DB, id, m.Schema(), m.TableName(), m.PrimaryKey(), fields, true) + return getRecord(ctx, r.DB, id, m.Schema(), m.TableName(), m.PrimaryKey(), fields, true) } -func (r *Repository[T]) DeadList(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { +func (r *Repository[T]) DeadListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { m := r.New() - return listRecords(r.DB, m.Schema(), m.TableName(), fields, limit, pageCursor, orderBy, order, filters, true) + return listRecords(ctx, r.DB, m.Schema(), m.TableName(), fields, limit, pageCursor, orderBy, order, filters, true) } -func (r *Repository[T]) Delete(m T) error { - return deleteRecord(r.DB, m.TableName(), m.PrimaryKey(), m.PrimaryKeyValue()) +func (r *Repository[T]) DeleteContext(ctx context.Context, m T) error { + return deleteRecord(ctx, r.DB, m.TableName(), m.PrimaryKey(), m.PrimaryKeyValue()) } -func (r *Repository[T]) Detail(id interface{}, fields []string) (map[string]any, error) { +func (r *Repository[T]) DetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) { m := r.New() - return getRecord(r.DB, id, m.Schema(), m.TableName(), m.PrimaryKey(), fields, false) + return getRecord(ctx, r.DB, id, m.Schema(), m.TableName(), m.PrimaryKey(), fields, false) } -func (r *Repository[T]) Edit(table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { - return editRecord(r.DB, table, pk, pkVal, cols, vals) +func (r *Repository[T]) EditContext(ctx context.Context, table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { + return editRecord(ctx, r.DB, table, pk, pkVal, cols, vals) } -func (r *Repository[T]) List(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { +func (r *Repository[T]) ListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { m := r.New() - return listRecords(r.DB, m.Schema(), m.TableName(), fields, limit, pageCursor, orderBy, order, filters, false) + return listRecords(ctx, r.DB, m.Schema(), m.TableName(), fields, limit, pageCursor, orderBy, order, filters, false) } -func (r *Repository[T]) ListOne(orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { - results, err := r.List(1, nil, orderBy, order, fields, filters) +func (r *Repository[T]) ListOneContext(ctx context.Context, orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { + results, err := r.ListContext(ctx, 1, nil, orderBy, order, fields, filters) if len(results) == 0 { return make(map[string]any), err } return results[0], err } -func (r *Repository[T]) Raw(query string, params map[string]any) ([]map[string]any, error) { +func (r *Repository[T]) RawContext(ctx context.Context, query string, params map[string]any) ([]map[string]any, error) { m := r.New() sqlText, args := helper.PrepareRawQuery(query, params) - return rawRecords(r.DB, m.Schema(), sqlText, args...) + return rawRecords(ctx, r.DB, m.Schema(), sqlText, args...) } -func (r *Repository[T]) Undelete(m T) error { - return undeleteRecord(r.DB, m.TableName(), m.PrimaryKey(), m.PrimaryKeyValue()) +func (r *Repository[T]) UndeleteContext(ctx context.Context, m T) error { + return undeleteRecord(ctx, r.DB, m.TableName(), m.PrimaryKey(), m.PrimaryKeyValue()) } diff --git a/app/repository/repository_actions.go b/app/repository/repository_actions.go index a11195b..5800466 100644 --- a/app/repository/repository_actions.go +++ b/app/repository/repository_actions.go @@ -1,6 +1,7 @@ package repository import ( + "context" "database/sql" "fmt" "strings" @@ -10,7 +11,7 @@ import ( var ScanFunc = helper.GenericScanToMap -func addRecord(db *sql.DB, m BaseModel) error { +func addRecord(ctx context.Context, db *sql.DB, m BaseModel) error { allCols := m.Columns() allVals := m.Values() @@ -30,11 +31,12 @@ func addRecord(db *sql.DB, m BaseModel) error { strings.Join(placeholders, ", "), ) - _, err := db.Exec(query, finalVals...) + _, err := db.ExecContext(ctx, query, finalVals...) return err } func bulkRecords( + ctx context.Context, db *sql.DB, schema map[string]string, table string, @@ -101,7 +103,7 @@ func bulkRecords( orderExpr, ) - rows, err := db.Query(query, args...) + rows, err := db.QueryContext(ctx, query, args...) if err != nil { return nil, err } @@ -118,7 +120,7 @@ func bulkRecords( return list, nil } -func bulkAddRecords(db *sql.DB, m []BaseModel) error { +func bulkAddRecords(ctx context.Context, db *sql.DB, m []BaseModel) error { first := m[0] table := first.TableName() allCols := first.Columns() @@ -143,21 +145,21 @@ func bulkAddRecords(db *sql.DB, m []BaseModel) error { strings.Join(rowsSQL, ", "), ) - _, err := db.Exec(query, args...) + _, err := db.ExecContext(ctx, query, args...) return err } -func deleteRecord(db *sql.DB, table, pk string, pkVal interface{}) error { +func deleteRecord(ctx context.Context, db *sql.DB, table, pk string, pkVal interface{}) error { query := fmt.Sprintf( "UPDATE %s SET `deleted_at` = NOW() WHERE `%s` = ? AND `deleted_at` IS NULL", table, pk, ) - _, err := db.Exec(query, pkVal) + _, err := db.ExecContext(ctx, query, pkVal) return err } -func editRecord(db *sql.DB, table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { +func editRecord(ctx context.Context, db *sql.DB, table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { if len(cols) == 0 { return nil } @@ -175,11 +177,11 @@ func editRecord(db *sql.DB, table, pk string, pkVal interface{}, cols []string, ) vals = append(vals, pkVal) - _, err := db.Exec(query, vals...) + _, err := db.ExecContext(ctx, query, vals...) return err } -func getRecord(db *sql.DB, id interface{}, schema map[string]string, table string, pk string, fields []string, deleted bool) (map[string]any, error) { +func getRecord(ctx context.Context, db *sql.DB, id interface{}, schema map[string]string, table string, pk string, fields []string, deleted bool) (map[string]any, error) { selected := helper.EscapeMysqlFields(helper.FilterFields(fields, helper.MapKeys(schema))) condition := "`deleted_at` IS NULL" if deleted { @@ -194,7 +196,7 @@ func getRecord(db *sql.DB, id interface{}, schema map[string]string, table strin condition, ) - rows, err := db.Query(query, id) + rows, err := db.QueryContext(ctx, query, id) if err != nil { return nil, err } @@ -207,6 +209,7 @@ func getRecord(db *sql.DB, id interface{}, schema map[string]string, table strin } func listRecords( + ctx context.Context, db *sql.DB, schema map[string]string, table string, @@ -267,7 +270,7 @@ func listRecords( ) args = append(args, limit) - rows, err := db.Query(query, args...) + rows, err := db.QueryContext(ctx, query, args...) if err != nil { return nil, err } @@ -284,8 +287,8 @@ func listRecords( return list, nil } -func rawRecords(db *sql.DB, _ map[string]string, sqlText string, args ...interface{}) ([]map[string]any, error) { - rows, err := db.Query(sqlText, args...) +func rawRecords(ctx context.Context, db *sql.DB, _ map[string]string, sqlText string, args ...interface{}) ([]map[string]any, error) { + rows, err := db.QueryContext(ctx, sqlText, args...) if err != nil { return nil, err } @@ -294,12 +297,12 @@ func rawRecords(db *sql.DB, _ map[string]string, sqlText string, args ...interfa return helper.SimpleScanRows(helper.NewRowsAdapter(rows)) } -func undeleteRecord(db *sql.DB, table, pk string, pkVal interface{}) error { +func undeleteRecord(ctx context.Context, db *sql.DB, table, pk string, pkVal interface{}) error { query := fmt.Sprintf( "UPDATE %s SET `deleted_at` = NULL WHERE `%s` = ? AND `deleted_at` IS NOT NULL", table, pk, ) - _, err := db.Exec(query, pkVal) + _, err := db.ExecContext(ctx, query, pkVal) return err } diff --git a/tests/unit/controller/base_controller_test.go b/tests/unit/controller/base_controller_test.go index 806f1c1..ac86673 100644 --- a/tests/unit/controller/base_controller_test.go +++ b/tests/unit/controller/base_controller_test.go @@ -2,6 +2,7 @@ package controller import ( "bytes" + "context" "encoding/json" "errors" "io" @@ -105,60 +106,108 @@ func (fr *fakeRepository) New() *fakeModel { return &fakeModel{} } -func (fr *fakeRepository) Add(m *fakeModel) error { +func (fr *fakeRepository) AddContext(ctx context.Context, m *fakeModel) error { fr.insertedModel = m return fr.insertedError } -func (fr *fakeRepository) Edit(table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { +func (fr *fakeRepository) EditContext(ctx context.Context, table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { fr.updateFieldsCalled = true fr.updateFieldsCols = cols fr.updateFieldsVals = vals return fr.updateFieldsError } -func (fr *fakeRepository) Delete(m *fakeModel) error { +func (fr *fakeRepository) DeleteContext(ctx context.Context, m *fakeModel) error { fr.deleteCalled = true return fr.deleteError } -func (fr *fakeRepository) Undelete(m *fakeModel) error { +func (fr *fakeRepository) UndeleteContext(ctx context.Context, m *fakeModel) error { fr.undeleteCalled = true return fr.deleteError } -func (fr *fakeRepository) Detail(id interface{}, fields []string) (map[string]any, error) { +func (fr *fakeRepository) DetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) { return fr.getResult, fr.getError } -func (fr *fakeRepository) DeadDetail(id interface{}, fields []string) (map[string]any, error) { +func (fr *fakeRepository) DeadDetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) { return fr.getDeletedResult, fr.getDeletedError } -func (fr *fakeRepository) List(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { +func (fr *fakeRepository) ListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { return fr.listActiveResult, fr.listActiveError } -func (fr *fakeRepository) DeadList(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { +func (fr *fakeRepository) DeadListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { return fr.listDeletedResult, fr.listDeletedError } -func (fr *fakeRepository) Bulk(ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { +func (fr *fakeRepository) BulkContext(ctx context.Context, ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { return fr.bulkGetResult, fr.bulkGetError } -func (fr *fakeRepository) ListOne(orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { +func (fr *fakeRepository) ListOneContext(ctx context.Context, orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { return fr.listOneResult, fr.listOneError } -func (fr *fakeRepository) Raw(query string, params map[string]any) ([]map[string]any, error) { +func (fr *fakeRepository) RawContext(ctx context.Context, query string, params map[string]any) ([]map[string]any, error) { return fr.rawResult, fr.rawError } -func (fr *fakeRepository) BulkAdd(m []*fakeModel) error { +func (fr *fakeRepository) BulkAddContext(ctx context.Context, m []*fakeModel) error { return fr.bulkAddError } +func (fr *fakeRepository) Add(m *fakeModel) error { + return fr.AddContext(context.Background(), m) +} + +func (fr *fakeRepository) Edit(table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { + return fr.EditContext(context.Background(), table, pk, pkVal, cols, vals) +} + +func (fr *fakeRepository) Delete(m *fakeModel) error { + return fr.DeleteContext(context.Background(), m) +} + +func (fr *fakeRepository) Undelete(m *fakeModel) error { + return fr.UndeleteContext(context.Background(), m) +} + +func (fr *fakeRepository) Detail(id interface{}, fields []string) (map[string]any, error) { + return fr.DetailContext(context.Background(), id, fields) +} + +func (fr *fakeRepository) DeadDetail(id interface{}, fields []string) (map[string]any, error) { + return fr.DeadDetailContext(context.Background(), id, fields) +} + +func (fr *fakeRepository) List(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { + return fr.ListContext(context.Background(), limit, pageCursor, orderBy, order, fields, filters) +} + +func (fr *fakeRepository) DeadList(limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { + return fr.DeadListContext(context.Background(), limit, pageCursor, orderBy, order, fields, filters) +} + +func (fr *fakeRepository) Bulk(ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { + return fr.BulkContext(context.Background(), ids, limit, pageCursor, orderBy, order, fields) +} + +func (fr *fakeRepository) ListOne(orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { + return fr.ListOneContext(context.Background(), orderBy, order, fields, filters) +} + +func (fr *fakeRepository) Raw(query string, params map[string]any) ([]map[string]any, error) { + return fr.RawContext(context.Background(), query, params) +} + +func (fr *fakeRepository) BulkAdd(m []*fakeModel) error { + return fr.BulkAddContext(context.Background(), m) +} + func TestNewBaseController(t *testing.T) { fr := &fakeRepository{} From 866064b7d88de5b83d3c92ec5050ef19cc89ada9 Mon Sep 17 00:00:00 2001 From: JhonatanEstabile Date: Mon, 7 Sep 2026 13:50:25 -0300 Subject: [PATCH 2/2] including tests to assert that context arrives at the repository layer --- .../controller/context_propagation_test.go | 241 ++++++++++++++++++ 1 file changed, 241 insertions(+) create mode 100644 tests/unit/controller/context_propagation_test.go diff --git a/tests/unit/controller/context_propagation_test.go b/tests/unit/controller/context_propagation_test.go new file mode 100644 index 0000000..47a8dc4 --- /dev/null +++ b/tests/unit/controller/context_propagation_test.go @@ -0,0 +1,241 @@ +package controller + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/not-empty/grit-microframework-go/app/controller" + "github.com/not-empty/grit-microframework-go/app/helper" + "github.com/stretchr/testify/require" + + ulidmock "github.com/not-empty/ulid-go-lib/mock" +) + +type ctxProbeKey string + +const ctxProbe ctxProbeKey = "ctx-probe" + +type ctxSpyRepository struct { + *fakeRepository + seen map[string]context.Context +} + +func newCtxSpy() *ctxSpyRepository { + return &ctxSpyRepository{ + fakeRepository: &fakeRepository{ + getResult: map[string]any{"id": "1", "field": "oldValue"}, + getDeletedResult: map[string]any{"id": "1", "field": "oldValue"}, + listActiveResult: []map[string]any{{"id": "1", "field": "value"}}, + listDeletedResult: []map[string]any{{"id": "1", "field": "value"}}, + bulkGetResult: []map[string]any{{"id": "1", "field": "value"}}, + listOneResult: map[string]any{"id": "1", "field": "value"}, + rawResult: []map[string]any{{"id": "1"}}, + }, + seen: map[string]context.Context{}, + } +} + +func (s *ctxSpyRepository) AddContext(ctx context.Context, m *fakeModel) error { + s.seen["AddContext"] = ctx + return s.fakeRepository.AddContext(ctx, m) +} + +func (s *ctxSpyRepository) BulkAddContext(ctx context.Context, m []*fakeModel) error { + s.seen["BulkAddContext"] = ctx + return s.fakeRepository.BulkAddContext(ctx, m) +} + +func (s *ctxSpyRepository) BulkContext(ctx context.Context, ids []string, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string) ([]map[string]any, error) { + s.seen["BulkContext"] = ctx + return s.fakeRepository.BulkContext(ctx, ids, limit, pageCursor, orderBy, order, fields) +} + +func (s *ctxSpyRepository) DeadDetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) { + s.seen["DeadDetailContext"] = ctx + return s.fakeRepository.DeadDetailContext(ctx, id, fields) +} + +func (s *ctxSpyRepository) DeadListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { + s.seen["DeadListContext"] = ctx + return s.fakeRepository.DeadListContext(ctx, limit, pageCursor, orderBy, order, fields, filters) +} + +func (s *ctxSpyRepository) DeleteContext(ctx context.Context, m *fakeModel) error { + s.seen["DeleteContext"] = ctx + return s.fakeRepository.DeleteContext(ctx, m) +} + +func (s *ctxSpyRepository) DetailContext(ctx context.Context, id interface{}, fields []string) (map[string]any, error) { + s.seen["DetailContext"] = ctx + return s.fakeRepository.DetailContext(ctx, id, fields) +} + +func (s *ctxSpyRepository) EditContext(ctx context.Context, table, pk string, pkVal interface{}, cols []string, vals []interface{}) error { + s.seen["EditContext"] = ctx + return s.fakeRepository.EditContext(ctx, table, pk, pkVal, cols, vals) +} + +func (s *ctxSpyRepository) ListContext(ctx context.Context, limit int, pageCursor *helper.PageCursor, orderBy, order string, fields []string, filters []helper.Filter) ([]map[string]any, error) { + s.seen["ListContext"] = ctx + return s.fakeRepository.ListContext(ctx, limit, pageCursor, orderBy, order, fields, filters) +} + +func (s *ctxSpyRepository) ListOneContext(ctx context.Context, orderBy, order string, fields []string, filters []helper.Filter) (map[string]any, error) { + s.seen["ListOneContext"] = ctx + return s.fakeRepository.ListOneContext(ctx, orderBy, order, fields, filters) +} + +func (s *ctxSpyRepository) RawContext(ctx context.Context, query string, params map[string]any) ([]map[string]any, error) { + s.seen["RawContext"] = ctx + return s.fakeRepository.RawContext(ctx, query, params) +} + +func (s *ctxSpyRepository) UndeleteContext(ctx context.Context, m *fakeModel) error { + s.seen["UndeleteContext"] = ctx + return s.fakeRepository.UndeleteContext(ctx, m) +} + +func newCtxSpyController(spy *ctxSpyRepository) *controller.BaseController[*fakeModel] { + return &controller.BaseController[*fakeModel]{ + Repo: spy, + Prefix: "/fake", + SetPK: func(m *fakeModel, id string) { m.ID = id }, + ULIDGen: &ulidmock.ULIDMock{ + GenerateFunc: func(ts int64) (string, error) { + return "00000000000000000000000000000", nil + }, + }, + } +} + +func mustJSON(t *testing.T, v any) *bytes.Buffer { + t.Helper() + b, err := json.Marshal(v) + require.NoError(t, err) + return bytes.NewBuffer(b) +} + +func TestBaseController_RequestContextReachesRepository(t *testing.T) { + helper.RegisterRawQueries("fake", map[string]string{"probe": "SELECT id FROM fake"}) + + cases := []struct { + name string + method string + expect string + body func(t *testing.T) *bytes.Buffer + call func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) + }{ + { + name: "Add", method: http.MethodPost, expect: "AddContext", + body: func(t *testing.T) *bytes.Buffer { return mustJSON(t, map[string]string{"field": "value"}) }, + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Add(rr, r) + }, + }, + { + name: "BulkAdd", method: http.MethodPost, expect: "BulkAddContext", + body: func(t *testing.T) *bytes.Buffer { + return mustJSON(t, []map[string]string{{"field": "value"}}) + }, + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.BulkAdd(rr, r) + }, + }, + { + name: "Bulk", method: http.MethodPost, expect: "BulkContext", + body: func(t *testing.T) *bytes.Buffer { + return mustJSON(t, map[string]any{"ids": []string{"1", "2"}}) + }, + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Bulk(rr, r) + }, + }, + { + name: "DeadDetail", method: http.MethodGet, expect: "DeadDetailContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.DeadDetail(rr, r) + }, + }, + { + name: "DeadList", method: http.MethodGet, expect: "DeadListContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.DeadList(rr, r) + }, + }, + { + name: "Delete", method: http.MethodDelete, expect: "DeleteContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Delete(rr, r) + }, + }, + { + name: "Detail", method: http.MethodGet, expect: "DetailContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Detail(rr, r) + }, + }, + { + name: "Edit", method: http.MethodPatch, expect: "EditContext", + body: func(t *testing.T) *bytes.Buffer { return mustJSON(t, map[string]any{"field": "newValue"}) }, + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Edit(rr, r) + }, + }, + { + name: "List", method: http.MethodGet, expect: "ListContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.List(rr, r) + }, + }, + { + name: "ListOne", method: http.MethodGet, expect: "ListOneContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.ListOne(rr, r) + }, + }, + { + name: "Raw", method: http.MethodPost, expect: "RawContext", + body: func(t *testing.T) *bytes.Buffer { + return mustJSON(t, map[string]any{"query": "probe", "params": map[string]any{}}) + }, + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Raw(rr, r) + }, + }, + { + name: "Undelete", method: http.MethodPatch, expect: "UndeleteContext", + call: func(bc *controller.BaseController[*fakeModel], rr *httptest.ResponseRecorder, r *http.Request) { + bc.Undelete(rr, r) + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + spy := newCtxSpy() + bc := newCtxSpyController(spy) + + var req *http.Request + if tc.body != nil { + req = httptest.NewRequest(tc.method, "/fake/probe/1", tc.body(t)) + } else { + req = httptest.NewRequest(tc.method, "/fake/probe/1", nil) + } + req = req.WithContext(context.WithValue(req.Context(), ctxProbe, "probe-value")) + + tc.call(bc, httptest.NewRecorder(), req) + + ctx, ok := spy.seen[tc.expect] + require.True(t, ok, "%s should call %s so the request context reaches the database layer", tc.name, tc.expect) + require.Equal(t, "probe-value", ctx.Value(ctxProbe), "request context values must survive until the repository") + require.Nil(t, ctx.Done(), "context must not be cancelable, GetContextWithoutCancel strips cancelation") + + deadline, hasDeadline := ctx.Deadline() + require.False(t, hasDeadline, "context must not carry a deadline, got %v", deadline) + }) + } +}