package api_test import ( "bytes" "context" "errors" "io" "log/slog" "net" "net/http" "net/http/httptest" "strings" "testing" "sneak.berlin/go/simplexcalc/internal/api" "sneak.berlin/go/simplexcalc/internal/simplex" ) const ( // credential is what the API under test is configured with. credential = "a-credential-for-these-tests" //nolint:gosec // G101: invented for tests bearer = "Bearer " + credential chatsPath = "/api/v1/chats" messagesPath = "/api/v1/chats/3/messages" unauthorized = `{"error":"unauthorized"}` + "\n" ) var errChat = errors.New("sqlite: database is locked at /var/lib/simplexcalc") // fakeClient stands in for the chat client. It answers with contacts, // items or sent, or with an error: contactsErr from Contacts, err from // the others. It remembers what it was asked. type fakeClient struct { contacts []simplex.Contact items []simplex.ChatItem sent simplex.ChatItem contactsErr error err error userID, contactID int64 count int text string hadDeadline bool } func (f *fakeClient) Contacts( ctx context.Context, userID int64, ) ([]simplex.Contact, error) { f.userID = userID _, f.hadDeadline = ctx.Deadline() return f.contacts, f.contactsErr } func (f *fakeClient) ChatItems( ctx context.Context, contactID int64, count int, ) ([]simplex.ChatItem, error) { f.contactID, f.count = contactID, count _, f.hadDeadline = ctx.Deadline() return f.items, f.err } func (f *fakeClient) SendMessage( ctx context.Context, contactID int64, text string, ) (simplex.ChatItem, error) { f.contactID, f.text = contactID, text _, f.hadDeadline = ctx.Deadline() return f.sent, f.err } func newAPI(token string, client api.ChatClient) *http.Server { return api.New(api.Params{ Log: slog.New(slog.DiscardHandler), Client: client, UserID: 1, Port: 8080, Token: token, }) } // request sends srv one request, with the Authorization header auth // unless that is empty. func request( t *testing.T, srv *http.Server, method, path, auth string, ) *httptest.ResponseRecorder { t.Helper() req := httptest.NewRequestWithContext(t.Context(), method, path, nil) if auth != "" { req.Header.Set("Authorization", auth) } rec := httptest.NewRecorder() srv.Handler.ServeHTTP(rec, req) return rec } // TestCredential: only the configured credential, sent as a bearer, // gets in. With none configured, nothing does, an empty bearer // included. func TestCredential(t *testing.T) { t.Parallel() for name, tc := range map[string]struct { token, auth string in bool }{ "right": {credential, bearer, true}, "right, scheme in lower case": {credential, "bearer " + credential, true}, "wrong": {credential, strings.ToUpper(bearer), false}, "missing": {credential, "", false}, "another scheme": {credential, "Basic " + credential, false}, "no scheme": {credential, credential, false}, "empty bearer": {credential, "Bearer ", false}, "none configured": {"", bearer, false}, "none configured, empty bearer": {"", "Bearer ", false}, } { t.Run(name, func(t *testing.T) { t.Parallel() rec := request(t, newAPI(tc.token, &fakeClient{}), http.MethodGet, chatsPath, tc.auth) if tc.in { if rec.Code != http.StatusOK { t.Errorf("status = %d, want 200", rec.Code) } return } if rec.Code != http.StatusUnauthorized { t.Fatalf("status = %d, want 401", rec.Code) } if got := rec.Header().Get("WWW-Authenticate"); got != "Bearer" { t.Errorf("WWW-Authenticate = %q, want Bearer", got) } if rec.Body.String() != unauthorized { t.Errorf("body = %q, want %q", rec.Body.String(), unauthorized) } }) } } // TestNoCredentialWarns: an API without a credential says at startup // that it refuses every request. func TestNoCredentialWarns(t *testing.T) { t.Parallel() var logged bytes.Buffer api.New(api.Params{ Log: slog.New(slog.NewJSONHandler(&logged, nil)), Client: &fakeClient{}, Port: 8080, }) if !strings.Contains(logged.String(), `"level":"WARN"`) || !strings.Contains(logged.String(), "API_TOKEN_FILE is not set") { t.Errorf("log = %q, want a warning that API_TOKEN_FILE is not set", logged.String()) } } // TestNoPathIsExempt: an unknown path or method needs the credential // like everything else, and then gets a JSON error. func TestNoPathIsExempt(t *testing.T) { t.Parallel() srv := newAPI(credential, &fakeClient{}) for _, tc := range []struct { method, path, auth string want int body string }{ {http.MethodGet, "/", "", http.StatusUnauthorized, unauthorized}, { http.MethodGet, "/.well-known/healthcheck", "", http.StatusUnauthorized, unauthorized, }, { http.MethodGet, "/api/v1/nothing", bearer, http.StatusNotFound, `{"error":"not found"}` + "\n", }, { http.MethodPost, chatsPath, bearer, http.StatusMethodNotAllowed, `{"error":"method not allowed"}` + "\n", }, { http.MethodPost, messagesPath, "", http.StatusUnauthorized, unauthorized, }, } { rec := request(t, srv, tc.method, tc.path, tc.auth) if rec.Code != tc.want || rec.Body.String() != tc.body { t.Errorf("%s %s: %d %q, want %d %q", tc.method, tc.path, rec.Code, rec.Body.String(), tc.want, tc.body) } } } // TestOptionsAsterisk: "OPTIONS *" is refused like any other request. // net/http would answer it before the handler, so this request goes to // a running server rather than to its handler. func TestOptionsAsterisk(t *testing.T) { t.Parallel() srv := newAPI("", &fakeClient{}) listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", "127.0.0.1:0") if err != nil { t.Fatal(err) } go func() { _ = srv.Serve(listener) }() t.Cleanup(func() { _ = srv.Close() }) req, err := http.NewRequestWithContext(t.Context(), http.MethodOptions, "http://"+listener.Addr().String(), nil) if err != nil { t.Fatal(err) } // The request line becomes "OPTIONS * HTTP/1.1". req.URL.Opaque = "*" resp, err := http.DefaultClient.Do(req) if err != nil { t.Fatal(err) } defer func() { _ = resp.Body.Close() }() body, err := io.ReadAll(resp.Body) if err != nil { t.Fatal(err) } if resp.StatusCode != http.StatusUnauthorized || string(body) != unauthorized { t.Errorf("OPTIONS *: %d %q, want 401 %q", resp.StatusCode, body, unauthorized) } } // TestHeaders: every response, whatever its status, carries the // security headers, and none lets another origin in. func TestHeaders(t *testing.T) { t.Parallel() want := map[string]string{ "X-Content-Type-Options": "nosniff", "Content-Security-Policy": "default-src 'none'; frame-ancestors 'none'", "X-Frame-Options": "DENY", "Referrer-Policy": "no-referrer", "Permissions-Policy": "camera=(), microphone=(), geolocation=()", "Strict-Transport-Security": "max-age=31536000; includeSubDomains", "Cache-Control": "no-store", } for _, tc := range []struct { client *fakeClient path, auth string }{ {&fakeClient{}, chatsPath, bearer}, {&fakeClient{}, chatsPath, ""}, {&fakeClient{}, "/nothing", bearer}, {&fakeClient{contactsErr: errChat}, chatsPath, bearer}, } { rec := request(t, newAPI(credential, tc.client), http.MethodGet, tc.path, tc.auth) for key, value := range want { if got := rec.Header().Get(key); got != value { t.Errorf("%d response: %s = %q, want %q", rec.Code, key, got, value) } } if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" { t.Errorf("%d response: Access-Control-Allow-Origin = %q", rec.Code, got) } } }