package funcspec import ( "context" "errors" "net/http" "net/http/httptest" "strings" "testing" "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/security" ) // txFactory returns a pool whose every transaction is a distinct MockDatabase. func txFactory(queries *int) *MockDatabase { return &MockDatabase{ RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error { return fn(&MockDatabase{ QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error { *queries++ if rows, ok := dest.(*[]map[string]interface{}); ok { *rows = []map[string]interface{}{{"id": float64(1)}} } return nil }, }) }, } } func TestOnTxBeginFirstAndBeforeResponseOnSecondTx(t *testing.T) { var queries int h := NewHandler(txFactory(&queries)) var order []string txs := map[HookType][]common.Database{} for _, ht := range []HookType{OnTxBegin, BeforeQuery, AfterQuery, BeforeResponse} { ht := ht h.Hooks().Register(ht, func(c *HookContext) error { order = append(order, string(ht)) txs[ht] = append(txs[ht], c.Tx) return nil }) } w := httptest.NewRecorder() h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{})(w, createTestRequest("GET", "/t", nil, nil, nil)) if w.Code != http.StatusOK { t.Fatalf("status %d body %s", w.Code, w.Body) } if got := strings.Join(order, ","); got != "on_tx_begin,before_query,after_query,on_tx_begin,before_response" { t.Fatalf("hook order %s", got) } if txs[BeforeQuery][0] != txs[OnTxBegin][0] || txs[AfterQuery][0] != txs[OnTxBegin][0] { t.Fatal("query hooks must run on the OnTxBegin transaction") } if txs[OnTxBegin][0] == txs[OnTxBegin][1] || txs[BeforeResponse][0] != txs[OnTxBegin][1] { t.Fatal("BeforeResponse must run on a second, distinct transaction") } } func TestOnTxBeginListBeforeResponseOnSecondTx(t *testing.T) { var queries int h := NewHandler(txFactory(&queries)) var txs []common.Database h.Hooks().Register(OnTxBegin, func(c *HookContext) error { txs = append(txs, c.Tx); return nil }) var respTx common.Database h.Hooks().Register(BeforeResponse, func(c *HookContext) error { respTx = c.Tx; return nil }) w := httptest.NewRecorder() h.SqlQueryList("SELECT * FROM users", SqlQueryOptions{NoCount: true})(w, createTestRequest("GET", "/t", nil, nil, nil)) if w.Code != http.StatusOK { t.Fatalf("status %d body %s", w.Code, w.Body) } if len(txs) != 2 || txs[0] == txs[1] || respTx != txs[1] { t.Fatalf("expected 2 distinct transactions with BeforeResponse on the second, got %d", len(txs)) } } func TestOnTxBeginErrorAnswersTransactionError(t *testing.T) { var queries int h := NewHandler(txFactory(&queries)) h.Hooks().Register(OnTxBegin, func(*HookContext) error { return errors.New("secret detail") }) for name, run := range map[string]HTTPFuncType{ "single": h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{}), "list": h.SqlQueryList("SELECT * FROM users", SqlQueryOptions{NoCount: true}), } { w := httptest.NewRecorder() run(w, createTestRequest("GET", "/t", nil, nil, nil)) if w.Code != http.StatusInternalServerError { t.Fatalf("%s: status %d body %s", name, w.Code, w.Body) } if strings.Contains(w.Body.String(), "secret detail") { t.Fatalf("%s: hook error leaked to client: %s", name, w.Body) } } if queries != 0 { t.Fatalf("no query may run after a failed OnTxBegin, ran %d", queries) } } type stubProvider struct{ security.SecurityProvider } func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) { var execs []string h := NewHandler(&MockDatabase{ RunInTransactionFunc: func(ctx context.Context, fn func(common.Database) error) error { return fn(&MockDatabase{ ExecFunc: func(ctx context.Context, query string, args ...interface{}) (common.Result, error) { execs = append(execs, query) return &MockResult{}, nil }, QueryFunc: func(ctx context.Context, dest interface{}, query string, args ...interface{}) error { execs = append(execs, "QUERY") if rows, ok := dest.(*[]map[string]interface{}); ok { *rows = []map[string]interface{}{{"id": float64(1)}} } return nil }, }) }, }) list, err := security.NewSecurityList(stubProvider{}) if err != nil { t.Fatal(err) } list.SetTxSettings(func(security.SecurityContext) (map[string]string, error) { return map[string]string{"app.user_id": "1"}, nil }) RegisterSecurityHooks(h, list) w := httptest.NewRecorder() h.SqlQuery("SELECT * FROM users WHERE id = 1", SqlQueryOptions{})(w, createTestRequest("GET", "/t", nil, nil, nil)) if w.Code != http.StatusOK { t.Fatalf("status %d body %s", w.Code, w.Body) } // tx 1: stamp, then the query; tx 2 (BeforeResponse): stamp again. if len(execs) != 3 || !strings.Contains(execs[0], "set_config('app.user_id'") || execs[1] != "QUERY" || !strings.Contains(execs[2], "set_config('app.user_id'") { t.Fatalf("each transaction must be stamped before any other SQL, got %v", execs) } }