diff options
Diffstat (limited to 'internal/clientcert')
| -rw-r--r-- | internal/clientcert/clientcert.go | 62 | ||||
| -rw-r--r-- | internal/clientcert/clientcert_test.go | 111 |
2 files changed, 173 insertions, 0 deletions
diff --git a/internal/clientcert/clientcert.go b/internal/clientcert/clientcert.go new file mode 100644 index 0000000..4bba835 --- /dev/null +++ b/internal/clientcert/clientcert.go @@ -0,0 +1,62 @@ +// Package clientcert manages client (mutual-TLS) certificates: which +// certificate mitmux presents to an upstream server that requires one, +// selected by matching the request's hostname the same way scope rules +// do (see internal/scope) - substring match by default, or a regex - so +// the "which rule applies to this host" mental model stays identical +// throughout the tool. +package clientcert + +import ( + "crypto/tls" + "fmt" + "regexp" + "strings" +) + +// Cert is one client certificate, scoped to hosts matching Pattern. +type Cert struct { + ID int64 + Enabled bool + Name string + Pattern string + IsRegex bool + CertPEM []byte + KeyPEM []byte +} + +func (c Cert) matches(host string) bool { + if c.IsRegex { + re, err := regexp.Compile(c.Pattern) + if err != nil { + return false + } + return re.MatchString(host) + } + return strings.Contains(strings.ToLower(host), strings.ToLower(c.Pattern)) +} + +// FindFor returns the first enabled cert whose pattern matches host, or +// nil if none applies - mitmux then just doesn't present a client +// certificate for that connection, same as if mutual TLS weren't +// configured at all. First-match-wins on ID order, same convention as +// match-and-replace rules' Position ordering, minus the extra field: +// client certs are keyed by host, not layered edits, so insertion order +// is a reasonable enough tiebreaker without adding one. +func FindFor(certs []Cert, host string) *Cert { + for i := range certs { + if certs[i].Enabled && certs[i].matches(host) { + return &certs[i] + } + } + return nil +} + +// TLSCertificate parses c's PEM-encoded cert/key pair into the form +// crypto/tls needs to present it during a handshake. +func (c Cert) TLSCertificate() (tls.Certificate, error) { + cert, err := tls.X509KeyPair(c.CertPEM, c.KeyPEM) + if err != nil { + return tls.Certificate{}, fmt.Errorf("parse client certificate %q: %w", c.Name, err) + } + return cert, nil +} diff --git a/internal/clientcert/clientcert_test.go b/internal/clientcert/clientcert_test.go new file mode 100644 index 0000000..712562b --- /dev/null +++ b/internal/clientcert/clientcert_test.go @@ -0,0 +1,111 @@ +package clientcert + +import ( + "bytes" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "testing" + "time" +) + +// selfSignedPEM generates a throwaway self-signed cert/key pair entirely +// in memory (crypto/tls + crypto/x509 stdlib only) rather than shelling +// out to openssl or checking in a fixture - keeps the test hermetic and +// portable. +func selfSignedPEM(t *testing.T) (certPEM, keyPEM []byte) { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("generate key: %v", err) + } + tmpl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "test"}, + NotBefore: time.Now(), + NotAfter: time.Now().Add(time.Hour), + } + der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key) + if err != nil { + t.Fatalf("create certificate: %v", err) + } + keyDER, err := x509.MarshalECPrivateKey(key) + if err != nil { + t.Fatalf("marshal key: %v", err) + } + var certBuf, keyBuf bytes.Buffer + pem.Encode(&certBuf, &pem.Block{Type: "CERTIFICATE", Bytes: der}) + pem.Encode(&keyBuf, &pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + return certBuf.Bytes(), keyBuf.Bytes() +} + +func TestFindForExactMatch(t *testing.T) { + certs := []Cert{{ID: 1, Enabled: true, Pattern: "internal.example.com"}} + got := FindFor(certs, "internal.example.com") + if got == nil || got.ID != 1 { + t.Errorf("got %+v, want match on id 1", got) + } +} + +func TestFindForSubstringMatch(t *testing.T) { + certs := []Cert{{ID: 1, Enabled: true, Pattern: "example.com"}} + if FindFor(certs, "api.example.com") == nil { + t.Error("expected subdomain to match by substring, same convention as scope.Rule") + } +} + +func TestFindForRegex(t *testing.T) { + certs := []Cert{{ID: 1, Enabled: true, Pattern: `^api-\d+\.example\.com$`, IsRegex: true}} + if FindFor(certs, "api-42.example.com") == nil { + t.Error("expected regex pattern to match") + } + if FindFor(certs, "api-x.example.com") != nil { + t.Error("expected regex pattern not to match a non-numeric suffix") + } +} + +func TestFindForSkipsDisabled(t *testing.T) { + certs := []Cert{{ID: 1, Enabled: false, Pattern: "example.com"}} + if FindFor(certs, "example.com") != nil { + t.Error("expected a disabled cert not to match") + } +} + +func TestFindForNoMatch(t *testing.T) { + certs := []Cert{{ID: 1, Enabled: true, Pattern: "example.com"}} + if FindFor(certs, "other.org") != nil { + t.Error("expected no match for an unrelated host") + } +} + +func TestFindForFirstMatchWins(t *testing.T) { + certs := []Cert{ + {ID: 1, Enabled: true, Pattern: "example.com"}, + {ID: 2, Enabled: true, Pattern: "example.com"}, + } + got := FindFor(certs, "example.com") + if got == nil || got.ID != 1 { + t.Errorf("got %+v, want the first matching cert (id 1)", got) + } +} + +func TestTLSCertificateValidPair(t *testing.T) { + certPEM, keyPEM := selfSignedPEM(t) + c := Cert{Name: "test", CertPEM: certPEM, KeyPEM: keyPEM} + if _, err := c.TLSCertificate(); err != nil { + t.Errorf("unexpected error: %v", err) + } +} + +func TestTLSCertificateMismatchedPairErrors(t *testing.T) { + certPEM, _ := selfSignedPEM(t) + _, otherKeyPEM := selfSignedPEM(t) + c := Cert{Name: "test", CertPEM: certPEM, KeyPEM: otherKeyPEM} + if _, err := c.TLSCertificate(); err == nil { + t.Error("expected an error pairing a cert with a key that doesn't match it") + } +} |