Files
cpone_midleware/internal/http/middleware_database_test.go

87 lines
3.1 KiB
Go

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())
}
}