package http import ( "context" "errors" "net/http" "net/http/httptest" "testing" "primaya-api/cpone-middleware/internal/databaseconfig" "primaya-api/cpone-middleware/internal/repository" ) type fakeDatabaseResolver struct { resolvedCode string err error } func (f *fakeDatabaseResolver) Resolve(_ context.Context, code string) (repository.MySQLLayananRepository, error) { f.resolvedCode = code return repository.MySQLLayananRepository{}, f.err } func TestSelectDatabasePriority(t *testing.T) { tests := []struct { name, target, header, want string }{ {name: "header", target: "/resource?kode_rs=QUERY", header: " rs_header ", want: "RS_HEADER"}, {name: "query", target: "/resource?kode_rs=rs_query", want: "RS_QUERY"}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { resolver := &fakeDatabaseResolver{} next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if _, ok := selectedRepository(r.Context()); !ok { t.Fatal("selected repository missing from context") } w.WriteHeader(http.StatusNoContent) }) request := httptest.NewRequest(http.MethodGet, test.target, nil) request.Header.Set("X-RS-Code", test.header) recorder := httptest.NewRecorder() SelectDatabase(resolver, next).ServeHTTP(recorder, request) if recorder.Code != http.StatusNoContent || resolver.resolvedCode != test.want { t.Fatalf("status=%d code=%q, want status=204 code=%q", recorder.Code, resolver.resolvedCode, test.want) } if recorder.Header().Get("X-RS-Code") != test.want { t.Fatalf("response X-RS-Code = %q", recorder.Header().Get("X-RS-Code")) } }) } } func TestSelectDatabaseRequiresRSCode(t *testing.T) { resolver := &fakeDatabaseResolver{} request := httptest.NewRequest(http.MethodGet, "/resource", nil) recorder := httptest.NewRecorder() SelectDatabase(resolver, http.HandlerFunc(func(http.ResponseWriter, *http.Request) { t.Fatal("next handler should not be called") })).ServeHTTP(recorder, request) if recorder.Code != http.StatusUnprocessableEntity { t.Fatalf("status = %d, body=%s", recorder.Code, recorder.Body.String()) } if resolver.resolvedCode != "" { t.Fatalf("resolver called with %q", resolver.resolvedCode) } } func TestSelectDatabaseRejectsWrongInstanceCode(t *testing.T) { resolver := &fakeDatabaseResolver{err: databaseconfig.ErrRSCodeMismatch} request := httptest.NewRequest(http.MethodGet, "/resource?kode_rs=unknown", nil) recorder := httptest.NewRecorder() SelectDatabase(resolver, http.HandlerFunc(func(http.ResponseWriter, *http.Request) { t.Fatal("next handler should not be called") })).ServeHTTP(recorder, request) if recorder.Code != http.StatusForbidden { t.Fatalf("status = %d, body=%s", recorder.Code, recorder.Body.String()) } resolver.err = errors.New("connection refused") recorder = httptest.NewRecorder() SelectDatabase(resolver, http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})).ServeHTTP(recorder, request) if recorder.Code != http.StatusServiceUnavailable { t.Fatalf("status = %d, body=%s", recorder.Code, recorder.Body.String()) } }