实现首轮协议与 Transit 签名 PoC #1

Merged
panxiao81 merged 3 commits from feat/initial-poc into main 2026-09-11 16:51:40 +00:00
13 changed files with 1023 additions and 0 deletions
Showing only changes of commit ab5dbebca7 - Show all commits
+14
View File
@@ -0,0 +1,14 @@
.PHONY: fmt test test-transit
fmt:
gofmt -w $$(find . -name '*.go' -type f)
test:
go test ./...
test-transit:
@test -n "$(WORKLOAD_STS_TEST_BAO_ADDR)" || \
(echo 'WORKLOAD_STS_TEST_BAO_ADDR is required' >&2; exit 1)
@test -n "$(WORKLOAD_STS_TEST_BAO_TOKEN)" || \
(echo 'WORKLOAD_STS_TEST_BAO_TOKEN is required' >&2; exit 1)
go test -v ./internal/signing -run TestTransitSignerProducesJWTVerifiedByPublishedJWKS
+40
View File
@@ -0,0 +1,40 @@
module git.ddupan.top/panxiao81/workload-sts
go 1.26.0
require (
github.com/go-chi/chi/v5 v5.3.2
github.com/go-jose/go-jose/v4 v4.1.5
github.com/openbao/openbao/api/v2 v2.7.0
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.71.0
go.opentelemetry.io/otel v1.46.0
go.opentelemetry.io/otel/metric v1.46.0
go.opentelemetry.io/otel/sdk v1.46.0
go.opentelemetry.io/otel/sdk/metric v1.46.0
go.opentelemetry.io/otel/trace v1.46.0
)
require (
github.com/cenkalti/backoff/v5 v5.0.3 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/felixge/httpsnoop v1.1.0 // indirect
github.com/go-logr/logr v1.4.4 // indirect
github.com/go-logr/stdr v1.2.2 // indirect
github.com/go-viper/mapstructure/v2 v2.5.0 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/hashicorp/errwrap v1.1.0 // indirect
github.com/hashicorp/go-cleanhttp v0.5.2 // indirect
github.com/hashicorp/go-multierror v1.1.1 // indirect
github.com/hashicorp/go-retryablehttp v0.7.8 // indirect
github.com/hashicorp/go-secure-stdlib/parseutil v0.2.0 // indirect
github.com/hashicorp/go-secure-stdlib/strutil v0.1.2 // indirect
github.com/hashicorp/go-sockaddr v1.0.7 // indirect
github.com/hashicorp/hcl v1.0.1-vault-7 // indirect
github.com/mitchellh/mapstructure v1.5.0 // indirect
github.com/ryanuber/go-glob v1.0.0 // indirect
go.opentelemetry.io/auto/sdk v1.2.1 // indirect
golang.org/x/net v0.58.0 // indirect
golang.org/x/sys v0.47.0 // indirect
golang.org/x/text v0.41.0 // indirect
golang.org/x/time v0.15.0 // indirect
)
+87
View File
@@ -0,0 +1,87 @@
github.com/cenkalti/backoff/v5 v5.0.3 h1:ZN+IMa753KfX5hd8vVaMixjnqRZ3y8CuJKRKj1xcsSM=
github.com/cenkalti/backoff/v5 v5.0.3/go.mod h1:rkhZdG3JZukswDf7f0cwqPNk4K0sa+F97BxZthm/crw=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/fatih/color v1.19.0 h1:Zp3PiM21/9Ld6FzSKyL5c/BULoe/ONr9KlbYVOfG8+w=
github.com/fatih/color v1.19.0/go.mod h1:zNk67I0ZUT1bEGsSGyCZYZNrHuTkJJB+r6Q9VuMi0LE=
github.com/felixge/httpsnoop v1.1.0 h1:3YtUj32ZZkqZtt3sZZsClsymw/QDuVfpNhoA31zeORc=
github.com/felixge/httpsnoop v1.1.0/go.mod h1:Zqxgdd+1Rkcz8euOqdr7lqgCRJztwr5hp9vDSi5UZCE=
github.com/go-chi/chi/v5 v5.3.2 h1:5YQkICvTCSZ25hoRsyJazN0scjzKGiu4VAUc7H1o1nY=
github.com/go-chi/chi/v5 v5.3.2/go.mod h1:R+tYY2hNuVUUjxoPtqUdgBqevM9s9njzkTLutVsOCto=
github.com/go-jose/go-jose/v4 v4.1.5 h1:RjgjO2LOtWOJKUC5wpwY9LR3B3vwVAz6JS2YHfYU6eA=
github.com/go-jose/go-jose/v4 v4.1.5/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08=
github.com/go-logr/logr v1.2.2/go.mod h1:jdQByPbusPIv2/zmleS9BjJVeZ6kBagPoEUsqbVz/1A=
github.com/go-logr/logr v1.4.4 h1:tG4xh9yMsRCAiodLVTxyrkzSZ9+o0L1Kg/+cPVcbP/8=
github.com/go-logr/logr v1.4.4/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/go-test/deep v1.1.1 h1:0r/53hagsehfO4bzD2Pgr/+RgHqhmf+k1Bpse2cTu1U=
github.com/go-test/deep v1.1.1/go.mod h1:5C2ZWiW0ErCdrYzpqxLbTX7MG14M9iiw8DgHncVwcsE=
github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro=
github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4=
github.com/hashicorp/errwrap v1.1.0 h1:OxrOeh75EUXMY8TBjag2fzXGZ40LB6IKw45YeGUDY2I=
github.com/hashicorp/errwrap v1.1.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4=
github.com/hashicorp/go-cleanhttp v0.5.2 h1:035FKYIWjmULyFRBKPs8TBQoi0x6d9G4xc9neXJWAZQ=
github.com/hashicorp/go-cleanhttp v0.5.2/go.mod h1:kO/YDlP8L1346E6Sodw+PrpBSV4/SoxCXGY6BqNFT48=
github.com/hashicorp/go-hclog v1.6.3 h1:Qr2kF+eVWjTiYmU7Y31tYlP1h0q/X3Nl3tPGdaB11/k=
github.com/hashicorp/go-hclog v1.6.3/go.mod h1:W4Qnvbt70Wk/zYJryRzDRU/4r0kIg0PVHBcfoyhpF5M=
github.com/hashicorp/go-multierror v1.1.1 h1:H5DkEtf6CXdFp0N0Em5UCwQpXMWke8IA0+lD48awMYo=
github.com/hashicorp/go-multierror v1.1.1/go.mod h1:iw975J/qwKPdAO1clOe2L8331t/9/fmwbPZ6JB6eMoM=
github.com/hashicorp/go-retryablehttp v0.7.8 h1:ylXZWnqa7Lhqpk0L1P1LzDtGcCR0rPVUrx/c8Unxc48=
github.com/hashicorp/go-retryablehttp v0.7.8/go.mod h1:rjiScheydd+CxvumBsIrFKlx3iS0jrZ7LvzFGFmuKbw=
github.com/hashicorp/go-secure-stdlib/parseutil v0.2.0 h1:U+kC2dOhMFQctRfhK0gRctKAPTloZdMU5ZJxaesJ/VM=
github.com/hashicorp/go-secure-stdlib/parseutil v0.2.0/go.mod h1:Ll013mhdmsVDuoIXVfBtvgGJsXDYkTw1kooNcoCXuE0=
github.com/hashicorp/go-secure-stdlib/strutil v0.1.2 h1:kes8mmyCpxJsI7FTwtzRqEy9CdjCtrXrXGuOpxEA7Ts=
github.com/hashicorp/go-secure-stdlib/strutil v0.1.2/go.mod h1:Gou2R9+il93BqX25LAKCLuM+y9U2T4hlwvT1yprcna4=
github.com/hashicorp/go-sockaddr v1.0.7 h1:G+pTkSO01HpR5qCxg7lxfsFEZaG+C0VssTy/9dbT+Fw=
github.com/hashicorp/go-sockaddr v1.0.7/go.mod h1:FZQbEYa1pxkQ7WLpyXJ6cbjpT8q0YgQaK/JakXqGyWw=
github.com/hashicorp/hcl v1.0.1-vault-7 h1:ag5OxFVy3QYTFTJODRzTKVZ6xvdfLLCA1cy/Y6xGI0I=
github.com/hashicorp/hcl v1.0.1-vault-7/go.mod h1:XYhtn6ijBSAj6n4YqAaf7RBPS4I06AItNorpy+MoQNM=
github.com/mattn/go-colorable v0.1.15 h1:+u9SLTRGnXv73cEsnsmoZBom+dMU88B2M0aDcWy0/jY=
github.com/mattn/go-colorable v0.1.15/go.mod h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stgPZH1UqBm1s8=
github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI=
github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A=
github.com/mitchellh/mapstructure v1.5.0 h1:jeMsZIYE/09sWLaz43PL7Gy6RuMjD2eJVyuac5Z2hdY=
github.com/mitchellh/mapstructure v1.5.0/go.mod h1:bFUtVrKA4DC2yAKiSyO/QUcy7e+RRV2QTWOzhPopBRo=
github.com/openbao/openbao/api/v2 v2.7.0 h1:3CD1l3tr39nQraCgFGAWA5vYvPFzZoZrt3NL7DMQKAc=
github.com/openbao/openbao/api/v2 v2.7.0/go.mod h1:uXbMoyH2pjSvNyTepinUvLde8pOJB82EuhUCfOKnKbo=
github.com/ryanuber/go-glob v1.0.0 h1:iQh3xXAumdQ+4Ufa5b25cRpC5TYKlno6hsv6Cb3pkBk=
github.com/ryanuber/go-glob v1.0.0/go.mod h1:807d1WSdnB0XRJzKNil9Om6lcp/3a0v4qIHxIXzX/Yc=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.71.0 h1:3g7B90UzBltIDKq1/5mrTGxTnOFDV0ICOhLoxiZ8jlg=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.71.0/go.mod h1:Ef8SuTh59BT7+ofpDxN9z+yOlc4t2GjLmKDgYNJL/NU=
go.opentelemetry.io/otel v1.46.0 h1:FHt5/CDyVxi/8IM1CH7VE/rRgq3kLHa2mSTVMO8AWyc=
go.opentelemetry.io/otel v1.46.0/go.mod h1:Gj3SEScelsNC45tp4nSxRYlS+f5iez7W8XPMCt905kE=
go.opentelemetry.io/otel/metric v1.46.0 h1:yBnkXvgV7AXFILZc5K6IZe/CBFF3OS7BJ8ov6/lj0K8=
go.opentelemetry.io/otel/metric v1.46.0/go.mod h1:iPmdWqifKUdzziPkvvzIJXITl56fQx2mGM/DHLB3/2o=
go.opentelemetry.io/otel/metric/x v0.68.0 h1:TA/cBT23D3MnxYPwHL7YFOdYGdx0A0v+s7Mzotpd1dU=
go.opentelemetry.io/otel/metric/x v0.68.0/go.mod h1:agudOmvWhwUTjgibWDzxD2PoWYnpw5Ht5jISYOD2Hd4=
go.opentelemetry.io/otel/sdk v1.46.0 h1:h5CNQQjEbuQXY/JfZtgt3i7HVFV3aHPO2OAwO2eTYPI=
go.opentelemetry.io/otel/sdk v1.46.0/go.mod h1:GAERFXFt5SYCEB+YiKUbMBeza6UaDH7GmGOZEfh2gSM=
go.opentelemetry.io/otel/sdk/metric v1.46.0 h1:0piZ26EG4RBfebb2jhDH6ERCYHoVWduc3kLgPCwSnSE=
go.opentelemetry.io/otel/sdk/metric v1.46.0/go.mod h1:I1PbKrdVc8Qu8HYVDNtqVIwLwjNrhsV/uFuxfwg8mO4=
go.opentelemetry.io/otel/trace v1.46.0 h1:OULy7ccdJnZtJ0UDYFOIGaCmiWzJ8Vi2G/Rsu60qs1c=
go.opentelemetry.io/otel/trace v1.46.0/go.mod h1:J7GAXweO77XSFkB/rmAqk9D6ihszhFjLU+d9WuUxDLI=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
+51
View File
@@ -0,0 +1,51 @@
package httpapi
import (
"errors"
"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 {
Metadata http.Handler
Token http.Handler
JWKS http.Handler
Health http.Handler
Ready http.Handler
}
type Telemetry struct {
TracerProvider trace.TracerProvider
MeterProvider metric.MeterProvider
Propagators propagation.TextMapPropagator
}
func NewRouter(endpoints Endpoints, telemetry Telemetry) (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")
}
router := chi.NewRouter()
router.Get("/.well-known/oauth-authorization-server", endpoints.Metadata.ServeHTTP)
router.Post("/oauth2/token", endpoints.Token.ServeHTTP)
router.Get("/oauth2/jwks", endpoints.JWKS.ServeHTTP)
router.Get("/healthz", endpoints.Health.ServeHTTP)
router.Get("/readyz", endpoints.Ready.ServeHTTP)
options := make([]otelhttp.Option, 0, 3)
panxiao81 marked this conversation as resolved
Review

