diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..f9de0b4 --- /dev/null +++ b/Makefile @@ -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 diff --git a/docs/architecture.md b/docs/architecture.md index e64492b..aa262fe 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -52,7 +52,8 @@ exchange、client store 适配和 Transit signer,不能减少本项目的核 不引入完整 Authorization Server 框架。业务代码仍将 credential verification、授权和 token materialization 保持为独立边界。详见 [`poc.md`](poc.md)。 -HTTP router 外层使用 `otelhttp` 生成 server span 和标准 HTTP metrics。span 与 metric 只使用 +启用 OpenTelemetry 时,启动层可以在纯 chi router 外使用 `otelhttp` 生成 server span 和 +标准 HTTP metrics;未启用时不安装该 wrapper。span 与 metric 只使用 固定路由模板及受控低基数字段,不捕获请求/响应 body、Authorization header、 `subject_token`、输出 token、`client_id`、principal 或 `jti`。trace 与 metric provider 通过启动依赖注入;未配置 exporter 时保持 no-op,不把 telemetry 输出到标准输出。 diff --git a/docs/specification.md b/docs/specification.md index 423beb3..f56d66a 100644 --- a/docs/specification.md +++ b/docs/specification.md @@ -483,7 +483,8 @@ metrics 只能使用 verifier、audience、profile 和结果等受控低基数 使用 `go-jose/v4` 做 JOSE/JWK 互操作,并通过官方 `openbao/api/v2` 调用 Transit。 - HTTP 路由层使用 `chi/v5`,保持 handler 和 middleware 与标准 `net/http` 兼容;OAuth 表单字段仍由协议层显式解析,不使用自动 request binding。 -- 可观测性使用 OpenTelemetry Go trace 与 metric SDK,HTTP server 使用 `otelhttp`;SDK +- 可选的可观测性使用 OpenTelemetry Go trace 与 metric SDK,启用时 HTTP server 使用 + `otelhttp`;未配置 telemetry 时纯 chi handler 不创建 OTel instrumentation。SDK exporter/reader 由进程启动配置注入,不在协议层固定 OTLP gRPC、OTLP HTTP 或具体后端。 第一阶段日志使用结构化日志并关联 trace/span ID,不要求启用 OpenTelemetry Logs SDK。 diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..e2db0ea --- /dev/null +++ b/go.mod @@ -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/go-viper/mapstructure/v2 v2.5.0 + 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/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 +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..4744215 --- /dev/null +++ b/go.sum @@ -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= diff --git a/internal/httpapi/router.go b/internal/httpapi/router.go new file mode 100644 index 0000000..21a7c35 --- /dev/null +++ b/internal/httpapi/router.go @@ -0,0 +1,31 @@ +package httpapi + +import ( + "errors" + "net/http" + + "github.com/go-chi/chi/v5" +) + +type Endpoints struct { + Metadata http.Handler + Token http.Handler + JWKS http.Handler + Health http.Handler + Ready http.Handler +} + +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") + } + + 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) + + return router, nil +} diff --git a/internal/httpapi/router_test.go b/internal/httpapi/router_test.go new file mode 100644 index 0000000..735fa4e --- /dev/null +++ b/internal/httpapi/router_test.go @@ -0,0 +1,55 @@ +package httpapi + +import ( + "net/http" + "net/http/httptest" + "testing" +) + +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, + }) + 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{}); err == nil { + t.Fatal("NewRouter() error = nil, want missing endpoint error") + } +} diff --git a/internal/protocol/token_exchange.go b/internal/protocol/token_exchange.go new file mode 100644 index 0000000..f0d5751 --- /dev/null +++ b/internal/protocol/token_exchange.go @@ -0,0 +1,81 @@ +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") + } + 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") + } + } + + return request, nil +} + +func requestError(code, description string) *RequestError { + return &RequestError{Code: code, Description: description} +} diff --git a/internal/protocol/token_exchange_test.go b/internal/protocol/token_exchange_test.go new file mode 100644 index 0000000..7a3564e --- /dev/null +++ b/internal/protocol/token_exchange_test.go @@ -0,0 +1,78 @@ +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: "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"}, + } + + 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 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}, + "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"}, + } +} diff --git a/internal/signing/jwt.go b/internal/signing/jwt.go new file mode 100644 index 0000000..d661613 --- /dev/null +++ b/internal/signing/jwt.go @@ -0,0 +1,123 @@ +package signing + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "strings" + "time" +) + +const maximumTTL = 5 * time.Minute + +type RS256Signer interface { + ActiveKey(context.Context) (SigningKey, error) + SignRS256(context.Context, SigningKey, []byte) ([]byte, error) +} + +type SigningKey struct { + ID string + Version int +} + +type IssueRequest struct { + Subject string + PrincipalName string + Audience string + JWTID string + Scope string +} + +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 { + issuer string + ttl time.Duration + signer RS256Signer +} + +func NewIssuer(issuer string, ttl time.Duration, signer RS256Signer) (*Issuer, error) { + if issuer == "" || signer == nil { + return nil, errors.New("issuer and signer are required") + } + if ttl <= 0 || ttl > maximumTTL { + return nil, errors.New("TTL must be between zero and five minutes") + } + return &Issuer{issuer: issuer, ttl: ttl, signer: signer}, nil +} + +func (i *Issuer) Issue(ctx context.Context, request IssueRequest) (string, error) { + return i.issueAt(ctx, request, time.Now()) +} + +func (i *Issuer) issueAt(ctx context.Context, request IssueRequest, now time.Time) (string, error) { + if err := validateIssueRequest(request); 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") + } + + claims := Claims{ + Issuer: i.issuer, + Subject: request.Subject, + PrincipalName: request.PrincipalName, + Audience: []string{request.Audience}, + IssuedAt: now.Unix(), + NotBefore: now.Unix(), + ExpiresAt: now.Add(i.ttl).Unix(), + JWTID: request.JWTID, + Scope: request.Scope, + } + signingInput, err := encodeSigningInput(key.ID, claims) + if err != nil { + return "", err + } + signature, err := i.signer.SignRS256(ctx, key, []byte(signingInput)) + if err != nil { + return "", fmt.Errorf("sign JWT: %w", err) + } + return serializeCompact(signingInput, signature), nil +} + +func validateIssueRequest(request IssueRequest) error { + if request.Subject == "" || request.Audience == "" || request.JWTID == "" || request.Scope == "" { + return errors.New("required JWT input is missing") + } + if strings.ContainsAny(request.Subject, "\r\n") || strings.ContainsAny(request.JWTID, "\r\n") { + return errors.New("invalid JWT input") + } + return nil +} + +func encodeSigningInput(keyID string, claims Claims) (string, error) { + header, err := json.Marshal(map[string]string{"alg": "RS256", "kid": keyID, "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) + } + return base64.RawURLEncoding.EncodeToString(header) + "." + base64.RawURLEncoding.EncodeToString(payload), nil +} + +func serializeCompact(signingInput string, signature []byte) string { + return signingInput + "." + base64.RawURLEncoding.EncodeToString(signature) +} diff --git a/internal/signing/jwt_test.go b/internal/signing/jwt_test.go new file mode 100644 index 0000000..beb6cf9 --- /dev/null +++ b/internal/signing/jwt_test.go @@ -0,0 +1,95 @@ +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("https://identity.ad.ddupan.top", 5*time.Minute, signer) + if err != nil { + t.Fatalf("NewIssuer() error = %v", err) + } + now := time.Unix(1_789_062_000, 0) + + token, err := issuer.issueAt(context.Background(), IssueRequest{ + Subject: "01993f4d-5e1a-7000-8000-000000000001", + PrincipalName: "ci/homelab-infra-plan", + Audience: "https://bao.ad.ddupan.top:8200", + JWTID: "01993f4d-5e1a-7000-8000-000000000002", + Scope: "bao.login", + }, now) + 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) + } + _, err = NewIssuer("https://identity.ad.ddupan.top", 5*time.Minute+time.Second, &localRS256Signer{keyID: "test", key: privateKey}) + if err == nil { + t.Fatal("NewIssuer() error = nil, want excessive lifetime error") + } +} + +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[:]) +} diff --git a/internal/signing/transit.go b/internal/signing/transit.go new file mode 100644 index 0000000..607f233 --- /dev/null +++ b/internal/signing/transit.go @@ -0,0 +1,154 @@ +package signing + +import ( + "context" + "crypto/rsa" + "crypto/x509" + "encoding/base64" + "encoding/pem" + "errors" + "fmt" + "sort" + "strconv" + "strings" + + "github.com/go-jose/go-jose/v4" + "github.com/go-viper/mapstructure/v2" + baoapi "github.com/openbao/openbao/api/v2" +) + +type TransitSigner struct { + client *baoapi.Client + mount string + key string +} + +type transitKeyData struct { + LatestVersion int `mapstructure:"latest_version"` + Keys map[string]transitKeyVersion `mapstructure:"keys"` +} + +type transitKeyVersion struct { + PublicKey string `mapstructure:"public_key"` +} + +func NewTransitSigner(client *baoapi.Client, mount, key string) (*TransitSigner, error) { + if client == nil { + return nil, errors.New("OpenBao client is required") + } + mount, key = strings.Trim(mount, "/"), 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) { + data, err := s.readKeyData(ctx) + if err != nil { + return SigningKey{}, err + } + if data.LatestVersion < 1 { + return SigningKey{}, errors.New("read Transit key: invalid latest_version") + } + return SigningKey{ID: s.keyID(data.LatestVersion), Version: data.LatestVersion}, nil +} + +func (s *TransitSigner) SignRS256(ctx context.Context, key SigningKey, input []byte) ([]byte, error) { + if key.Version < 1 || key.ID != s.keyID(key.Version) { + 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") + } + return decodeTransitSignature(encoded, key.Version) +} + +func (s *TransitSigner) JWKS(ctx context.Context) (jose.JSONWebKeySet, error) { + data, err := s.readKeyData(ctx) + if err != nil { + return jose.JSONWebKeySet{}, err + } + versions := make([]int, 0, len(data.Keys)) + for text := range data.Keys { + version, err := strconv.Atoi(text) + 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(versions))} + for _, version := range versions { + publicKey, err := parseRSAPublicKey(data.Keys[strconv.Itoa(version)].PublicKey) + 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) readKeyData(ctx context.Context) (transitKeyData, error) { + secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath()) + if err != nil { + return transitKeyData{}, fmt.Errorf("read Transit key: %w", err) + } + if secret == nil || secret.Data == nil { + return transitKeyData{}, errors.New("read Transit key: empty response") + } + var data transitKeyData + if err := mapstructure.Decode(secret.Data, &data); err != nil { + return transitKeyData{}, fmt.Errorf("decode Transit key metadata: %w", err) + } + return data, nil +} + +func decodeTransitSignature(encoded string, version int) ([]byte, error) { + parts := strings.SplitN(encoded, ":", 3) + if len(parts) != 3 || parts[0] != "vault" || parts[1] != "v"+strconv.Itoa(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) 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 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 +} diff --git a/internal/signing/transit_integration_test.go b/internal/signing/transit_integration_test.go new file mode 100644 index 0000000..939df42 --- /dev/null +++ b/internal/signing/transit_integration_test.go @@ -0,0 +1,75 @@ +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("https://identity.ad.ddupan.top", 5*time.Minute, signer) + if err != nil { + t.Fatalf("NewIssuer() error = %v", err) + } + now := time.Now().UTC().Truncate(time.Second) + + compact, err := issuer.issueAt(context.Background(), IssueRequest{ + Subject: "01993f4d-5e1a-7000-8000-000000000001", + PrincipalName: "ci/homelab-infra-plan", + Audience: "https://bao.ad.ddupan.top:8200", + JWTID: "01993f4d-5e1a-7000-8000-000000000002", + Scope: "bao.login", + }, now) + 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) + } +} diff --git a/internal/telemetry/http.go b/internal/telemetry/http.go new file mode 100644 index 0000000..3443a0c --- /dev/null +++ b/internal/telemetry/http.go @@ -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...) +} diff --git a/internal/telemetry/http_test.go b/internal/telemetry/http_test.go new file mode 100644 index 0000000..313b0af --- /dev/null +++ b/internal/telemetry/http_test.go @@ -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) + } + } +} diff --git a/internal/telemetry/providers.go b/internal/telemetry/providers.go new file mode 100644 index 0000000..7666d3f --- /dev/null +++ b/internal/telemetry/providers.go @@ -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) +} diff --git a/internal/telemetry/providers_test.go b/internal/telemetry/providers_test.go new file mode 100644 index 0000000..ca098f6 --- /dev/null +++ b/internal/telemetry/providers_test.go @@ -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") + } +}