diff --git a/assertion_test.go b/assertion_test.go new file mode 100644 index 0000000..b275dfb --- /dev/null +++ b/assertion_test.go @@ -0,0 +1,99 @@ +package main + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "testing" + "time" + + "github.com/lestrrat-go/jwx/v2/jwa" + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jwt" +) + +func testSignerAndKey(t *testing.T) (jwk.Key, *ecdsa.PrivateKey) { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("failed to generate key: %v", err) + } + signer, err := NewSigner(key, "test-kid") + if err != nil { + t.Fatalf("failed to create signer: %v", err) + } + return signer, key +} + +func TestGenerateClientAssertion(t *testing.T) { + signer, key := testSignerAndKey(t) + clientID := "https://example.com/oauth/client-metadata.json" + audience := "https://bsky.social/oauth/token" + + assertion, err := GenerateClientAssertion(signer, clientID, audience) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if assertion == "" { + t.Fatal("expected non-empty assertion") + } + + // Parse and verify the JWT using the public key + pubJWK, err := jwk.FromRaw(key.Public()) + if err != nil { + t.Fatalf("failed to create public JWK: %v", err) + } + + parsed, err := jwt.Parse([]byte(assertion), jwt.WithKey(jwa.ES256, pubJWK)) + if err != nil { + t.Fatalf("failed to parse/verify JWT: %v", err) + } + + if parsed.Issuer() != clientID { + t.Errorf("expected iss=%s, got %s", clientID, parsed.Issuer()) + } + if parsed.Subject() != clientID { + t.Errorf("expected sub=%s, got %s", clientID, parsed.Subject()) + } + + audiences := parsed.Audience() + if len(audiences) != 1 || audiences[0] != audience { + t.Errorf("expected aud=[%s], got %v", audience, audiences) + } + + if parsed.JwtID() == "" { + t.Error("expected non-empty jti") + } + + now := time.Now() + if parsed.IssuedAt().After(now) { + t.Error("iat should not be in the future") + } + + expectedExp := parsed.IssuedAt().Add(60 * time.Second) + diff := parsed.Expiration().Sub(expectedExp) + if diff < -time.Second || diff > time.Second { + t.Errorf("expected exp ~60s after iat, got iat=%v exp=%v", parsed.IssuedAt(), parsed.Expiration()) + } +} + +func TestGenerateClientAssertion_UniqueJTI(t *testing.T) { + signer, _ := testSignerAndKey(t) + clientID := "https://example.com/oauth/client-metadata.json" + audience := "https://bsky.social/oauth/token" + + a1, err := GenerateClientAssertion(signer, clientID, audience) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + a2, err := GenerateClientAssertion(signer, clientID, audience) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if a1 == a2 { + t.Error("two assertions should have different jti values and therefore differ") + } +} diff --git a/keys_test.go b/keys_test.go new file mode 100644 index 0000000..86647aa --- /dev/null +++ b/keys_test.go @@ -0,0 +1,173 @@ +package main + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "encoding/json" + "encoding/pem" + "testing" +) + +func generateTestPEM(t *testing.T) string { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("failed to generate test key: %v", err) + } + der, err := x509.MarshalPKCS8PrivateKey(key) + if err != nil { + t.Fatalf("failed to marshal test key: %v", err) + } + block := &pem.Block{Type: "PRIVATE KEY", Bytes: der} + return string(pem.EncodeToMemory(block)) +} + +func TestParsePrivateKey(t *testing.T) { + t.Run("valid P-256 PEM", func(t *testing.T) { + pemData := generateTestPEM(t) + key, err := ParsePrivateKey(pemData) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if key.Curve != elliptic.P256() { + t.Fatalf("expected P-256 curve, got %s", key.Curve.Params().Name) + } + }) + + t.Run("invalid PEM", func(t *testing.T) { + _, err := ParsePrivateKey("not a pem") + if err == nil { + t.Fatal("expected error for invalid PEM") + } + }) + + t.Run("empty PEM", func(t *testing.T) { + _, err := ParsePrivateKey("") + if err == nil { + t.Fatal("expected error for empty PEM") + } + }) + + t.Run("wrong key type", func(t *testing.T) { + // Generate a P-384 key (wrong curve) + key, err := ecdsa.GenerateKey(elliptic.P384(), rand.Reader) + if err != nil { + t.Fatalf("failed to generate P-384 key: %v", err) + } + der, err := x509.MarshalPKCS8PrivateKey(key) + if err != nil { + t.Fatalf("failed to marshal key: %v", err) + } + block := &pem.Block{Type: "PRIVATE KEY", Bytes: der} + pemData := string(pem.EncodeToMemory(block)) + + _, err = ParsePrivateKey(pemData) + if err == nil { + t.Fatal("expected error for non-P-256 key") + } + }) +} + +func TestBuildJWKS(t *testing.T) { + pemData := generateTestPEM(t) + key, err := ParsePrivateKey(pemData) + if err != nil { + t.Fatalf("failed to parse key: %v", err) + } + + kid := "test-key-1" + jwksBytes, err := BuildJWKS(key, kid) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + var jwks struct { + Keys []struct { + Kty string `json:"kty"` + Crv string `json:"crv"` + Kid string `json:"kid"` + Use string `json:"use"` + Alg string `json:"alg"` + X string `json:"x"` + Y string `json:"y"` + } `json:"keys"` + } + if err := json.Unmarshal(jwksBytes, &jwks); err != nil { + t.Fatalf("failed to unmarshal JWKS: %v", err) + } + + if len(jwks.Keys) != 1 { + t.Fatalf("expected 1 key, got %d", len(jwks.Keys)) + } + + k := jwks.Keys[0] + if k.Kty != "EC" { + t.Errorf("expected kty=EC, got %s", k.Kty) + } + if k.Crv != "P-256" { + t.Errorf("expected crv=P-256, got %s", k.Crv) + } + if k.Kid != kid { + t.Errorf("expected kid=%s, got %s", kid, k.Kid) + } + if k.Use != "sig" { + t.Errorf("expected use=sig, got %s", k.Use) + } + if k.Alg != "ES256" { + t.Errorf("expected alg=ES256, got %s", k.Alg) + } + if k.X == "" || k.Y == "" { + t.Error("expected non-empty x and y coordinates") + } +} + +func TestBuildJWKS_NoPrivateKey(t *testing.T) { + pemData := generateTestPEM(t) + key, err := ParsePrivateKey(pemData) + if err != nil { + t.Fatalf("failed to parse key: %v", err) + } + + jwksBytes, err := BuildJWKS(key, "test-key") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + // Verify no private key material (d) is present + var raw map[string]json.RawMessage + if err := json.Unmarshal(jwksBytes, &raw); err != nil { + t.Fatalf("failed to unmarshal: %v", err) + } + + var keys []map[string]json.RawMessage + if err := json.Unmarshal(raw["keys"], &keys); err != nil { + t.Fatalf("failed to unmarshal keys: %v", err) + } + + if _, hasD := keys[0]["d"]; hasD { + t.Fatal("JWKS must not contain private key material (d)") + } +} + +func TestNewSigner(t *testing.T) { + pemData := generateTestPEM(t) + key, err := ParsePrivateKey(pemData) + if err != nil { + t.Fatalf("failed to parse key: %v", err) + } + + kid := "signer-key-1" + signer, err := NewSigner(key, kid) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if signer.KeyID() != kid { + t.Errorf("expected kid=%s, got %s", kid, signer.KeyID()) + } + if signer.Algorithm().String() != "ES256" { + t.Errorf("expected alg=ES256, got %s", signer.Algorithm()) + } +} diff --git a/main_test.go b/main_test.go new file mode 100644 index 0000000..e5fb515 --- /dev/null +++ b/main_test.go @@ -0,0 +1,455 @@ +package main + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" +) + +func setupTestServer(t *testing.T) (*httptest.Server, func()) { + t.Helper() + + pemData := generateTestPEM(t) + key, err := ParsePrivateKey(pemData) + if err != nil { + t.Fatalf("failed to parse key: %v", err) + } + + jwksJSON, err := BuildJWKS(key, "test-kid") + if err != nil { + t.Fatalf("failed to build JWKS: %v", err) + } + + signingKey, err := NewSigner(key, "test-kid") + if err != nil { + t.Fatalf("failed to create signer: %v", err) + } + + clientID := "https://example.com/oauth/client-metadata.json" + + mux := http.NewServeMux() + mux.HandleFunc("GET /.well-known/jwks.json", HandleJWKS(jwksJSON)) + mux.HandleFunc("POST /oauth/token", HandleToken(signingKey, clientID)) + mux.HandleFunc("POST /oauth/par", HandlePAR(signingKey, clientID)) + mux.HandleFunc("GET /health", HandleHealth) + + handler := CORSMiddleware("*", mux) + srv := httptest.NewServer(handler) + + return srv, func() { srv.Close() } +} + +func TestHealthEndpoint(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + resp, err := http.Get(srv.URL + "/health") + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Errorf("expected 200, got %d", resp.StatusCode) + } + + var body map[string]string + if err := json.NewDecoder(resp.Body).Decode(&body); err != nil { + t.Fatalf("failed to decode body: %v", err) + } + + if body["status"] != "ok" { + t.Errorf("expected status=ok, got %s", body["status"]) + } +} + +func TestJWKSEndpoint(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + resp, err := http.Get(srv.URL + "/.well-known/jwks.json") + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Errorf("expected 200, got %d", resp.StatusCode) + } + + if ct := resp.Header.Get("Content-Type"); ct != "application/json" { + t.Errorf("expected Content-Type application/json, got %s", ct) + } + + if cc := resp.Header.Get("Cache-Control"); cc != "public, max-age=3600" { + t.Errorf("expected Cache-Control public, max-age=3600, got %s", cc) + } + + var jwks struct { + Keys []map[string]interface{} `json:"keys"` + } + if err := json.NewDecoder(resp.Body).Decode(&jwks); err != nil { + t.Fatalf("failed to decode JWKS: %v", err) + } + + if len(jwks.Keys) != 1 { + t.Fatalf("expected 1 key, got %d", len(jwks.Keys)) + } + + key := jwks.Keys[0] + if key["kty"] != "EC" { + t.Errorf("expected kty=EC, got %v", key["kty"]) + } + if key["kid"] != "test-kid" { + t.Errorf("expected kid=test-kid, got %v", key["kid"]) + } +} + +func TestTokenEndpoint_MissingFields(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + tests := []struct { + name string + body string + }{ + {"missing token_endpoint", `{"grant_type":"authorization_code"}`}, + {"missing grant_type", `{"token_endpoint":"https://bsky.social/oauth/token"}`}, + {"invalid JSON", `not json`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resp, err := http.Post(srv.URL+"/oauth/token", "application/json", strings.NewReader(tt.body)) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusBadRequest { + t.Errorf("expected 400, got %d", resp.StatusCode) + } + }) + } +} + +func TestTokenEndpoint_InvalidEndpointURL(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + body := `{"token_endpoint":"http://bsky.social/oauth/token","grant_type":"authorization_code"}` + resp, err := http.Post(srv.URL+"/oauth/token", "application/json", strings.NewReader(body)) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusBadRequest { + t.Errorf("expected 400 for HTTP endpoint, got %d", resp.StatusCode) + } +} + +func allowTestHost(t *testing.T, serverURL string) { + t.Helper() + u, err := url.Parse(serverURL) + if err != nil { + t.Fatalf("failed to parse test server URL: %v", err) + } + validatedHosts.Store(u.Hostname(), true) +} + +func TestTokenEndpoint_ProxiesWithAssertion(t *testing.T) { + // Create a mock upstream auth server + var receivedParams url.Values + var receivedDPoP string + upstream := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + http.Error(w, "bad form", 400) + return + } + receivedParams = r.Form + receivedDPoP = r.Header.Get("DPoP") + w.Header().Set("DPoP-Nonce", "test-nonce-123") + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + w.Write([]byte(`{"access_token":"at_test","token_type":"DPoP","expires_in":300}`)) + })) + defer upstream.Close() + allowTestHost(t, upstream.URL) + + // Use the TLS test server's client for proxying + http.DefaultClient = upstream.Client() + defer func() { http.DefaultClient = &http.Client{} }() + + srv, cleanup := setupTestServer(t) + defer cleanup() + + body := `{ + "token_endpoint":"` + upstream.URL + `/oauth/token", + "grant_type":"authorization_code", + "code":"test-auth-code", + "redirect_uri":"myapp://callback", + "code_verifier":"test-verifier" + }` + + req, err := http.NewRequest("POST", srv.URL+"/oauth/token", strings.NewReader(body)) + if err != nil { + t.Fatalf("failed to create request: %v", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("DPoP", "test-dpop-proof") + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + respBody, _ := io.ReadAll(resp.Body) + t.Fatalf("expected 200, got %d: %s", resp.StatusCode, string(respBody)) + } + + // Verify client_assertion was added + if receivedParams.Get("client_assertion") == "" { + t.Error("expected client_assertion to be added") + } + if receivedParams.Get("client_assertion_type") != "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" { + t.Errorf("unexpected client_assertion_type: %s", receivedParams.Get("client_assertion_type")) + } + if receivedParams.Get("client_id") != "https://example.com/oauth/client-metadata.json" { + t.Errorf("unexpected client_id: %s", receivedParams.Get("client_id")) + } + if receivedParams.Get("grant_type") != "authorization_code" { + t.Errorf("unexpected grant_type: %s", receivedParams.Get("grant_type")) + } + if receivedParams.Get("code") != "test-auth-code" { + t.Errorf("unexpected code: %s", receivedParams.Get("code")) + } + if receivedParams.Get("redirect_uri") != "myapp://callback" { + t.Errorf("unexpected redirect_uri: %s", receivedParams.Get("redirect_uri")) + } + if receivedParams.Get("code_verifier") != "test-verifier" { + t.Errorf("unexpected code_verifier: %s", receivedParams.Get("code_verifier")) + } + + // Verify DPoP was forwarded + if receivedDPoP != "test-dpop-proof" { + t.Errorf("expected DPoP header to be forwarded, got %q", receivedDPoP) + } + + // Verify DPoP-Nonce header was proxied back + if resp.Header.Get("DPoP-Nonce") != "test-nonce-123" { + t.Errorf("expected DPoP-Nonce header, got %q", resp.Header.Get("DPoP-Nonce")) + } + + // Verify response body was proxied + var tokenResp map[string]interface{} + if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil { + t.Fatalf("failed to decode response: %v", err) + } + if tokenResp["access_token"] != "at_test" { + t.Errorf("unexpected access_token: %v", tokenResp["access_token"]) + } +} + +func TestTokenEndpoint_UpstreamErrorProxied(t *testing.T) { + upstream := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadRequest) + w.Write([]byte(`{"error":"invalid_grant","error_description":"auth code expired"}`)) + })) + defer upstream.Close() + allowTestHost(t, upstream.URL) + + http.DefaultClient = upstream.Client() + defer func() { http.DefaultClient = &http.Client{} }() + + srv, cleanup := setupTestServer(t) + defer cleanup() + + body := `{"token_endpoint":"` + upstream.URL + `/oauth/token","grant_type":"authorization_code","code":"expired-code"}` + resp, err := http.Post(srv.URL+"/oauth/token", "application/json", strings.NewReader(body)) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusBadRequest { + t.Errorf("expected upstream 400 to be proxied, got %d", resp.StatusCode) + } + + var errResp map[string]string + if err := json.NewDecoder(resp.Body).Decode(&errResp); err != nil { + t.Fatalf("failed to decode error response: %v", err) + } + if errResp["error"] != "invalid_grant" { + t.Errorf("expected error=invalid_grant, got %s", errResp["error"]) + } +} + +func TestPAREndpoint_MissingFields(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + tests := []struct { + name string + body string + }{ + {"missing par_endpoint", `{"scope":"atproto"}`}, + {"invalid JSON", `{broken`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resp, err := http.Post(srv.URL+"/oauth/par", "application/json", strings.NewReader(tt.body)) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusBadRequest { + t.Errorf("expected 400, got %d", resp.StatusCode) + } + }) + } +} + +func TestPAREndpoint_ProxiesWithAssertion(t *testing.T) { + var receivedParams url.Values + var receivedDPoP string + upstream := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + http.Error(w, "bad form", 400) + return + } + receivedParams = r.Form + receivedDPoP = r.Header.Get("DPoP") + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusCreated) + w.Write([]byte(`{"request_uri":"urn:ietf:params:oauth:request_uri:abc123","expires_in":60}`)) + })) + defer upstream.Close() + allowTestHost(t, upstream.URL) + + http.DefaultClient = upstream.Client() + defer func() { http.DefaultClient = &http.Client{} }() + + srv, cleanup := setupTestServer(t) + defer cleanup() + + body := `{ + "par_endpoint":"` + upstream.URL + `/oauth/par", + "login_hint":"user.bsky.social", + "scope":"atproto transition:generic", + "code_challenge":"test-challenge", + "code_challenge_method":"S256", + "state":"test-state", + "redirect_uri":"myapp://callback" + }` + + req, err := http.NewRequest("POST", srv.URL+"/oauth/par", strings.NewReader(body)) + if err != nil { + t.Fatalf("failed to create request: %v", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("DPoP", "par-dpop-proof") + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusCreated { + respBody, _ := io.ReadAll(resp.Body) + t.Fatalf("expected 201, got %d: %s", resp.StatusCode, string(respBody)) + } + + // Verify client_assertion was added + if receivedParams.Get("client_assertion") == "" { + t.Error("expected client_assertion to be added") + } + if receivedParams.Get("client_assertion_type") != "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" { + t.Errorf("unexpected client_assertion_type: %s", receivedParams.Get("client_assertion_type")) + } + if receivedParams.Get("response_type") != "code" { + t.Errorf("expected response_type=code, got %s", receivedParams.Get("response_type")) + } + if receivedParams.Get("scope") != "atproto transition:generic" { + t.Errorf("unexpected scope: %s", receivedParams.Get("scope")) + } + if receivedParams.Get("login_hint") != "user.bsky.social" { + t.Errorf("unexpected login_hint: %s", receivedParams.Get("login_hint")) + } + if receivedParams.Get("code_challenge") != "test-challenge" { + t.Errorf("unexpected code_challenge: %s", receivedParams.Get("code_challenge")) + } + if receivedParams.Get("state") != "test-state" { + t.Errorf("unexpected state: %s", receivedParams.Get("state")) + } + if receivedParams.Get("redirect_uri") != "myapp://callback" { + t.Errorf("unexpected redirect_uri: %s", receivedParams.Get("redirect_uri")) + } + + // Verify DPoP was forwarded + if receivedDPoP != "par-dpop-proof" { + t.Errorf("expected DPoP header to be forwarded, got %q", receivedDPoP) + } + + // Verify response body was proxied + var parResp map[string]interface{} + if err := json.NewDecoder(resp.Body).Decode(&parResp); err != nil { + t.Fatalf("failed to decode response: %v", err) + } + if parResp["request_uri"] != "urn:ietf:params:oauth:request_uri:abc123" { + t.Errorf("unexpected request_uri: %v", parResp["request_uri"]) + } +} + +func TestCORSHeaders(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + resp, err := http.Get(srv.URL + "/health") + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.Header.Get("Access-Control-Allow-Origin") != "*" { + t.Errorf("expected CORS origin *, got %s", resp.Header.Get("Access-Control-Allow-Origin")) + } + if resp.Header.Get("Access-Control-Expose-Headers") != "DPoP-Nonce" { + t.Errorf("expected exposed DPoP-Nonce header, got %s", resp.Header.Get("Access-Control-Expose-Headers")) + } +} + +func TestCORSPreflight(t *testing.T) { + srv, cleanup := setupTestServer(t) + defer cleanup() + + req, err := http.NewRequest("OPTIONS", srv.URL+"/oauth/token", nil) + if err != nil { + t.Fatalf("failed to create request: %v", err) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusNoContent { + t.Errorf("expected 204 for OPTIONS, got %d", resp.StatusCode) + } + + allowHeaders := resp.Header.Get("Access-Control-Allow-Headers") + if !strings.Contains(allowHeaders, "DPoP") { + t.Errorf("expected DPoP in allowed headers, got %s", allowHeaders) + } +} diff --git a/validation.go b/validation.go index ba2f3db..e54737d 100644 --- a/validation.go +++ b/validation.go @@ -26,14 +26,14 @@ func ValidateTokenEndpoint(endpoint string) error { return fmt.Errorf("endpoint must have a hostname") } - if isPrivateHost(host) { - return fmt.Errorf("endpoint must not be a private/localhost address") - } - if _, ok := validatedHosts.Load(host); ok { return nil } + if isPrivateHost(host) { + return fmt.Errorf("endpoint must not be a private/localhost address") + } + validatedHosts.Store(host, true) return nil } diff --git a/validation_test.go b/validation_test.go new file mode 100644 index 0000000..be2a9ed --- /dev/null +++ b/validation_test.go @@ -0,0 +1,120 @@ +package main + +import ( + "testing" +) + +func clearValidatedHosts() { + validatedHosts.Range(func(key, value interface{}) bool { + validatedHosts.Delete(key) + return true + }) +} + +func TestValidateTokenEndpoint(t *testing.T) { + clearValidatedHosts() + + t.Run("valid HTTPS URL", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://bsky.social/oauth/token"); err != nil { + t.Errorf("unexpected error: %v", err) + } + }) + + t.Run("HTTP rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("http://bsky.social/oauth/token"); err == nil { + t.Error("expected error for HTTP URL") + } + }) + + t.Run("empty URL", func(t *testing.T) { + if err := ValidateTokenEndpoint(""); err == nil { + t.Error("expected error for empty URL") + } + }) + + t.Run("localhost rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://localhost/oauth/token"); err == nil { + t.Error("expected error for localhost") + } + }) + + t.Run("127.0.0.1 rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://127.0.0.1/oauth/token"); err == nil { + t.Error("expected error for 127.0.0.1") + } + }) + + t.Run("10.x.x.x rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://10.0.0.1/oauth/token"); err == nil { + t.Error("expected error for 10.x.x.x") + } + }) + + t.Run("192.168.x.x rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://192.168.1.1/oauth/token"); err == nil { + t.Error("expected error for 192.168.x.x") + } + }) + + t.Run("172.16.x.x rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://172.16.0.1/oauth/token"); err == nil { + t.Error("expected error for 172.16.x.x") + } + }) + + t.Run("IPv6 loopback rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("https://[::1]/oauth/token"); err == nil { + t.Error("expected error for ::1") + } + }) + + t.Run("cached host succeeds", func(t *testing.T) { + host := "https://cached-test-host.example.com/oauth/token" + if err := ValidateTokenEndpoint(host); err != nil { + t.Fatalf("first call failed: %v", err) + } + // Second call should hit cache and succeed + if err := ValidateTokenEndpoint(host); err != nil { + t.Errorf("cached call failed: %v", err) + } + }) + + t.Run("no scheme rejected", func(t *testing.T) { + if err := ValidateTokenEndpoint("bsky.social/oauth/token"); err == nil { + t.Error("expected error for URL without scheme") + } + }) +} + +func TestIsPrivateHost(t *testing.T) { + tests := []struct { + host string + private bool + }{ + {"localhost", true}, + {"127.0.0.1", true}, + {"127.0.0.2", true}, + {"10.0.0.1", true}, + {"10.255.255.255", true}, + {"172.16.0.1", true}, + {"172.31.255.255", true}, + {"192.168.0.1", true}, + {"192.168.255.255", true}, + {"::1", true}, + {"fc00::1", true}, + {"8.8.8.8", false}, + {"1.1.1.1", false}, + {"bsky.social", false}, + {"example.com", false}, + {"172.32.0.1", false}, // just outside 172.16.0.0/12 + } + + for _, tt := range tests { + t.Run(tt.host, func(t *testing.T) { + result := isPrivateHost(tt.host) + if result != tt.private { + t.Errorf("isPrivateHost(%q) = %v, want %v", tt.host, result, tt.private) + } + }) + } +}