otel应该作为可选引入,不强制安装

otel应该作为可选引入,不强制安装
Review

已处理:纯 chi router 已移除 otelhttp 依赖路径;OTel 现在是启动层显式选择的独立 wrapper,未启用时不安装 instrumentation。

已处理:纯 chi router 已移除 otelhttp 依赖路径;OTel 现在是启动层显式选择的独立 wrapper,未启用时不安装 instrumentation。
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
}
+95
View File
@@ -0,0 +1,95 @@
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) {
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{})
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
}
tests := []struct {
method string
path string
status int
}{
{http.MethodGet, "/.well-known/oauth-authorization-server", http.StatusNoContent},
{http.MethodPost, "/oauth2/token", http.StatusNoContent},
{http.MethodGet, "/oauth2/jwks", http.StatusNoContent},
{http.MethodGet, "/healthz", http.StatusNoContent},
{http.MethodGet, "/readyz", http.StatusNoContent},
{http.MethodGet, "/oauth2/token", http.StatusMethodNotAllowed},
{http.MethodPost, "/oauth2/jwks", http.StatusMethodNotAllowed},
{http.MethodGet, "/unknown", http.StatusNotFound},
}
for _, test := range tests {
t.Run(test.method+" "+test.path, func(t *testing.T) {
request := httptest.NewRequest(test.method, test.path, nil)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != test.status {
t.Fatalf("status = %d, want %d", response.Code, test.status)
}
})
}
}
func TestRouterRequiresEveryEndpoint(t *testing.T) {
if _, err := NewRouter(Endpoints{}, Telemetry{}); 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)
}
}
}
+78
View File
@@ -0,0 +1,78 @@
package protocol
import (
"fmt"
"net/url"
"regexp"
"strings"
)
const (
TokenExchangeGrantType = "urn:ietf:params:oauth:grant-type:token-exchange"
JWTTokenType = "urn:ietf:params:oauth:token-type:jwt"
AccessTokenType = "urn:ietf:params:oauth:token-type:access_token"
)
var clientIDPattern = regexp.MustCompile(`^[a-z0-9](?:[a-z0-9-]{1,61}[a-z0-9])$`)
type ExchangeRequest struct {
ClientID string
SubjectToken string
SubjectTokenType string
RequestedTokenType string
Audience string
Scopes []string
}
type RequestError struct {
Code string
Description string
}
func (e *RequestError) Error() string {
return fmt.Sprintf("%s: %s", e.Code, e.Description)
}
func ParseExchangeRequest(form url.Values) (ExchangeRequest, error) {
request := ExchangeRequest{
ClientID: form.Get("client_id"),
SubjectToken: form.Get("subject_token"),
SubjectTokenType: form.Get("subject_token_type"),
RequestedTokenType: form.Get("requested_token_type"),
Audience: form.Get("audience"),
Scopes: strings.Fields(form.Get("scope")),
}
if form.Get("grant_type") != TokenExchangeGrantType {
return ExchangeRequest{}, requestError("unsupported_grant_type", "unsupported grant_type")
}
if !clientIDPattern.MatchString(request.ClientID) {
return ExchangeRequest{}, requestError("invalid_request", "client_id must use 3-63 lowercase letters, digits, or interior hyphens")
}
if request.SubjectToken == "" {
return ExchangeRequest{}, requestError("invalid_request", "subject_token is required")
}
if request.SubjectTokenType != JWTTokenType {
return ExchangeRequest{}, requestError("invalid_request", "unsupported subject_token_type")
}
if request.RequestedTokenType != AccessTokenType {
return ExchangeRequest{}, requestError("invalid_request", "unsupported requested_token_type")
}
if request.Audience == "" || len(form["audience"]) != 1 {
return ExchangeRequest{}, requestError("invalid_target", "exactly one audience is required")
}
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) {
panxiao81 marked this conversation as resolved Outdated
Outdated
Review

除非标准中明确规定需要返回错误,否则直接忽略就好了

除非标准中明确规定需要返回错误,否则直接忽略就好了
Outdated
Review

已处理:未知扩展参数按 RFC 6749 忽略;resource 和 actor_token 是 RFC 8693 已定义参数,resource 改为 invalid_target,actor_token 因 RFC 要求验证而在 v1 返回 invalid_request,并补充测试。

已处理:未知扩展参数按 RFC 6749 忽略;resource 和 actor_token 是 RFC 8693 已定义参数,resource 改为 invalid_target,actor_token 因 RFC 要求验证而在 v1 返回 invalid_request,并补充测试。
return ExchangeRequest{}, requestError("invalid_request", unsupported+" is not supported")
}
}
return request, nil
}
func requestError(code, description string) *RequestError {
return &RequestError{Code: code, Description: description}
}
+69
View File
@@ -0,0 +1,69 @@
package protocol
import (
"errors"
"net/url"
"reflect"
"testing"
)
func TestParseExchangeRequest(t *testing.T) {
request, err := ParseExchangeRequest(validForm())
if err != nil {
t.Fatalf("ParseExchangeRequest() error = %v", err)
}
if request.ClientID != "homelab-infra-ci" {
t.Fatalf("ClientID = %q", request.ClientID)
}
if !reflect.DeepEqual(request.Scopes, []string{"bao.login", "ssh.certificate"}) {
t.Fatalf("Scopes = %#v", request.Scopes)
}
}
func TestParseExchangeRequestRejectsInvalidInput(t *testing.T) {
tests := []struct {
name string
mutate func(url.Values)
code string
}{
{name: "missing client id", mutate: func(v url.Values) { v.Del("client_id") }, code: "invalid_request"},
{name: "short client id", mutate: func(v url.Values) { v.Set("client_id", "ci") }, code: "invalid_request"},
{name: "uppercase client id", mutate: func(v url.Values) { v.Set("client_id", "CI-worker") }, code: "invalid_request"},
{name: "leading hyphen", mutate: func(v url.Values) { v.Set("client_id", "-ci-worker") }, code: "invalid_request"},
{name: "missing token", mutate: func(v url.Values) { v.Del("subject_token") }, code: "invalid_request"},
{name: "wrong grant", mutate: func(v url.Values) { v.Set("grant_type", "client_credentials") }, code: "unsupported_grant_type"},
{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: "actor token", mutate: func(v url.Values) { v.Set("actor_token", "not-supported") }, code: "invalid_request"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
form := validForm()
test.mutate(form)
_, err := ParseExchangeRequest(form)
var requestErr *RequestError
if !errors.As(err, &requestErr) {
t.Fatalf("error = %v, want RequestError", err)
}
if requestErr.Code != test.code {
t.Fatalf("error code = %q, want %q", requestErr.Code, test.code)
}
})
}
}
func validForm() url.Values {
return url.Values{
"grant_type": {TokenExchangeGrantType},
"client_id": {"homelab-infra-ci"},
"subject_token": {"test-subject-token"},
"subject_token_type": {JWTTokenType},
"requested_token_type": {AccessTokenType},
"audience": {"https://bao.ad.ddupan.top:8200"},
"scope": {"bao.login ssh.certificate"},
}
}
+97
View File
@@ -0,0 +1,97 @@
package signing
panxiao81 marked this conversation as resolved
Review

