Archived
修正协议处理并解耦可选遥测
This commit is contained in:
@@ -5,10 +5,6 @@ import (
|
||||
"net/http"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
"go.opentelemetry.io/otel/propagation"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
type Endpoints struct {
|
||||
@@ -19,13 +15,7 @@ type Endpoints struct {
|
||||
Ready http.Handler
|
||||
}
|
||||
|
||||
type Telemetry struct {
|
||||
TracerProvider trace.TracerProvider
|
||||
MeterProvider metric.MeterProvider
|
||||
Propagators propagation.TextMapPropagator
|
||||
}
|
||||
|
||||
func NewRouter(endpoints Endpoints, telemetry Telemetry) (http.Handler, error) {
|
||||
func NewRouter(endpoints Endpoints) (http.Handler, error) {
|
||||
if endpoints.Metadata == nil || endpoints.Token == nil || endpoints.JWKS == nil || endpoints.Health == nil || endpoints.Ready == nil {
|
||||
return nil, errors.New("all HTTP endpoints are required")
|
||||
}
|
||||
@@ -37,15 +27,5 @@ func NewRouter(endpoints Endpoints, telemetry Telemetry) (http.Handler, error) {
|
||||
router.Get("/healthz", endpoints.Health.ServeHTTP)
|
||||
router.Get("/readyz", endpoints.Ready.ServeHTTP)
|
||||
|
||||
options := make([]otelhttp.Option, 0, 3)
|
||||
if telemetry.TracerProvider != nil {
|
||||
options = append(options, otelhttp.WithTracerProvider(telemetry.TracerProvider))
|
||||
}
|
||||
if telemetry.MeterProvider != nil {
|
||||
options = append(options, otelhttp.WithMeterProvider(telemetry.MeterProvider))
|
||||
}
|
||||
if telemetry.Propagators != nil {
|
||||
options = append(options, otelhttp.WithPropagators(telemetry.Propagators))
|
||||
}
|
||||
return otelhttp.NewHandler(router, "workload-sts.http", options...), nil
|
||||
return router, nil
|
||||
}
|
||||
|
||||
@@ -3,11 +3,7 @@ package httpapi
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
"go.opentelemetry.io/otel/sdk/trace/tracetest"
|
||||
)
|
||||
|
||||
func TestRouterExposesOnlySpecifiedMethods(t *testing.T) {
|
||||
@@ -20,7 +16,7 @@ func TestRouterExposesOnlySpecifiedMethods(t *testing.T) {
|
||||
JWKS: endpoint,
|
||||
Health: endpoint,
|
||||
Ready: endpoint,
|
||||
}, Telemetry{})
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewRouter() error = %v", err)
|
||||
}
|
||||
@@ -53,43 +49,7 @@ func TestRouterExposesOnlySpecifiedMethods(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRouterRequiresEveryEndpoint(t *testing.T) {
|
||||
if _, err := NewRouter(Endpoints{}, Telemetry{}); err == nil {
|
||||
if _, err := NewRouter(Endpoints{}); err == nil {
|
||||
t.Fatal("NewRouter() error = nil, want missing endpoint error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouterTelemetryDoesNotCaptureCredentialValues(t *testing.T) {
|
||||
const canary = "canary-secret-token"
|
||||
recorder := tracetest.NewSpanRecorder()
|
||||
tracerProvider := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(recorder))
|
||||
t.Cleanup(func() { _ = tracerProvider.Shutdown(t.Context()) })
|
||||
|
||||
endpoint := http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
|
||||
response.WriteHeader(http.StatusNoContent)
|
||||
})
|
||||
router, err := NewRouter(Endpoints{
|
||||
Metadata: endpoint,
|
||||
Token: endpoint,
|
||||
JWKS: endpoint,
|
||||
Health: endpoint,
|
||||
Ready: endpoint,
|
||||
}, Telemetry{TracerProvider: tracerProvider})
|
||||
if err != nil {
|
||||
t.Fatalf("NewRouter() error = %v", err)
|
||||
}
|
||||
|
||||
request := httptest.NewRequest(http.MethodGet, "/oauth2/jwks?subject_token="+canary, nil)
|
||||
request.Header.Set("Authorization", "Bearer "+canary)
|
||||
response := httptest.NewRecorder()
|
||||
router.ServeHTTP(response, request)
|
||||
|
||||
spans := recorder.Ended()
|
||||
if len(spans) != 1 {
|
||||
t.Fatalf("ended spans = %d, want 1", len(spans))
|
||||
}
|
||||
for _, attribute := range spans[0].Attributes() {
|
||||
if strings.Contains(attribute.Value.Emit(), canary) {
|
||||
t.Fatalf("span attribute %q captured credential canary", attribute.Key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -64,9 +64,12 @@ func ParseExchangeRequest(form url.Values) (ExchangeRequest, error) {
|
||||
if len(request.Scopes) == 0 {
|
||||
return ExchangeRequest{}, requestError("invalid_scope", "scope is required")
|
||||
}
|
||||
for _, unsupported := range []string{"resource", "actor_token", "actor_token_type"} {
|
||||
if form.Has(unsupported) {
|
||||
return ExchangeRequest{}, requestError("invalid_request", unsupported+" is not supported")
|
||||
if form.Has("resource") {
|
||||
return ExchangeRequest{}, requestError("invalid_target", "resource target is not supported")
|
||||
}
|
||||
for _, actorParameter := range []string{"actor_token", "actor_token_type"} {
|
||||
if form.Has(actorParameter) {
|
||||
return ExchangeRequest{}, requestError("invalid_request", "actor token is not supported")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -36,6 +36,7 @@ func TestParseExchangeRequestRejectsInvalidInput(t *testing.T) {
|
||||
{name: "missing audience", mutate: func(v url.Values) { v.Del("audience") }, code: "invalid_target"},
|
||||
{name: "multiple audiences", mutate: func(v url.Values) { v.Add("audience", "second") }, code: "invalid_target"},
|
||||
{name: "missing scope", mutate: func(v url.Values) { v.Del("scope") }, code: "invalid_scope"},
|
||||
{name: "resource target", mutate: func(v url.Values) { v.Set("resource", "https://example.test") }, code: "invalid_target"},
|
||||
{name: "actor token", mutate: func(v url.Values) { v.Set("actor_token", "not-supported") }, code: "invalid_request"},
|
||||
}
|
||||
|
||||
@@ -56,6 +57,14 @@ func TestParseExchangeRequestRejectsInvalidInput(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseExchangeRequestIgnoresUnknownParameters(t *testing.T) {
|
||||
form := validForm()
|
||||
form.Set("future_extension", "ignored")
|
||||
if _, err := ParseExchangeRequest(form); err != nil {
|
||||
t.Fatalf("ParseExchangeRequest() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func validForm() url.Values {
|
||||
return url.Values{
|
||||
"grant_type": {TokenExchangeGrantType},
|
||||
|
||||
+37
-14
@@ -2,15 +2,14 @@ package signing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var rawURLEncoding = base64.RawURLEncoding
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
)
|
||||
|
||||
type RS256Signer interface {
|
||||
ActiveKey(ctx context.Context) (SigningKey, error)
|
||||
@@ -59,25 +58,49 @@ func (i *Issuer) Sign(ctx context.Context, claims Claims) (string, error) {
|
||||
return "", errors.New("signer returned an invalid signing key")
|
||||
}
|
||||
|
||||
header, err := json.Marshal(map[string]string{
|
||||
"alg": "RS256",
|
||||
"kid": key.ID,
|
||||
"typ": "at+jwt",
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("encode protected header: %w", err)
|
||||
}
|
||||
payload, err := json.Marshal(claims)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("encode claims: %w", err)
|
||||
}
|
||||
|
||||
signingInput := rawURLEncoding.EncodeToString(header) + "." + rawURLEncoding.EncodeToString(payload)
|
||||
signature, err := i.signer.SignRS256(ctx, key, []byte(signingInput))
|
||||
opaque := &contextSigner{ctx: ctx, signer: i.signer, key: key}
|
||||
options := (&jose.SignerOptions{}).
|
||||
WithType(jose.ContentType("at+jwt")).
|
||||
WithHeader(jose.HeaderKey("kid"), key.ID)
|
||||
signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: opaque}, options)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create JWT signer: %w", err)
|
||||
}
|
||||
jws, err := signer.Sign(payload)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("sign JWT: %w", err)
|
||||
}
|
||||
return signingInput + "." + rawURLEncoding.EncodeToString(signature), nil
|
||||
compact, err := jws.CompactSerialize()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("serialize JWT: %w", err)
|
||||
}
|
||||
return compact, nil
|
||||
}
|
||||
|
||||
type contextSigner struct {
|
||||
ctx context.Context
|
||||
signer RS256Signer
|
||||
key SigningKey
|
||||
}
|
||||
|
||||
func (*contextSigner) Public() *jose.JSONWebKey {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (*contextSigner) Algs() []jose.SignatureAlgorithm {
|
||||
return []jose.SignatureAlgorithm{jose.RS256}
|
||||
}
|
||||
|
||||
func (s *contextSigner) SignPayload(payload []byte, algorithm jose.SignatureAlgorithm) ([]byte, error) {
|
||||
if algorithm != jose.RS256 {
|
||||
return nil, errors.New("unsupported signing algorithm")
|
||||
}
|
||||
return s.signer.SignRS256(s.ctx, s.key, payload)
|
||||
}
|
||||
|
||||
func validateClaims(claims Claims, now time.Time) error {
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
package telemetry
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
|
||||
"go.opentelemetry.io/otel/propagation"
|
||||
)
|
||||
|
||||
func WrapHTTP(handler http.Handler, providers *Providers, propagators propagation.TextMapPropagator) http.Handler {
|
||||
if providers == nil {
|
||||
return handler
|
||||
}
|
||||
|
||||
options := []otelhttp.Option{
|
||||
otelhttp.WithTracerProvider(providers.Tracer),
|
||||
otelhttp.WithMeterProvider(providers.Meter),
|
||||
}
|
||||
if propagators != nil {
|
||||
options = append(options, otelhttp.WithPropagators(propagators))
|
||||
}
|
||||
return otelhttp.NewHandler(handler, "workload-sts.http", options...)
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package telemetry
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
metricnoop "go.opentelemetry.io/otel/metric/noop"
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
"go.opentelemetry.io/otel/sdk/trace/tracetest"
|
||||
)
|
||||
|
||||
func TestWrapHTTPWithoutProvidersReturnsOriginalHandler(t *testing.T) {
|
||||
handler := http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})
|
||||
if got := WrapHTTP(handler, nil, nil); got == nil {
|
||||
t.Fatal("WrapHTTP() returned nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrapHTTPDoesNotCaptureCredentialValues(t *testing.T) {
|
||||
const canary = "canary-secret-token"
|
||||
recorder := tracetest.NewSpanRecorder()
|
||||
tracerProvider := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(recorder))
|
||||
t.Cleanup(func() { _ = tracerProvider.Shutdown(t.Context()) })
|
||||
|
||||
handler := WrapHTTP(http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
|
||||
response.WriteHeader(http.StatusNoContent)
|
||||
}), &Providers{Tracer: tracerProvider, Meter: metricnoop.NewMeterProvider()}, nil)
|
||||
|
||||
request := httptest.NewRequest(http.MethodGet, "/oauth2/jwks?subject_token="+canary, nil)
|
||||
request.Header.Set("Authorization", "Bearer "+canary)
|
||||
response := httptest.NewRecorder()
|
||||
handler.ServeHTTP(response, request)
|
||||
|
||||
spans := recorder.Ended()
|
||||
if len(spans) != 1 {
|
||||
t.Fatalf("ended spans = %d, want 1", len(spans))
|
||||
}
|
||||
for _, attribute := range spans[0].Attributes() {
|
||||
if strings.Contains(attribute.Value.Emit(), canary) {
|
||||
t.Fatalf("span attribute %q captured credential canary", attribute.Key)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user