diff --git a/CHANGELOG.md b/CHANGELOG.md index 4abae17..00cab0a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,6 +14,16 @@ Jwtinfo: add sentinel and typed errors for errors.Is/As, and route CLI/MCP display through errdisp domain leaves. + Jwks: add sentinel errors for PEM decode, unsupported key, and non-public key; route CLI display through errdisp. + + Requests: add sentinel and typed errors for client/validation leaves (nil client, empty args, serverName URL, wrong transport, timeout, proxyproto, URI, transport URL). + + Mcp: add sentinel and typed errors for tool-boundary input rules (config/token sources, required fields, TLS info, encrypted key env, validation joins). + + Certinfo: add InvalidTLSEndpointError for host:port parse failures. + + Jwtinfo: add typed errors for base64 JWT parts, JWT parse sources, and invalid request-values JSON. + ### Fix Devenv: prefer httpbin on 127.0.0.1:8081 and proxy nginx upstreams through the allocated httpbin port so `devenv test` keeps working when the preferred port is already taken; fail fast in request integration tests with `set -e`, enable `pipefail` on success-case request leaf pipelines, and assert request exit status separately from expected error text. diff --git a/internal/certinfo/certinfo.go b/internal/certinfo/certinfo.go index 6310368..5358819 100644 --- a/internal/certinfo/certinfo.go +++ b/internal/certinfo/certinfo.go @@ -187,7 +187,7 @@ func (c *Config) SetTLSEndpoint(ctx context.Context, hostport string) error { if hostport != emptyString { eHost, ePort, err := net.SplitHostPort(hostport) if err != nil { - return fmt.Errorf("invalid TLS endpoint %q: %w", hostport, err) + return fmt.Errorf("%w: %w", &InvalidTLSEndpointError{Endpoint: hostport}, err) } c.TLSEndpoint = hostport diff --git a/internal/certinfo/errors.go b/internal/certinfo/errors.go index 24e84a5..bc8e4a1 100644 --- a/internal/certinfo/errors.go +++ b/internal/certinfo/errors.go @@ -17,6 +17,7 @@ var ( ErrEmptyArg = errors.New("empty string provided as argument") ErrNoCertsInFile = errors.New("no valid certificates found in file") ErrUnrecognizedKeyType = errors.New("unrecognized private key type") + ErrInvalidTLSEndpoint = errors.New("invalid TLS endpoint") ) // EmptyArgError is returned when a required string argument is empty. @@ -66,3 +67,19 @@ func (e *UnrecognizedKeyTypeError) Error() string { func (*UnrecognizedKeyTypeError) Is(target error) bool { return target == ErrUnrecognizedKeyType } + +// InvalidTLSEndpointError is returned when host:port cannot be split. +// errors.Is(err, ErrInvalidTLSEndpoint) is true. +type InvalidTLSEndpointError struct { + Endpoint string +} + +// Error returns a message including the invalid endpoint. +func (e *InvalidTLSEndpointError) Error() string { + return fmt.Sprintf("invalid TLS endpoint %q", e.Endpoint) +} + +// Is reports whether target is ErrInvalidTLSEndpoint. +func (*InvalidTLSEndpointError) Is(target error) bool { + return target == ErrInvalidTLSEndpoint +} diff --git a/internal/certinfo/errors_test.go b/internal/certinfo/errors_test.go index becd32f..853d662 100644 --- a/internal/certinfo/errors_test.go +++ b/internal/certinfo/errors_test.go @@ -8,6 +8,7 @@ import ( "github.com/stretchr/testify/require" ) +// TestCertinfo_errorSentinels_Is checks errors.Is against certinfo sentinels. func TestCertinfo_errorSentinels_Is(t *testing.T) { t.Parallel() @@ -66,6 +67,16 @@ func TestCertinfo_errorSentinels_Is(t *testing.T) { err: ErrUnsupportedPublicKey, target: ErrUnsupportedPublicKey, }, + { + name: "InvalidTLSEndpointError", + err: &InvalidTLSEndpointError{Endpoint: "bad"}, + target: ErrInvalidTLSEndpoint, + }, + { + name: "InvalidTLSEndpointError wrapped", + err: fmt.Errorf("set: %w", &InvalidTLSEndpointError{Endpoint: "x"}), + target: ErrInvalidTLSEndpoint, + }, } for _, tt := range tests { @@ -76,6 +87,7 @@ func TestCertinfo_errorSentinels_Is(t *testing.T) { } } +// TestCertinfo_errorTypes_AsType checks errors.AsType for typed certinfo errors. func TestCertinfo_errorTypes_AsType(t *testing.T) { t.Parallel() @@ -93,4 +105,9 @@ func TestCertinfo_errorTypes_AsType(t *testing.T) { gotKeyType, ok := errors.AsType[*UnrecognizedKeyTypeError](keyType) require.True(t, ok) require.Equal(t, "CERTIFICATE", gotKeyType.Type) + + tlsEp := fmt.Errorf("wrap: %w", &InvalidTLSEndpointError{Endpoint: "no-port"}) + gotTLS, ok := errors.AsType[*InvalidTLSEndpointError](tlsEp) + require.True(t, ok) + require.Equal(t, "no-port", gotTLS.Endpoint) } diff --git a/internal/cmd/jwks.go b/internal/cmd/jwks.go index 07c99b4..fcc7c91 100644 --- a/internal/cmd/jwks.go +++ b/internal/cmd/jwks.go @@ -4,6 +4,7 @@ import ( "fmt" "github.com/spf13/cobra" + "github.com/xenos76/https-wrench/internal/errdisp" "github.com/xenos76/https-wrench/internal/jwks" "github.com/xenos76/https-wrench/internal/style" ) @@ -31,7 +32,7 @@ Examples: Run: func(cmd *cobra.Command, _ []string) { jwksJSON, err := jwks.GenerateJWKS(cmd.Context(), jwksPublicKeyFile, jwksKID) if err != nil { - cmd.PrintErrf("Error generating JWKS: %s\n", err) + cmd.PrintErrf("Error generating JWKS: %s\n", errdisp.FormatCause(err)) return } diff --git a/internal/errdisp/errdisp.go b/internal/errdisp/errdisp.go index 0c0678f..da08099 100644 --- a/internal/errdisp/errdisp.go +++ b/internal/errdisp/errdisp.go @@ -3,8 +3,10 @@ Copyright © 2025 Zeno Belli xeno@os76.xyz */ // Package errdisp formats errors for CLI and MCP user boundaries. -// It holds no sentinels; domain identity stays in packages such as certinfo -// and jwtinfo. +// It holds no sentinels; domain identity stays in packages such as certinfo, +// jwtinfo, jwks, and requests. MCP tool-boundary errors stay in package mcp +// (errdisp cannot import mcp without a cycle); bare MCP leaves still format +// via Error() when Format finds no registered domain leaf. package errdisp import ( @@ -12,7 +14,9 @@ import ( "strings" "github.com/xenos76/https-wrench/internal/certinfo" + "github.com/xenos76/https-wrench/internal/jwks" "github.com/xenos76/https-wrench/internal/jwtinfo" + "github.com/xenos76/https-wrench/internal/requests" ) // Cause returns the deepest single-cause unwrap of err. @@ -73,6 +77,23 @@ func Format(err error) string { // domainLeaf returns a domain leaf message when err matches a known failure. func domainLeaf(err error) (string, bool) { + if msg, ok := certinfoTypedLeaf(err); ok { + return msg, true + } + + if msg, ok := jwtinfoTypedLeaf(err); ok { + return msg, true + } + + if msg, ok := requestsTypedLeaf(err); ok { + return msg, true + } + + return domainSentinelLeaf(err) +} + +// certinfoTypedLeaf returns a certinfo typed-error leaf message when matched. +func certinfoTypedLeaf(err error) (string, bool) { if empty, ok := errors.AsType[*certinfo.EmptyArgError](err); ok { return empty.Error(), true } @@ -85,6 +106,15 @@ func domainLeaf(err error) (string, bool) { return keyType.Error(), true } + if tlsEp, ok := errors.AsType[*certinfo.InvalidTLSEndpointError](err); ok { + return tlsEp.Error(), true + } + + return "", false +} + +// jwtinfoTypedLeaf returns a jwtinfo typed-error leaf message when matched. +func jwtinfoTypedLeaf(err error) (string, bool) { if empty, ok := errors.AsType[*jwtinfo.EmptyArgError](err); ok { return empty.Error(), true } @@ -117,31 +147,49 @@ func domainLeaf(err error) (string, bool) { return thr.Error(), true } - for _, s := range []error{ - certinfo.ErrNilReader, - certinfo.ErrPEMDecode, - certinfo.ErrCertPoolFromFile, - certinfo.ErrNoCertsInConfig, - certinfo.ErrUnsupportedKey, - certinfo.ErrUnsupportedPublicKey, - certinfo.ErrEmptyArg, - certinfo.ErrNoCertsInFile, - certinfo.ErrUnrecognizedKeyType, - jwtinfo.ErrNilBodyReader, - jwtinfo.ErrEmptyRequestValues, - jwtinfo.ErrEmptyArg, - jwtinfo.ErrInvalidJWTFormat, - jwtinfo.ErrInvalidHeaderJSON, - jwtinfo.ErrInvalidClaimsJSON, - jwtinfo.ErrEmptyClaims, - jwtinfo.ErrClaimMissing, - jwtinfo.ErrClaimNotNumeric, - jwtinfo.ErrInvalidKV, - jwtinfo.ErrEmptyParamName, - jwtinfo.ErrInvalidRenewThreshold, - jwtinfo.ErrTokenLifetimeInvalid, - jwtinfo.ErrTokenRequestStatus, - } { + if b64, ok := errors.AsType[*jwtinfo.InvalidBase64PartError](err); ok { + return b64.Error(), true + } + + if parse, ok := errors.AsType[*jwtinfo.JWTParseError](err); ok { + return parse.Error(), true + } + + return "", false +} + +// requestsTypedLeaf returns a requests typed-error leaf message when matched. +func requestsTypedLeaf(err error) (string, bool) { + if empty, ok := errors.AsType[*requests.EmptyArgError](err); ok { + return empty.Error(), true + } + + if snURL, ok := errors.AsType[*requests.ServerNameURLError](err); ok { + return snURL.Error(), true + } + + if wt, ok := errors.AsType[*requests.WrongTransportError](err); ok { + return wt.Error(), true + } + + if to, ok := errors.AsType[*requests.InvalidTimeoutError](err); ok { + return to.Error(), true + } + + if uri, ok := errors.AsType[*requests.InvalidURIError](err); ok { + return uri.Error(), true + } + + if turl, ok := errors.AsType[*requests.InvalidTransportURLError](err); ok { + return turl.Error(), true + } + + return "", false +} + +// domainSentinelLeaf returns a package-sentinel leaf message when errors.Is matches. +func domainSentinelLeaf(err error) (string, bool) { + for _, s := range domainSentinels { if errors.Is(err, s) { return s.Error(), true } @@ -150,6 +198,52 @@ func domainLeaf(err error) (string, bool) { return "", false } +// domainSentinels lists stable package sentinels matched with errors.Is. +var domainSentinels = []error{ + certinfo.ErrNilReader, + certinfo.ErrPEMDecode, + certinfo.ErrCertPoolFromFile, + certinfo.ErrNoCertsInConfig, + certinfo.ErrUnsupportedKey, + certinfo.ErrUnsupportedPublicKey, + certinfo.ErrEmptyArg, + certinfo.ErrNoCertsInFile, + certinfo.ErrUnrecognizedKeyType, + certinfo.ErrInvalidTLSEndpoint, + jwtinfo.ErrNilBodyReader, + jwtinfo.ErrEmptyRequestValues, + jwtinfo.ErrEmptyArg, + jwtinfo.ErrInvalidJWTFormat, + jwtinfo.ErrInvalidHeaderJSON, + jwtinfo.ErrInvalidClaimsJSON, + jwtinfo.ErrEmptyClaims, + jwtinfo.ErrClaimMissing, + jwtinfo.ErrClaimNotNumeric, + jwtinfo.ErrInvalidKV, + jwtinfo.ErrEmptyParamName, + jwtinfo.ErrInvalidRenewThreshold, + jwtinfo.ErrTokenLifetimeInvalid, + jwtinfo.ErrTokenRequestStatus, + jwtinfo.ErrInvalidBase64Header, + jwtinfo.ErrInvalidBase64Claims, + jwtinfo.ErrJWTParse, + jwtinfo.ErrInvalidRequestJSON, + jwks.ErrPEMDecode, + jwks.ErrUnsupportedPublicKey, + jwks.ErrNotPublicKey, + requests.ErrMethodNotFound, + requests.ErrNilClient, + requests.ErrEmptyArg, + requests.ErrServerNameIsURL, + requests.ErrWrongTransport, + requests.ErrInvalidTimeout, + requests.ErrProxyProtoNeedsOverride, + requests.ErrProxyProtoDisabled, + requests.ErrTransportOverrideRequired, + requests.ErrInvalidURI, + requests.ErrInvalidTransportURL, +} + // topLabel returns the outermost wrap text without the unwrapped suffix. func topLabel(err error) string { u := errors.Unwrap(err) diff --git a/internal/errdisp/errdisp_test.go b/internal/errdisp/errdisp_test.go index 3d09e99..778fd29 100644 --- a/internal/errdisp/errdisp_test.go +++ b/internal/errdisp/errdisp_test.go @@ -11,9 +11,12 @@ import ( "github.com/stretchr/testify/require" "github.com/xenos76/https-wrench/internal/certinfo" + "github.com/xenos76/https-wrench/internal/jwks" "github.com/xenos76/https-wrench/internal/jwtinfo" + "github.com/xenos76/https-wrench/internal/requests" ) +// TestCause checks Cause unwraps to the deepest single-cause leaf. func TestCause(t *testing.T) { t.Parallel() @@ -80,6 +83,27 @@ func TestFormatCause(t *testing.T) { err := fmt.Errorf("claims: %w", &jwtinfo.ClaimError{Claim: "exp", Kind: jwtinfo.ClaimMissing}) require.Equal(t, "exp claim missing", FormatCause(err)) }) + + t.Run("jwks PEM sentinel", func(t *testing.T) { + t.Parallel() + + err := fmt.Errorf("generate: %w", jwks.ErrPEMDecode) + require.Equal(t, jwks.ErrPEMDecode.Error(), FormatCause(err)) + }) + + t.Run("requests empty arg", func(t *testing.T) { + t.Parallel() + + err := fmt.Errorf("SetServerName error: %w", &requests.EmptyArgError{Name: "serverName"}) + require.Equal(t, "empty string provided as serverName", FormatCause(err)) + }) + + t.Run("certinfo invalid TLS endpoint", func(t *testing.T) { + t.Parallel() + + err := fmt.Errorf("set: %w", &certinfo.InvalidTLSEndpointError{Endpoint: "bad"}) + require.Equal(t, `invalid TLS endpoint "bad"`, FormatCause(err)) + }) } // TestFormat checks Format surfaces domain leaves or top label plus cause. diff --git a/internal/jwks/errors.go b/internal/jwks/errors.go new file mode 100644 index 0000000..8887996 --- /dev/null +++ b/internal/jwks/errors.go @@ -0,0 +1,16 @@ +package jwks + +import ( + "errors" +) + +// Package-level sentinels for stable jwks failure conditions. +// Match them with errors.Is after wrapping; Error() strings stay human-facing. +var ( + ErrPEMDecode = errors.New("failed to decode PEM block from public key file") + ErrUnsupportedPublicKey = errors.New("unsupported or invalid public key format") + ErrNotPublicKey = errors.New( + "the provided file does not contain a supported public key " + + "(it might be a private key or an unsupported format)", + ) +) diff --git a/internal/jwks/errors_test.go b/internal/jwks/errors_test.go new file mode 100644 index 0000000..08b0388 --- /dev/null +++ b/internal/jwks/errors_test.go @@ -0,0 +1,48 @@ +package jwks + +import ( + "errors" + "fmt" + "testing" + + "github.com/stretchr/testify/require" +) + +// TestJwks_errorSentinels_Is checks errors.Is against jwks sentinels. +func TestJwks_errorSentinels_Is(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + target error + }{ + { + name: "ErrPEMDecode", + err: ErrPEMDecode, + target: ErrPEMDecode, + }, + { + name: "ErrPEMDecode wrapped", + err: fmt.Errorf("generate: %w", ErrPEMDecode), + target: ErrPEMDecode, + }, + { + name: "ErrUnsupportedPublicKey wrapped", + err: fmt.Errorf("%w: %w", ErrUnsupportedPublicKey, errors.New("bad key")), + target: ErrUnsupportedPublicKey, + }, + { + name: "ErrNotPublicKey", + err: ErrNotPublicKey, + target: ErrNotPublicKey, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.ErrorIs(t, tt.err, tt.target) + }) + } +} diff --git a/internal/jwks/jwks.go b/internal/jwks/jwks.go index f90832c..6f2c08b 100644 --- a/internal/jwks/jwks.go +++ b/internal/jwks/jwks.go @@ -12,7 +12,6 @@ import ( "encoding/base64" "encoding/json" "encoding/pem" - "errors" "fmt" "os" @@ -29,12 +28,12 @@ func GenerateJWKS(ctx context.Context, publicKeyFile string, kid string) (string block, _ := pem.Decode(keyPEM) if block == nil { - return "", errors.New("failed to decode PEM block from public key file") + return "", ErrPEMDecode } key, err := jwkset.LoadX509KeyInfer(block) if err != nil { - return "", fmt.Errorf("unsupported or invalid public key format: %w", err) + return "", fmt.Errorf("%w: %w", ErrUnsupportedPublicKey, err) } // Ensure the key is a public key @@ -42,8 +41,7 @@ func GenerateJWKS(ctx context.Context, publicKeyFile string, kid string) (string case *rsa.PublicKey, *ecdsa.PublicKey, ed25519.PublicKey: // Valid public key types default: - return "", errors.New("the provided file does not contain a supported public key " + - "(it might be a private key or an unsupported format)") + return "", ErrNotPublicKey } if kid == "" { diff --git a/internal/jwks/jwks_test.go b/internal/jwks/jwks_test.go index b35642c..6a63f0a 100644 --- a/internal/jwks/jwks_test.go +++ b/internal/jwks/jwks_test.go @@ -104,8 +104,7 @@ func TestGenerateJWKS_Errors(t *testing.T) { require.NoError(t, err) _, err = GenerateJWKS(context.Background(), invalidFile, "") - require.Error(t, err) - require.ErrorContains(t, err, "failed to decode PEM block") + require.ErrorIs(t, err, ErrPEMDecode) }) t.Run("Unsupported block type", func(t *testing.T) { @@ -116,8 +115,7 @@ func TestGenerateJWKS_Errors(t *testing.T) { file.Close() _, err := GenerateJWKS(context.Background(), path, "") - require.Error(t, err) - require.ErrorContains(t, err, "unsupported or invalid public key format") + require.ErrorIs(t, err, ErrUnsupportedPublicKey) }) t.Run("Private key rejected", func(t *testing.T) { @@ -129,7 +127,6 @@ func TestGenerateJWKS_Errors(t *testing.T) { file.Close() _, err := GenerateJWKS(context.Background(), path, "") - require.Error(t, err) - require.ErrorContains(t, err, "does not contain a supported public key") + require.ErrorIs(t, err, ErrNotPublicKey) }) } diff --git a/internal/jwtinfo/errors.go b/internal/jwtinfo/errors.go index a34a2cc..d543fe7 100644 --- a/internal/jwtinfo/errors.go +++ b/internal/jwtinfo/errors.go @@ -23,6 +23,10 @@ var ( ErrInvalidRenewThreshold = errors.New("renewThreshold must be between 0 and 100") ErrTokenLifetimeInvalid = errors.New("token lifetime is zero or negative") ErrTokenRequestStatus = errors.New("token request returned non-OK status") + ErrInvalidBase64Header = errors.New("unable to decode base64 header") + ErrInvalidBase64Claims = errors.New("unable to decode base64 claims") + ErrJWTParse = errors.New("unable to parse JWT") + ErrInvalidRequestJSON = errors.New("unable to parse JSON request values") ) // EmptyArgError is returned when a required string argument is empty. @@ -181,3 +185,56 @@ func (e *InvalidRenewThresholdError) Error() string { func (*InvalidRenewThresholdError) Is(target error) bool { return target == ErrInvalidRenewThreshold } + +// InvalidBase64PartError is returned when a JWT part fails base64 decoding. +// errors.Is matches ErrInvalidBase64Header or ErrInvalidBase64Claims from Part. +type InvalidBase64PartError struct { + Name string + Part string // "header" or "claims" +} + +// Error returns a message naming the failed base64 part. +func (e *InvalidBase64PartError) Error() string { + return fmt.Sprintf("unable to decode base64 %s from %s", e.Part, e.Name) +} + +// Is reports whether target matches the header or claims sentinel for Part. +func (e *InvalidBase64PartError) Is(target error) bool { + switch e.Part { + case "header": + return target == ErrInvalidBase64Header + case "claims": + return target == ErrInvalidBase64Claims + default: + return false + } +} + +// JWTParseError is returned when jwt library parse/verify fails. +// errors.Is(err, ErrJWTParse) is true. +type JWTParseError struct { + Source string // "AccessTokenRaw", "file", "HTTP response", "JWKS" + URL string // optional JWKS URL when Source is "JWKS" +} + +// Error returns a human-facing parse failure message for Source. +func (e *JWTParseError) Error() string { + switch e.Source { + case "file": + return "unable to parse JWT token from file" + case "HTTP response": + return "unable to parse JWT token from HTTP response" + case "JWKS": + return fmt.Sprintf( + "failed to parse the JWT AccessTokenRaw against JWKS URL %s", + e.URL, + ) + default: + return "unable to parse AccessTokenRaw" + } +} + +// Is reports whether target is ErrJWTParse. +func (*JWTParseError) Is(target error) bool { + return target == ErrJWTParse +} diff --git a/internal/jwtinfo/errors_test.go b/internal/jwtinfo/errors_test.go index 6d5924f..519f4dc 100644 --- a/internal/jwtinfo/errors_test.go +++ b/internal/jwtinfo/errors_test.go @@ -94,6 +94,31 @@ func TestJwtinfo_errorSentinels_Is(t *testing.T) { err: ErrTokenLifetimeInvalid, target: ErrTokenLifetimeInvalid, }, + { + name: "InvalidBase64PartError header", + err: &InvalidBase64PartError{Name: "AccessToken", Part: "header"}, + target: ErrInvalidBase64Header, + }, + { + name: "InvalidBase64PartError claims", + err: &InvalidBase64PartError{Name: "AccessToken", Part: "claims"}, + target: ErrInvalidBase64Claims, + }, + { + name: "JWTParseError", + err: &JWTParseError{Source: "file"}, + target: ErrJWTParse, + }, + { + name: "JWTParseError wrapped", + err: fmt.Errorf("%w: %w", &JWTParseError{Source: "HTTP response"}, errors.New("bad")), + target: ErrJWTParse, + }, + { + name: "ErrInvalidRequestJSON wrapped", + err: fmt.Errorf("%w: %w", ErrInvalidRequestJSON, errors.New("bad json")), + target: ErrInvalidRequestJSON, + }, } for _, tt := range tests { @@ -148,4 +173,15 @@ func TestJwtinfo_errorTypes_AsType(t *testing.T) { gotThr, ok := errors.AsType[*InvalidRenewThresholdError](thr) require.True(t, ok) require.InDelta(t, -1.0, gotThr.Value, 0.001) + + b64 := fmt.Errorf("wrap: %w", &InvalidBase64PartError{Name: "AccessToken", Part: "header"}) + gotB64, ok := errors.AsType[*InvalidBase64PartError](b64) + require.True(t, ok) + require.Equal(t, "header", gotB64.Part) + + parse := fmt.Errorf("wrap: %w", &JWTParseError{Source: "JWKS", URL: "https://example"}) + gotParse, ok := errors.AsType[*JWTParseError](parse) + require.True(t, ok) + require.Equal(t, "JWKS", gotParse.Source) + require.Equal(t, "https://example", gotParse.URL) } diff --git a/internal/jwtinfo/jwtinfo.go b/internal/jwtinfo/jwtinfo.go index d9e4101..987c599 100644 --- a/internal/jwtinfo/jwtinfo.go +++ b/internal/jwtinfo/jwtinfo.go @@ -134,10 +134,7 @@ func RequestToken(ctx context.Context, reqURL string, reqValues map[string]strin &jwt.RegisteredClaims{}, ) if err != nil { - return nil, fmt.Errorf( - "unable to parse JWT token from HTTP response: %w", - err, - ) + return nil, fmt.Errorf("%w: %w", &JWTParseError{Source: "HTTP response"}, err) } return t, nil @@ -158,10 +155,7 @@ func ReadTokenFromFile(fileName string) (*JwtTokenData, error) { &jwt.RegisteredClaims{}, ) if err != nil { - return nil, fmt.Errorf( - "unable to parse JWT token from file: %w", - err, - ) + return nil, fmt.Errorf("%w: %w", &JWTParseError{Source: "file"}, err) } return td, nil @@ -184,7 +178,7 @@ func ParseRequestJSONValues( err := json.Unmarshal([]byte(reqValues), &objmap) if err != nil { - return nil, fmt.Errorf("unable to parse Json request values: %w", err) + return nil, fmt.Errorf("%w: %w", ErrInvalidRequestJSON, err) } newMap := maps.Clone(reqValuesMap) @@ -273,7 +267,7 @@ func decodeToken(name, raw string) (header []byte, claims []byte, err error) { header, err = base64.RawURLEncoding.DecodeString(tokenB64Elements[0]) if err != nil { - return nil, nil, fmt.Errorf("unable to decode base64 header from %s: %w", name, err) + return nil, nil, fmt.Errorf("%w: %w", &InvalidBase64PartError{Name: name, Part: "header"}, err) } if !isValidJSON(header) { @@ -282,7 +276,7 @@ func decodeToken(name, raw string) (header []byte, claims []byte, err error) { claims, err = base64.RawURLEncoding.DecodeString(tokenB64Elements[1]) if err != nil { - return nil, nil, fmt.Errorf("unable to decode base64 claims from %s: %w", name, err) + return nil, nil, fmt.Errorf("%w: %w", &InvalidBase64PartError{Name: name, Part: "claims"}, err) } if !isValidJSON(claims) { @@ -299,10 +293,7 @@ func (jtd *JwtTokenData) ParseUnverified() error { &jwt.RegisteredClaims{}, ) if err != nil { - return fmt.Errorf( - "unable to parse AccessTokenRaw: %w", - err, - ) + return fmt.Errorf("%w: %w", &JWTParseError{Source: "AccessTokenRaw"}, err) } jtd.AccessTokenJwt = token @@ -335,11 +326,7 @@ func (jtd *JwtTokenData) ParseWithJWKS(ctx context.Context, jwksURL string, keyf jwks.Keyfunc, ) if err != nil { - return fmt.Errorf( - "failed to parse the JWT AccessTokenRaw against JWKS Url %s: %w", - jwksURL, - err, - ) + return fmt.Errorf("%w: %w", &JWTParseError{Source: "JWKS", URL: jwksURL}, err) } jtd.AccessTokenJwt = token diff --git a/internal/jwtinfo/jwtinfo_test.go b/internal/jwtinfo/jwtinfo_test.go index e3936bf..6dbd8b7 100644 --- a/internal/jwtinfo/jwtinfo_test.go +++ b/internal/jwtinfo/jwtinfo_test.go @@ -122,7 +122,7 @@ func TestParseRequestJSONValues(t *testing.T) { jsonStr: "{\"testKey2 :\"testValue2\", \"testKey3\":\"testValue3\"}", jsonRefMap: mapToValidJSON, requireError: true, - errorMsg: "unable to parse Json request values: invalid character 't' after object key", + errorMsg: "unable to parse JSON request values: invalid character 't' after object key", }, { name: "emptyJsonString", diff --git a/internal/mcp/errors.go b/internal/mcp/errors.go new file mode 100644 index 0000000..9a62c5d --- /dev/null +++ b/internal/mcp/errors.go @@ -0,0 +1,83 @@ +//nolint:revive // max-public-structs: typed domain errors for errors.Is/As +package mcp + +import ( + "errors" + "fmt" + "strings" +) + +// Package-level sentinels for stable MCP tool-boundary failure conditions. +// Match them with errors.Is after wrapping; Error() strings stay human-facing. +var ( + ErrExactlyOneConfigSource = errors.New("provide exactly one of configYaml or configPath") + ErrConfigSourceRequired = errors.New("configYaml or configPath is required") + ErrExactlyOneTokenSource = errors.New("provide exactly one of tokenFile or requestUrl") + ErrTokenSourceRequired = errors.New("tokenFile or requestUrl is required") + ErrRequestValuesRequired = errors.New("requestValues is required with requestUrl") + ErrPublicKeyFileRequired = errors.New("publicKeyFile is required") + ErrCertinfoInputRequired = errors.New( + "one of tlsEndpoint, certBundle, keyFile, or caBundle is required", + ) + ErrTLSInfoNeedsEndpoint = errors.New("tlsInfo requires tlsEndpoint") + ErrEncryptedKeyNeedsEnv = errors.New( + "encrypted private keys require CERTINFO_PKEY_PW under MCP", + ) + ErrNoJWTTokenData = errors.New("no JWT token data available") + ErrInvalidConfig = errors.New("invalid config") + ErrValidation = errors.New("validation failed") +) + +// ValidationError collects one or more validation messages. +// errors.Is(err, ErrValidation) is true; ErrInvalidConfig when Prefixed. +type ValidationError struct { + Messages []string + Prefixed bool // when true, Error() is "invalid config: ..." +} + +// Error returns the joined validation messages. +func (e *ValidationError) Error() string { + joined := strings.Join(e.Messages, "; ") + if e.Prefixed { + return "invalid config: " + joined + } + + return joined +} + +// Is reports whether target is ErrValidation or ErrInvalidConfig when Prefixed. +func (e *ValidationError) Is(target error) bool { + if target == ErrValidation { + return true + } + + return e.Prefixed && target == ErrInvalidConfig +} + +// RequiredFieldError is returned when a required tool input field is empty. +// errors.Is matches the sentinel for Field when known. +type RequiredFieldError struct { + Field string +} + +// Error returns a message naming the required field. +func (e *RequiredFieldError) Error() string { + switch e.Field { + case "requestValues": + return ErrRequestValuesRequired.Error() + default: + return fmt.Sprintf("%s is required", e.Field) + } +} + +// Is reports whether target matches the sentinel for Field. +func (e *RequiredFieldError) Is(target error) bool { + switch e.Field { + case "publicKeyFile": + return target == ErrPublicKeyFileRequired + case "requestValues": + return target == ErrRequestValuesRequired + default: + return false + } +} diff --git a/internal/mcp/errors_test.go b/internal/mcp/errors_test.go new file mode 100644 index 0000000..76b12d4 --- /dev/null +++ b/internal/mcp/errors_test.go @@ -0,0 +1,112 @@ +package mcp + +import ( + "errors" + "fmt" + "testing" + + "github.com/stretchr/testify/require" +) + +// TestMCP_errorSentinels_Is checks errors.Is against mcp sentinels. +// +//nolint:revive // function-length: table-driven sentinel coverage +func TestMCP_errorSentinels_Is(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + target error + }{ + { + name: "ErrExactlyOneConfigSource", + err: ErrExactlyOneConfigSource, + target: ErrExactlyOneConfigSource, + }, + { + name: "ErrConfigSourceRequired wrapped", + err: fmt.Errorf("load: %w", ErrConfigSourceRequired), + target: ErrConfigSourceRequired, + }, + { + name: "ErrExactlyOneTokenSource", + err: ErrExactlyOneTokenSource, + target: ErrExactlyOneTokenSource, + }, + { + name: "ErrTokenSourceRequired", + err: ErrTokenSourceRequired, + target: ErrTokenSourceRequired, + }, + { + name: "RequiredFieldError publicKeyFile", + err: &RequiredFieldError{Field: "publicKeyFile"}, + target: ErrPublicKeyFileRequired, + }, + { + name: "RequiredFieldError requestValues", + err: &RequiredFieldError{Field: "requestValues"}, + target: ErrRequestValuesRequired, + }, + { + name: "ErrCertinfoInputRequired", + err: ErrCertinfoInputRequired, + target: ErrCertinfoInputRequired, + }, + { + name: "ErrTLSInfoNeedsEndpoint", + err: ErrTLSInfoNeedsEndpoint, + target: ErrTLSInfoNeedsEndpoint, + }, + { + name: "ErrEncryptedKeyNeedsEnv", + err: ErrEncryptedKeyNeedsEnv, + target: ErrEncryptedKeyNeedsEnv, + }, + { + name: "ErrNoJWTTokenData", + err: ErrNoJWTTokenData, + target: ErrNoJWTTokenData, + }, + { + name: "ValidationError ErrValidation", + err: &ValidationError{Messages: []string{"a", "b"}}, + target: ErrValidation, + }, + { + name: "ValidationError ErrInvalidConfig when Prefixed", + err: &ValidationError{Messages: []string{"x"}, Prefixed: true}, + target: ErrInvalidConfig, + }, + { + name: "ValidationError Prefixed also ErrValidation", + err: &ValidationError{Messages: []string{"x"}, Prefixed: true}, + target: ErrValidation, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.ErrorIs(t, tt.err, tt.target) + }) + } +} + +// TestMCP_errorTypes_AsType checks errors.AsType for typed mcp errors. +func TestMCP_errorTypes_AsType(t *testing.T) { + t.Parallel() + + ve := fmt.Errorf("wrap: %w", &ValidationError{Messages: []string{"a"}, Prefixed: true}) + gotVE, ok := errors.AsType[*ValidationError](ve) + require.True(t, ok) + require.True(t, gotVE.Prefixed) + require.Equal(t, []string{"a"}, gotVE.Messages) + require.Equal(t, "invalid config: a", gotVE.Error()) + + rf := fmt.Errorf("wrap: %w", &RequiredFieldError{Field: "publicKeyFile"}) + gotRF, ok := errors.AsType[*RequiredFieldError](rf) + require.True(t, ok) + require.Equal(t, "publicKeyFile", gotRF.Field) +} diff --git a/internal/mcp/prompts.go b/internal/mcp/prompts.go index f88d6cc..29414c9 100644 --- a/internal/mcp/prompts.go +++ b/internal/mcp/prompts.go @@ -2,7 +2,6 @@ package mcp import ( "context" - "fmt" "strings" sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" @@ -37,7 +36,7 @@ func authorRequestsConfigPrompt(_ context.Context, req *sdkmcp.GetPromptRequest) yaml, errs := buildRequestsConfigYAML(input) if len(errs) > 0 { - return nil, fmt.Errorf("%s", strings.Join(errs, "; ")) + return nil, &ValidationError{Messages: errs} } exampleHints := exampleResourceHints(input) diff --git a/internal/mcp/tools.go b/internal/mcp/tools.go index 92957b0..81f9c5a 100644 --- a/internal/mcp/tools.go +++ b/internal/mcp/tools.go @@ -91,7 +91,7 @@ func requestsConfigTemplateHandler( ) (*sdkmcp.CallToolResult, requestsConfigTemplateOutput, error) { yaml, errs := buildRequestsConfigYAML(input) if len(errs) > 0 { - return nil, requestsConfigTemplateOutput{}, fmt.Errorf("%s", strings.Join(errs, "; ")) + return nil, requestsConfigTemplateOutput{}, &ValidationError{Messages: errs} } return nil, requestsConfigTemplateOutput{ConfigYAML: yaml}, nil diff --git a/internal/mcp/tools_exec.go b/internal/mcp/tools_exec.go index 45de98e..9847afb 100644 --- a/internal/mcp/tools_exec.go +++ b/internal/mcp/tools_exec.go @@ -3,7 +3,6 @@ package mcp import ( "bytes" "context" - "errors" "fmt" "io" "net/http" @@ -78,7 +77,7 @@ func (mcpFileReader) ReadFile(name string) ([]byte, error) { func (mcpFileReader) NoPasswordPrompt() bool { return true } func (mcpFileReader) ReadPassword(_ int) ([]byte, error) { - return nil, errors.New("encrypted private keys require CERTINFO_PKEY_PW under MCP") + return nil, ErrEncryptedKeyNeedsEnv } func registerExecTools(server *sdkmcp.Server) { @@ -193,7 +192,7 @@ func runRequestsExec(ctx context.Context, input runRequestsInput) (execToolOutpu valid, errs := validateRequestsConfig(yamlContent) if !valid { - return execToolOutput{}, fmt.Errorf("invalid config: %s", strings.Join(errs, "; ")) + return execToolOutput{}, &ValidationError{Messages: errs, Prefixed: true} } loaded, _, err := loadRequestsConfigYAML(yamlContent) @@ -234,13 +233,11 @@ func certinfoExec(ctx context.Context, input certinfoInput) (execToolOutput, err } if !certinfoInputProvided(input) { - return execToolOutput{}, errors.New( - "one of tlsEndpoint, certBundle, keyFile, or caBundle is required", - ) + return execToolOutput{}, ErrCertinfoInputRequired } if input.TLSInfo && input.TLSEndpoint == "" { - return execToolOutput{}, errors.New("tlsInfo requires tlsEndpoint") + return execToolOutput{}, ErrTLSInfoNeedsEndpoint } cfg, err := certinfo.New() @@ -299,7 +296,7 @@ func executeJwtinfo(ctx context.Context, input jwtinfoInput) (execToolOutput, er } if tokenData == nil || tokenData.AccessTokenRaw == "" { - return execToolOutput{}, errors.New("no JWT token data available") + return execToolOutput{}, ErrNoJWTTokenData } if err = tokenData.DecodeBase64(); err != nil { @@ -324,7 +321,7 @@ func executeJwtinfo(ctx context.Context, input jwtinfoInput) (execToolOutput, er func executeGenerateJWKS(ctx context.Context, input generateJWKSInput) (execToolOutput, error) { if strings.TrimSpace(input.PublicKeyFile) == "" { - return execToolOutput{}, errors.New("publicKeyFile is required") + return execToolOutput{}, &RequiredFieldError{Field: "publicKeyFile"} } jwksJSON, err := jwks.GenerateJWKS(ctx, input.PublicKeyFile, input.Kid) @@ -341,9 +338,9 @@ func loadConfigYAML(configYAML, configPath string) (string, error) { switch { case hasYAML && hasPath: - return "", errors.New("provide exactly one of configYaml or configPath") + return "", ErrExactlyOneConfigSource case !hasYAML && !hasPath: - return "", errors.New("configYaml or configPath is required") + return "", ErrConfigSourceRequired case hasPath: data, err := os.ReadFile(configPath) if err != nil { @@ -410,14 +407,14 @@ func loadJwtTokenData(ctx context.Context, input jwtinfoInput) (*jwtinfo.JwtToke switch { case hasFile && hasURL: - return nil, errors.New("provide exactly one of tokenFile or requestUrl") + return nil, ErrExactlyOneTokenSource case !hasFile && !hasURL: - return nil, errors.New("tokenFile or requestUrl is required") + return nil, ErrTokenSourceRequired case hasFile: return jwtinfo.ReadTokenFromFile(input.TokenFile) default: if len(input.RequestValues) == 0 { - return nil, errors.New("requestValues is required with requestUrl") + return nil, &RequiredFieldError{Field: "requestValues"} } client := &http.Client{Timeout: execToolTimeout(input.TimeoutSec)} diff --git a/internal/requests/errors.go b/internal/requests/errors.go new file mode 100644 index 0000000..eb5b4b7 --- /dev/null +++ b/internal/requests/errors.go @@ -0,0 +1,128 @@ +//nolint:revive // max-public-structs: typed domain errors for errors.Is/As +package requests + +import ( + "errors" + "fmt" +) + +// Package-level sentinels for stable requests failure conditions. +// Match them with errors.Is after wrapping; Error() strings stay human-facing. +var ( + // ErrMethodNotFound is returned when an unsupported HTTP method is specified. + ErrMethodNotFound = errors.New("HTTP method not found") + + ErrNilClient = errors.New( + "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize", + ) + ErrEmptyArg = errors.New("empty string provided as argument") + ErrServerNameIsURL = errors.New("serverName should be a hostname, not a URL") + ErrWrongTransport = errors.New("expected *http.Transport") + ErrInvalidTimeout = errors.New("timeout value must be positive") + ErrProxyProtoNeedsOverride = errors.New( + "if EnableProxyProtocolV2 is true, a TransportOverrideURL must be set", + ) + ErrProxyProtoDisabled = errors.New("proxy protocol v2 is not enabled for this request") + ErrTransportOverrideRequired = errors.New( + "SetProxyProtocolHeader failed: transportOverrideURL not set", + ) + ErrInvalidURI = errors.New("invalid uri") + ErrInvalidTransportURL = errors.New("failed to parse transport override url") +) + +// EmptyArgError is returned when a required string argument is empty. +// errors.Is(err, ErrEmptyArg) is true. +type EmptyArgError struct { + Name string +} + +// Error returns a message naming the empty argument. +func (e *EmptyArgError) Error() string { + return fmt.Sprintf("empty string provided as %s", e.Name) +} + +// Is reports whether target is ErrEmptyArg. +func (*EmptyArgError) Is(target error) bool { + return target == ErrEmptyArg +} + +// ServerNameURLError is returned when serverName looks like a URL. +// errors.Is(err, ErrServerNameIsURL) is true. +type ServerNameURLError struct { + Value string +} + +// Error returns a message including the invalid serverName. +func (e *ServerNameURLError) Error() string { + return fmt.Sprintf("serverName should be a hostname, not a URL: %s", e.Value) +} + +// Is reports whether target is ErrServerNameIsURL. +func (*ServerNameURLError) Is(target error) bool { + return target == ErrServerNameIsURL +} + +// WrongTransportError is returned when the client transport is not *http.Transport. +// errors.Is(err, ErrWrongTransport) is true. +type WrongTransportError struct { + Got any +} + +// Error returns a message including the unexpected transport type. +func (e *WrongTransportError) Error() string { + return fmt.Sprintf("expected *http.Transport, got %T", e.Got) +} + +// Is reports whether target is ErrWrongTransport. +func (*WrongTransportError) Is(target error) bool { + return target == ErrWrongTransport +} + +// InvalidTimeoutError is returned when client timeout is negative. +// errors.Is(err, ErrInvalidTimeout) is true. +type InvalidTimeoutError struct { + Value int +} + +// Error returns a message including the invalid timeout. +func (e *InvalidTimeoutError) Error() string { + return fmt.Sprintf("timeout value must be positive: %v provided", e.Value) +} + +// Is reports whether target is ErrInvalidTimeout. +func (*InvalidTimeoutError) Is(target error) bool { + return target == ErrInvalidTimeout +} + +// InvalidURIError is returned when a host URI fails Parse. +// errors.Is(err, ErrInvalidURI) is true. +type InvalidURIError struct { + URI string + Host string +} + +// Error returns a message including the invalid URI and host. +func (e *InvalidURIError) Error() string { + return fmt.Sprintf("invalid uri %s for host %s", e.URI, e.Host) +} + +// Is reports whether target is ErrInvalidURI. +func (*InvalidURIError) Is(target error) bool { + return target == ErrInvalidURI +} + +// InvalidTransportURLError is returned when a transport override URL cannot be parsed. +// errors.Is(err, ErrInvalidTransportURL) is true. +type InvalidTransportURLError struct { + URL string +} + +// Error returns a message including the invalid transport URL. +func (e *InvalidTransportURLError) Error() string { + return fmt.Sprintf("failed to parse transport override url: %s", e.URL) +} + +// Is reports whether target is ErrInvalidTransportURL. +func (*InvalidTransportURLError) Is(target error) bool { + return target == ErrInvalidTransportURL +} diff --git a/internal/requests/errors_test.go b/internal/requests/errors_test.go new file mode 100644 index 0000000..741ad19 --- /dev/null +++ b/internal/requests/errors_test.go @@ -0,0 +1,114 @@ +package requests + +import ( + "errors" + "fmt" + "testing" + + "github.com/stretchr/testify/require" +) + +// TestRequests_errorSentinels_Is checks errors.Is against requests sentinels. +func TestRequests_errorSentinels_Is(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + target error + }{ + { + name: "EmptyArgError", + err: &EmptyArgError{Name: "serverName"}, + target: ErrEmptyArg, + }, + { + name: "EmptyArgError wrapped", + err: fmt.Errorf("SetServerName error: %w", &EmptyArgError{Name: "serverName"}), + target: ErrEmptyArg, + }, + { + name: "ServerNameURLError", + err: &ServerNameURLError{Value: "https://x"}, + target: ErrServerNameIsURL, + }, + { + name: "WrongTransportError", + err: &WrongTransportError{Got: nil}, + target: ErrWrongTransport, + }, + { + name: "InvalidTimeoutError", + err: &InvalidTimeoutError{Value: -1}, + target: ErrInvalidTimeout, + }, + { + name: "InvalidURIError", + err: &InvalidURIError{URI: "/bad", Host: "h"}, + target: ErrInvalidURI, + }, + { + name: "InvalidTransportURLError", + err: &InvalidTransportURLError{URL: "bad"}, + target: ErrInvalidTransportURL, + }, + { + name: "ErrNilClient wrapped", + err: fmt.Errorf("op: %w", ErrNilClient), + target: ErrNilClient, + }, + { + name: "ErrProxyProtoNeedsOverride", + err: ErrProxyProtoNeedsOverride, + target: ErrProxyProtoNeedsOverride, + }, + { + name: "ErrProxyProtoDisabled", + err: ErrProxyProtoDisabled, + target: ErrProxyProtoDisabled, + }, + { + name: "ErrTransportOverrideRequired", + err: ErrTransportOverrideRequired, + target: ErrTransportOverrideRequired, + }, + { + name: "ErrMethodNotFound", + err: fmt.Errorf("FOO: %w", ErrMethodNotFound), + target: ErrMethodNotFound, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.ErrorIs(t, tt.err, tt.target) + }) + } +} + +// TestRequests_errorTypes_AsType checks errors.AsType for typed requests errors. +func TestRequests_errorTypes_AsType(t *testing.T) { + t.Parallel() + + empty := fmt.Errorf("wrap: %w", &EmptyArgError{Name: "transportURL"}) + gotEmpty, ok := errors.AsType[*EmptyArgError](empty) + require.True(t, ok) + require.Equal(t, "transportURL", gotEmpty.Name) + + urlErr := fmt.Errorf("wrap: %w", &ServerNameURLError{Value: "https://x"}) + gotURL, ok := errors.AsType[*ServerNameURLError](urlErr) + require.True(t, ok) + require.Equal(t, "https://x", gotURL.Value) + + timeout := fmt.Errorf("wrap: %w", &InvalidTimeoutError{Value: -1}) + gotTimeout, ok := errors.AsType[*InvalidTimeoutError](timeout) + require.True(t, ok) + require.Equal(t, -1, gotTimeout.Value) + + uri := fmt.Errorf("wrap: %w", &InvalidURIError{URI: "/x", Host: "h"}) + gotURI, ok := errors.AsType[*InvalidURIError](uri) + require.True(t, ok) + require.Equal(t, "/x", gotURI.URI) + require.Equal(t, "h", gotURI.Host) +} diff --git a/internal/requests/requests.go b/internal/requests/requests.go index aa054e3..36cefac 100644 --- a/internal/requests/requests.go +++ b/internal/requests/requests.go @@ -54,9 +54,6 @@ var defaultCurvePreferences = []tls.CurveID{ tls.CurveP521, } -// ErrMethodNotFound is returned when an unsupported HTTP method is specified. -var ErrMethodNotFound = errors.New("HTTP method not found") - var allowedHTTPMethods = map[string]string{ "GET": http.MethodGet, "HEAD": http.MethodHead, @@ -368,21 +365,20 @@ func NewRequestHTTPClient() *RequestHTTPClient { // SetServerName sets the ServerName for SNI in the TLS configuration. func (rc *RequestHTTPClient) SetServerName(serverName string) (*RequestHTTPClient, error) { if rc.client == nil { - return nil, errors.New( - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize") + return nil, ErrNilClient } if serverName == emptyString { - return nil, errors.New("serverName cannot be empty") + return nil, &EmptyArgError{Name: "serverName"} } if strings.Contains(serverName, "://") { - return nil, fmt.Errorf("serverName should be a hostname, not a URL: %s", serverName) + return nil, &ServerNameURLError{Value: serverName} } transport, ok := rc.client.Transport.(*http.Transport) if !ok { - return nil, fmt.Errorf("expected *http.Transport, got %T", rc.client.Transport) + return nil, &WrongTransportError{Got: rc.client.Transport} } tr := transport.Clone() @@ -399,8 +395,7 @@ func (rc *RequestHTTPClient) SetServerName(serverName string) (*RequestHTTPClien // SetCACertsPool sets the CA certificate pool for the HTTP transport. func (rc *RequestHTTPClient) SetCACertsPool(caPool *x509.CertPool) (*RequestHTTPClient, error) { if rc.client == nil { - return nil, errors.New( - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize") + return nil, ErrNilClient } if caPool == nil { @@ -414,7 +409,7 @@ func (rc *RequestHTTPClient) SetCACertsPool(caPool *x509.CertPool) (*RequestHTTP transport, ok := rc.client.Transport.(*http.Transport) if !ok { - return nil, fmt.Errorf("expected *http.Transport, got %T", rc.client.Transport) + return nil, &WrongTransportError{Got: rc.client.Transport} } tr := transport.Clone() @@ -431,13 +426,12 @@ func (rc *RequestHTTPClient) SetCACertsPool(caPool *x509.CertPool) (*RequestHTTP // SetInsecureSkipVerify sets whether to skip TLS certificate verification. func (rc *RequestHTTPClient) SetInsecureSkipVerify(isInsecure bool) (*RequestHTTPClient, error) { if rc.client == nil { - return nil, errors.New( - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize") + return nil, ErrNilClient } transport, ok := rc.client.Transport.(*http.Transport) if !ok { - return nil, fmt.Errorf("expected *http.Transport, got %T", rc.client.Transport) + return nil, &WrongTransportError{Got: rc.client.Transport} } tr := transport.Clone() @@ -475,13 +469,12 @@ func (rc *RequestHTTPClient) SetTransportOverride(transportURL string) (*Request } if rc.client == nil { - return nil, errors.New( - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize") + return nil, ErrNilClient } transportAddress, err := transportAddressFromURLString(transportURL) if err != nil { - return nil, fmt.Errorf("failed to parse transport override url: %s", transportURL) + return nil, &InvalidTransportURLError{URL: transportURL} } rc.transportAddress = transportAddress @@ -493,7 +486,7 @@ func (rc *RequestHTTPClient) SetTransportOverride(transportURL string) (*Request transport, ok := rc.client.Transport.(*http.Transport) if !ok { - return nil, fmt.Errorf("expected *http.Transport, got %T", rc.client.Transport) + return nil, &WrongTransportError{Got: rc.client.Transport} } tr := transport.Clone() @@ -525,12 +518,11 @@ func (rc *RequestHTTPClient) SetProxyProtocolV2(enable bool) *RequestHTTPClient // SetProxyProtocolHeader sets a custom PROXY protocol header for the client's dialer. func (rc *RequestHTTPClient) SetProxyProtocolHeader(header proxyproto.Header) (*RequestHTTPClient, error) { if rc.transportAddress == emptyString { - return nil, errors.New("SetProxyProtocolHeader failed: transportOverrideURL not set") + return nil, ErrTransportOverrideRequired } if rc.client == nil { - return nil, errors.New( - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize") + return nil, ErrNilClient } dialer := &net.Dialer{ @@ -540,7 +532,7 @@ func (rc *RequestHTTPClient) SetProxyProtocolHeader(header proxyproto.Header) (* transport, ok := rc.client.Transport.(*http.Transport) if !ok { - return nil, fmt.Errorf("expected *http.Transport, got %T", rc.client.Transport) + return nil, &WrongTransportError{Got: rc.client.Transport} } tr := transport.Clone() @@ -571,12 +563,11 @@ func (rc *RequestHTTPClient) SetProxyProtocolHeader(header proxyproto.Header) (* // SetClientTimeout sets the timeout for the HTTP client in seconds. func (rc *RequestHTTPClient) SetClientTimeout(timeout int) (*RequestHTTPClient, error) { if rc.client == nil { - return nil, errors.New( - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize") + return nil, ErrNilClient } if timeout < 0 { - return nil, fmt.Errorf("timeout value must be positive: %v provided", timeout) + return nil, &InvalidTimeoutError{Value: timeout} } t := time.Duration(timeout) * time.Second @@ -626,8 +617,7 @@ func NewHTTPClientFromRequestConfig( reqClient.SetProxyProtocolV2(r.EnableProxyProtocolV2) if r.EnableProxyProtocolV2 && r.TransportOverrideURL == emptyString { - return nil, errors.New( - "if EnableProxyProtocolV2 is true, a TransportOverrideURL must be set") + return nil, ErrProxyProtoNeedsOverride } if r.EnableProxyProtocolV2 && reqClient.transportAddress != emptyString { diff --git a/internal/requests/requests_handlers.go b/internal/requests/requests_handlers.go index a37b6c5..f433bea 100644 --- a/internal/requests/requests_handlers.go +++ b/internal/requests/requests_handlers.go @@ -5,7 +5,6 @@ import ( "context" "crypto/tls" "encoding/json" - "errors" "fmt" "io" "net" @@ -113,7 +112,7 @@ func getUrlsFromHost(h Host) ([]string, error) { for _, uri := range h.URIList { if parsed := uri.Parse(); !parsed { - return nil, fmt.Errorf("invalid uri %s for host %s", uri, h.Name) + return nil, &InvalidURIError{URI: string(uri), Host: h.Name} } s := fmt.Sprintf("%s://%s%s", httpClientDefaultScheme, h.Name, uri) @@ -128,7 +127,7 @@ func transportAddressFromURLString(transportURL string) (string, error) { var addr string if transportURL == emptyString { - return emptyString, errors.New("empty string provided as transportURL") + return emptyString, &EmptyArgError{Name: "transportURL"} } // Add HTTPS scheme if missing from transportURL @@ -154,7 +153,7 @@ func transportAddressFromURLString(transportURL string) (string, error) { // proxyProtoHeaderFromRequest generates a PROXY protocol v2 header for the given request and server name. func proxyProtoHeaderFromRequest(r RequestConfig, serverName string) (proxyproto.Header, error) { if !r.EnableProxyProtocolV2 { - return proxyproto.Header{}, errors.New("proxy protocol v2 is not enabled for this request") + return proxyproto.Header{}, ErrProxyProtoDisabled } headerSrcIP := net.ParseIP(proxyProtoDefaultSrcIPv4) diff --git a/internal/requests/requests_test.go b/internal/requests/requests_test.go index 62f7033..3b8905e 100644 --- a/internal/requests/requests_test.go +++ b/internal/requests/requests_test.go @@ -276,7 +276,7 @@ func TestNewHTTPClientFromRequestConfig_Error(t *testing.T) { desc string reqConf RequestConfig serverName string - errMsg string + target error }{ { desc: "EnableProxyProtocolV2", @@ -284,7 +284,7 @@ func TestNewHTTPClientFromRequestConfig_Error(t *testing.T) { EnableProxyProtocolV2: true, }, serverName: "localhost", - errMsg: "if EnableProxyProtocolV2 is true, a TransportOverrideURL must be set", + target: ErrProxyProtoNeedsOverride, }, { desc: "EnableProxyProtoNoServerName", @@ -293,7 +293,7 @@ func TestNewHTTPClientFromRequestConfig_Error(t *testing.T) { EnableProxyProtocolV2: true, }, serverName: emptyString, - errMsg: "SetServerName error: serverName cannot be empty", + target: ErrEmptyArg, }, } @@ -307,11 +307,7 @@ func TestNewHTTPClientFromRequestConfig_Error(t *testing.T) { tt.serverName, nil, ) - require.Error(t, err) - assert.Equal(t, - tt.errMsg, - err.Error(), - ) + require.ErrorIs(t, err, tt.target) }) } } @@ -485,17 +481,17 @@ func TestNewRequestHTTPClient_SetServerName_Error(t *testing.T) { testsError := []struct { desc string serverName string - errMsg string + target error }{ { "empty serverName", emptyString, - "serverName cannot be empty", + ErrEmptyArg, }, { "url as serverName", "https://localhost", - "serverName should be a hostname, not a URL: https://localhost", + ErrServerNameIsURL, }, } @@ -506,8 +502,7 @@ func TestNewRequestHTTPClient_SetServerName_Error(t *testing.T) { c := NewRequestHTTPClient() _, err := c.SetServerName(tt.serverName) - require.Error(t, err) - assert.Equal(t, tt.errMsg, err.Error()) + require.ErrorIs(t, err, tt.target) }) } @@ -517,11 +512,7 @@ func TestNewRequestHTTPClient_SetServerName_Error(t *testing.T) { var c RequestHTTPClient _, err := c.SetServerName("localhost") - require.Error(t, err) - assert.Equal(t, - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize", - err.Error(), - ) + require.ErrorIs(t, err, ErrNilClient) }) } @@ -561,8 +552,7 @@ func TestNewRequestHTTPClient_SetClientTimeout_Error(t *testing.T) { timeout := -1 _, err := c.SetClientTimeout(timeout) - require.Error(t, err) - assert.Equal(t, "timeout value must be positive: -1 provided", err.Error()) + require.ErrorIs(t, err, ErrInvalidTimeout) }) t.Run("Nil Timeout", func(t *testing.T) { @@ -573,12 +563,7 @@ func TestNewRequestHTTPClient_SetClientTimeout_Error(t *testing.T) { timeout := 10 _, err := c.SetClientTimeout(timeout) - require.Error(t, err) - assert.Equal( - t, - "*RequestHTTPClient.client is nil. Use NewRequestHTTPClient to initialize", - err.Error(), - ) + require.ErrorIs(t, err, ErrNilClient) }) } @@ -1241,15 +1226,14 @@ func TestProcessHTTPRequestsByHost_Errors(t *testing.T) { }, } _, err := processHTTPRequestsByHost(context.Background(), io.Discard, reqConf, nil, false) - require.Error(t, err) - require.ErrorContains(t, err, "invalid uri") + require.ErrorIs(t, err, ErrInvalidURI) }) } func TestProxyProtoHeaderFromRequest_Errors(t *testing.T) { t.Run("not enabled", func(t *testing.T) { _, err := proxyProtoHeaderFromRequest(RequestConfig{}, "localhost") - require.ErrorContains(t, err, "proxy protocol v2 is not enabled") + require.ErrorIs(t, err, ErrProxyProtoDisabled) }) // url.Parse won't fail for typical invalid URLs, but let's try a control character