等会你整个文件在用标准库重新实现JWT库?

等会你整个文件在用标准库重新实现JWT库?
Review

已处理:JWT protected header、JWS signing input 和 compact serialization 已改由 go-jose OpaqueSigner 完成,不再自行实现 JWT/JWS 编码。

已处理:JWT protected header、JWS signing input 和 compact serialization 已改由 go-jose OpaqueSigner 完成,不再自行实现 JWT/JWS 编码。
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"strings"
"time"
)
var rawURLEncoding = base64.RawURLEncoding
type RS256Signer interface {
ActiveKey(ctx context.Context) (SigningKey, error)
SignRS256(ctx context.Context, key SigningKey, signingInput []byte) ([]byte, error)
}
type SigningKey struct {
ID string
Version int
}
type Claims struct {
Issuer string `json:"iss"`
Subject string `json:"sub"`
PrincipalName string `json:"principal_name,omitempty"`
Audience []string `json:"aud"`
IssuedAt int64 `json:"iat"`
NotBefore int64 `json:"nbf"`
ExpiresAt int64 `json:"exp"`
JWTID string `json:"jti"`
Scope string `json:"scope"`
}
type Issuer struct {
signer RS256Signer
panxiao81 marked this conversation as resolved
Review

何意味?这个还有替换的可能吗

何意味?这个还有替换的可能吗
Review

已处理:移除了用于测试替换时间的 Issuer.now 字段。Issuer 现在固定持有 issuer/TTL,公开 Issue 使用当前时间,测试通过内部 issueAt 固定时间。

已处理:移除了用于测试替换时间的 Issuer.now 字段。Issuer 现在固定持有 issuer/TTL,公开 Issue 使用当前时间,测试通过内部 issueAt 固定时间。
now func() time.Time
}
func NewIssuer(signer RS256Signer) (*Issuer, error) {
if signer == nil {
return nil, errors.New("signer is required")
}
return &Issuer{signer: signer, now: time.Now}, nil
}
panxiao81 marked this conversation as resolved
Review

流程看上去对,但代码太面条了,说不定这里建模是有问题的

流程看上去对,但代码太面条了,说不定这里建模是有问题的
Review

已处理:流程重构为本地 encode → Transit.Sign(ctx) → 本地 compact;调用方改为 IssueRequest,issuer 统一生成时效 claims,职责和控制流已拆开。

已处理:流程重构为本地 encode → Transit.Sign(ctx) → 本地 compact;调用方改为 IssueRequest,issuer 统一生成时效 claims,职责和控制流已拆开。
func (i *Issuer) Sign(ctx context.Context, claims Claims) (string, error) {
if err := validateClaims(claims, i.now()); err != nil {
return "", err
}
key, err := i.signer.ActiveKey(ctx)
if err != nil {
return "", fmt.Errorf("select signing key: %w", err)
}
if key.ID == "" || key.Version < 1 {
return "", errors.New("signer returned an invalid signing key")
}
header, err := json.Marshal(map[string]string{
"alg": "RS256",
"kid": key.ID,
"typ": "at+jwt",
})
panxiao81 marked this conversation as resolved
Review

这部分是?

这部分是?
Review

已处理:删除了为 go-jose OpaqueSigner 捕获 context 的 contextSigner。JOSE 编码是纯本地计算,context 现在只传给 Transit RPC。

已处理:删除了为 go-jose OpaqueSigner 捕获 context 的 contextSigner。JOSE 编码是纯本地计算,context 现在只传给 Transit RPC。
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))
if err != nil {
return "", fmt.Errorf("sign JWT: %w", err)
}
return signingInput + "." + rawURLEncoding.EncodeToString(signature), nil
}
func validateClaims(claims Claims, now time.Time) error {
if claims.Issuer == "" || claims.Subject == "" || len(claims.Audience) != 1 || claims.JWTID == "" || claims.Scope == "" {
return errors.New("required JWT claim is missing")
}
if strings.ContainsAny(claims.Subject, "\r\n") || strings.ContainsAny(claims.JWTID, "\r\n") {
return errors.New("invalid JWT claim")
}
if claims.ExpiresAt <= claims.NotBefore || claims.NotBefore < claims.IssuedAt {
return errors.New("invalid JWT time window")
}
if time.Unix(claims.ExpiresAt, 0).After(now.Add(5 * time.Minute)) {
return errors.New("JWT lifetime exceeds five minutes")
}
return nil
}
+124
View File
@@ -0,0 +1,124 @@
package signing
import (
"context"
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"strings"
"testing"
"time"
"github.com/go-jose/go-jose/v4"
)
func TestIssuerProducesInteroperableRS256JWT(t *testing.T) {
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatalf("generate key: %v", err)
}
signer := &localRS256Signer{keyID: "workload-sts-v1", key: privateKey}
issuer, err := NewIssuer(signer)
if err != nil {
t.Fatalf("NewIssuer() error = %v", err)
}
now := time.Unix(1_789_062_000, 0)
issuer.now = func() time.Time { return now }
token, err := issuer.Sign(context.Background(), Claims{
Issuer: "https://identity.ad.ddupan.top",
Subject: "01993f4d-5e1a-7000-8000-000000000001",
PrincipalName: "ci/homelab-infra-plan",
Audience: []string{"https://bao.ad.ddupan.top:8200"},
IssuedAt: now.Unix(),
NotBefore: now.Unix(),
ExpiresAt: now.Add(5 * time.Minute).Unix(),
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
Scope: "bao.login",
})
if err != nil {
t.Fatalf("Sign() error = %v", err)
}
jws, err := jose.ParseSigned(token, []jose.SignatureAlgorithm{jose.RS256})
if err != nil {
t.Fatalf("go-jose rejected compact JWT: %v", err)
}
payload, err := jws.Verify(&privateKey.PublicKey)
if err != nil {
t.Fatalf("go-jose rejected signature: %v", err)
}
var got Claims
if err := json.Unmarshal(payload, &got); err != nil {
t.Fatalf("decode claims: %v", err)
}
if got.Subject != "01993f4d-5e1a-7000-8000-000000000001" {
t.Fatalf("subject = %q", got.Subject)
}
parts := strings.Split(token, ".")
headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
t.Fatalf("decode header: %v", err)
}
var header map[string]string
if err := json.Unmarshal(headerJSON, &header); err != nil {
t.Fatalf("unmarshal header: %v", err)
}
if header["typ"] != "at+jwt" || header["kid"] != "workload-sts-v1" {
t.Fatalf("header = %#v", header)
}
}
func TestIssuerRejectsExcessiveLifetime(t *testing.T) {
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatalf("generate key: %v", err)
}
issuer, err := NewIssuer(&localRS256Signer{keyID: "test", key: privateKey})
if err != nil {
t.Fatalf("NewIssuer() error = %v", err)
}
now := time.Unix(1_789_062_000, 0)
issuer.now = func() time.Time { return now }
_, err = issuer.Sign(context.Background(), Claims{
Issuer: "https://identity.ad.ddupan.top",
Subject: "principal",
Audience: []string{"audience"},
IssuedAt: now.Unix(),
NotBefore: now.Unix(),
ExpiresAt: now.Add(5*time.Minute + time.Second).Unix(),
JWTID: "jti",
Scope: "scope",
})
if err == nil {
t.Fatal("Sign() error = nil, want excessive lifetime error")
}
}
func TestIntegerRejectsNonIntegralJSONNumber(t *testing.T) {
for _, value := range []any{1.5, json.Number("1.5")} {
if _, err := integer(value); err == nil {
t.Fatalf("integer(%v) error = nil, want invalid integer error", value)
}
}
}
type localRS256Signer struct {
keyID string
key *rsa.PrivateKey
}
func (s *localRS256Signer) ActiveKey(context.Context) (SigningKey, error) {
return SigningKey{ID: s.keyID, Version: 1}, nil
}
func (s *localRS256Signer) SignRS256(_ context.Context, _ SigningKey, input []byte) ([]byte, error) {
digest := sha256.Sum256(input)
return rsa.SignPKCS1v15(rand.Reader, s.key, crypto.SHA256, digest[:])
}
+185
View File
@@ -0,0 +1,185 @@
package signing
panxiao81 marked this conversation as resolved
Review

整个文件大概看了一眼问题跟jwt.go差不多

整个文件大概看了一眼问题跟jwt.go差不多
Review

已处理:Transit metadata 改为类型化 mapstructure 解码,key metadata 读取、签名响应解析和 JWKS 转换已拆分;并通过真实 OpenBao v2.6.1 集成测试。

已处理:Transit metadata 改为类型化 mapstructure 解码,key metadata 读取、签名响应解析和 JWKS 转换已拆分;并通过真实 OpenBao v2.6.1 集成测试。
import (
"context"
"crypto/rsa"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"math"
"sort"
"strconv"
"strings"
"github.com/go-jose/go-jose/v4"
baoapi "github.com/openbao/openbao/api/v2"
)
type TransitSigner struct {
client *baoapi.Client
mount string
key string
}
func NewTransitSigner(client *baoapi.Client, mount, key string) (*TransitSigner, error) {
if client == nil {
return nil, errors.New("OpenBao client is required")
}
mount = strings.Trim(mount, "/")
key = strings.TrimSpace(key)
if mount == "" || key == "" || strings.Contains(key, "/") {
return nil, errors.New("valid Transit mount and key are required")
}
return &TransitSigner{client: client, mount: mount, key: key}, nil
}
func (s *TransitSigner) ActiveKey(ctx context.Context) (SigningKey, error) {
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
if err != nil {
return SigningKey{}, fmt.Errorf("read Transit key: %w", err)
}
if secret == nil || secret.Data == nil {
return SigningKey{}, errors.New("read Transit key: empty response")
}
version, err := integer(secret.Data["latest_version"])
if err != nil || version < 1 {
return SigningKey{}, errors.New("read Transit key: invalid latest_version")
}
return SigningKey{ID: s.keyID(version), Version: version}, nil
}
func (s *TransitSigner) SignRS256(ctx context.Context, key SigningKey, input []byte) ([]byte, error) {
if key.ID != s.keyID(key.Version) || key.Version < 1 {
return nil, errors.New("sign with Transit: invalid signing key")
}
secret, err := s.client.Logical().WriteWithContext(ctx, s.signPath(), map[string]any{
"input": base64.StdEncoding.EncodeToString(input),
"key_version": key.Version,
"signature_algorithm": "pkcs1v15",
})
if err != nil {
return nil, fmt.Errorf("sign with Transit: %w", err)
}
if secret == nil || secret.Data == nil {
return nil, errors.New("sign with Transit: empty response")
}
encoded, ok := secret.Data["signature"].(string)
if !ok {
return nil, errors.New("sign with Transit: missing signature")
}
parts := strings.SplitN(encoded, ":", 3)
if len(parts) != 3 || parts[0] != "vault" || parts[1] != "v"+strconv.Itoa(key.Version) {
return nil, errors.New("sign with Transit: unexpected signature version")
}
signature, err := base64.StdEncoding.DecodeString(parts[2])
if err != nil {
return nil, errors.New("sign with Transit: invalid signature encoding")
}
return signature, nil
}
func (s *TransitSigner) JWKS(ctx context.Context) (jose.JSONWebKeySet, error) {
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
if err != nil {
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys: %w", err)
}
if secret == nil || secret.Data == nil {
return jose.JSONWebKeySet{}, errors.New("read Transit keys: empty response")
}
keys, ok := secret.Data["keys"].(map[string]any)
if !ok {
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid keys")
}
versions := make([]int, 0, len(keys))
for versionText := range keys {
version, err := strconv.Atoi(versionText)
if err != nil || version < 1 {
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid key version")
}
versions = append(versions, version)
}
sort.Ints(versions)
set := jose.JSONWebKeySet{Keys: make([]jose.JSONWebKey, 0, len(keys))}
for _, version := range versions {
value := keys[strconv.Itoa(version)]
metadata, ok := value.(map[string]any)
if !ok {
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid key metadata")
}
publicPEM, ok := metadata["public_key"].(string)
if !ok {
return jose.JSONWebKeySet{}, errors.New("read Transit keys: missing public key")
}
publicKey, err := parseRSAPublicKey(publicPEM)
if err != nil {
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys version %d: %w", version, err)
}
set.Keys = append(set.Keys, jose.JSONWebKey{
Key: publicKey,
KeyID: s.keyID(version),
Algorithm: string(jose.RS256),
Use: "sig",
})
}
return set, nil
}
func (s *TransitSigner) keyPath() string {
return s.mount + "/keys/" + s.key
}
func (s *TransitSigner) signPath() string {
return s.mount + "/sign/" + s.key + "/sha2-256"
}
func (s *TransitSigner) keyID(version int) string {
return s.key + "-v" + strconv.Itoa(version)
}
func integer(value any) (int, error) {
switch value := value.(type) {
case int:
return value, nil
case int64:
if value > int64(^uint(0)>>1) || value < -int64(^uint(0)>>1)-1 {
return 0, errors.New("integer out of range")
}
return int(value), nil
case float64:
if math.IsNaN(value) || math.IsInf(value, 0) || math.Trunc(value) != value || value > float64(^uint(0)>>1) || value < -float64(^uint(0)>>1)-1 {
return 0, errors.New("invalid integer value")
}
return int(value), nil
case json.Number:
parsed, err := value.Int64()
if err != nil {
return 0, fmt.Errorf("invalid integer: %w", err)
}
return integer(parsed)
default:
return 0, fmt.Errorf("not an integer: %T", value)
}
}
func parseRSAPublicKey(value string) (*rsa.PublicKey, error) {
block, _ := pem.Decode([]byte(value))
if block == nil {
return nil, errors.New("invalid PEM public key")
}
parsed, err := x509.ParsePKIXPublicKey(block.Bytes)
if err != nil {
return nil, fmt.Errorf("parse public key: %w", err)
}
publicKey, ok := parsed.(*rsa.PublicKey)
if !ok {
return nil, errors.New("Transit key is not RSA")
}
return publicKey, nil
}
@@ -0,0 +1,80 @@
package signing
import (
"context"
"encoding/json"
"os"
"testing"
"time"
"github.com/go-jose/go-jose/v4"
baoapi "github.com/openbao/openbao/api/v2"
)
func TestTransitSignerProducesJWTVerifiedByPublishedJWKS(t *testing.T) {
address := os.Getenv("WORKLOAD_STS_TEST_BAO_ADDR")
token := os.Getenv("WORKLOAD_STS_TEST_BAO_TOKEN")
if address == "" || token == "" {
t.Skip("set WORKLOAD_STS_TEST_BAO_ADDR and WORKLOAD_STS_TEST_BAO_TOKEN")
}
config := baoapi.DefaultConfig()
config.Address = address
client, err := baoapi.NewClient(config)
if err != nil {
t.Fatalf("create OpenBao client: %v", err)
}
client.SetToken(token)
signer, err := NewTransitSigner(client, "transit", "workload-sts")
if err != nil {
t.Fatalf("NewTransitSigner() error = %v", err)
}
issuer, err := NewIssuer(signer)
if err != nil {
t.Fatalf("NewIssuer() error = %v", err)
}
now := time.Now().UTC().Truncate(time.Second)
issuer.now = func() time.Time { return now }
compact, err := issuer.Sign(context.Background(), Claims{
Issuer: "https://identity.ad.ddupan.top",
Subject: "01993f4d-5e1a-7000-8000-000000000001",
PrincipalName: "ci/homelab-infra-plan",
Audience: []string{"https://bao.ad.ddupan.top:8200"},
IssuedAt: now.Unix(),
NotBefore: now.Unix(),
ExpiresAt: now.Add(5 * time.Minute).Unix(),
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
Scope: "bao.login",
})
if err != nil {
t.Fatalf("Sign() error = %v", err)
}
set, err := signer.JWKS(context.Background())
if err != nil {
t.Fatalf("JWKS() error = %v", err)
}
encodedSet, err := json.Marshal(set)
if err != nil {
t.Fatalf("marshal JWKS: %v", err)
}
var decodedSet jose.JSONWebKeySet
if err := json.Unmarshal(encodedSet, &decodedSet); err != nil {
t.Fatalf("decode JWKS with go-jose: %v", err)
}
jws, err := jose.ParseSigned(compact, []jose.SignatureAlgorithm{jose.RS256})
if err != nil {
t.Fatalf("parse compact JWT: %v", err)
}
keyID := jws.Signatures[0].Header.KeyID
keys := decodedSet.Key(keyID)
if len(keys) != 1 {
t.Fatalf("JWKS keys for %q = %d, want 1", keyID, len(keys))
}
if _, err := jws.Verify(keys[0].Key); err != nil {
t.Fatalf("published JWKS rejected Transit signature: %v", err)
}
}
+76
View File
@@ -0,0 +1,76 @@
package telemetry
import (
"context"
"errors"
"fmt"
"go.opentelemetry.io/otel/metric"
metricnoop "go.opentelemetry.io/otel/metric/noop"
sdkmetric "go.opentelemetry.io/otel/sdk/metric"
sdkresource "go.opentelemetry.io/otel/sdk/resource"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
"go.opentelemetry.io/otel/trace"
tracenoop "go.opentelemetry.io/otel/trace/noop"
)
type Providers struct {
Tracer trace.TracerProvider
Meter metric.MeterProvider
shutdown func(context.Context) error
}
type Config struct {
Resource *sdkresource.Resource
TraceExporter sdktrace.SpanExporter
MetricReader sdkmetric.Reader
}
func New(config Config) (*Providers, error) {
if config.Resource == nil {
return nil, errors.New("OpenTelemetry resource is required")
}
providers := &Providers{
Tracer: tracenoop.NewTracerProvider(),
Meter: metricnoop.NewMeterProvider(),
}
var tracerProvider *sdktrace.TracerProvider
if config.TraceExporter != nil {
tracerProvider = sdktrace.NewTracerProvider(
sdktrace.WithBatcher(config.TraceExporter),
sdktrace.WithResource(config.Resource),
)
providers.Tracer = tracerProvider
}
var meterProvider *sdkmetric.MeterProvider
if config.MetricReader != nil {
meterProvider = sdkmetric.NewMeterProvider(
sdkmetric.WithReader(config.MetricReader),
sdkmetric.WithResource(config.Resource),
)
providers.Meter = meterProvider
}
providers.shutdown = func(ctx context.Context) error {
var shutdownErrors []error
if meterProvider != nil {
shutdownErrors = append(shutdownErrors, meterProvider.Shutdown(ctx))
}
if tracerProvider != nil {
shutdownErrors = append(shutdownErrors, tracerProvider.Shutdown(ctx))
}
if err := errors.Join(shutdownErrors...); err != nil {
return fmt.Errorf("shutdown OpenTelemetry providers: %w", err)
}
return nil
}
return providers, nil
}
func (p *Providers) Shutdown(ctx context.Context) error {
return p.shutdown(ctx)
}
+27
View File
@@ -0,0 +1,27 @@
package telemetry
import (
"context"
"testing"
sdkresource "go.opentelemetry.io/otel/sdk/resource"
)
func TestNewWithoutExportersUsesNoopProviders(t *testing.T) {
providers, err := New(Config{Resource: sdkresource.Empty()})
if err != nil {
t.Fatalf("New() error = %v", err)
}
if providers.Tracer == nil || providers.Meter == nil {
t.Fatal("New() returned nil provider")
}
if err := providers.Shutdown(context.Background()); err != nil {
t.Fatalf("Shutdown() error = %v", err)
}
}
func TestNewRequiresResource(t *testing.T) {
if _, err := New(Config{}); err == nil {
t.Fatal("New() error = nil, want missing resource error")
}
}