package main import ( "context" "crypto/tls" "crypto/x509" "fmt" "io" "log" "net" "net/http" "os" "strings" "time" ) const allowedHost = "api.demo.local" func main() { secret := os.Getenv("DEMO_API_KEY") if secret == "" { log.Fatal("DEMO_API_KEY is required") } transport := newUpstreamTransport() go func() { health := http.NewServeMux() health.HandleFunc("/healthz", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) }) log.Fatal(http.ListenAndServe("127.0.0.1:18080", health)) }() cert, err := tls.LoadX509KeyPair("/certs/mitm.crt", "/certs/mitm.key") if err != nil { log.Fatalf("load MITM certificate: %v", err) } tlsConfig := &tls.Config{ Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12, NextProtos: []string{"http/1.1"}, } listener, err := net.Listen("tcp", ":15001") if err != nil { log.Fatal(err) } server := &http.Server{ Handler: injectHandler(secret, transport), ReadHeaderTimeout: 5 * time.Second, } log.Printf("MITM proxy listening on :15001 for %s", allowedHost) log.Fatal(server.Serve(tls.NewListener(listener, tlsConfig))) } func newUpstreamTransport() *http.Transport { pem, err := os.ReadFile("/certs/upstream-ca.crt") if err != nil { log.Fatalf("read upstream CA: %v", err) } roots := x509.NewCertPool() if !roots.AppendCertsFromPEM(pem) { log.Fatal("parse upstream CA") } return &http.Transport{ Proxy: nil, TLSClientConfig: &tls.Config{ RootCAs: roots, ServerName: allowedHost, MinVersion: tls.VersionTLS12, }, ForceAttemptHTTP2: false, } } func injectHandler(secret string, transport http.RoundTripper) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { sni := strings.ToLower(strings.TrimSuffix(r.TLS.ServerName, ".")) host := canonicalHost(r.Host) // Bind all destination identities together before touching the secret. if sni == "" || host == "" || sni != host { log.Printf("drop injection: SNI/Host mismatch sni=%q host=%q", sni, host) http.Error(w, "SNI and Host must match", http.StatusMisdirectedRequest) return } if sni != allowedHost { http.Error(w, "destination is not allowed", http.StatusForbidden) return } out := r.Clone(context.Background()) out.RequestURI = "" out.URL.Scheme = "https" out.URL.Host = allowedHost + ":443" out.Host = allowedHost removeHopByHopHeaders(out.Header) out.Header.Set("Authorization", "Bearer "+secret) log.Printf("inject Authorization for sni=%s path=%s", sni, out.URL.Path) resp, err := transport.RoundTrip(out) if err != nil { http.Error(w, fmt.Sprintf("upstream failed: %v", err), http.StatusBadGateway) return } defer resp.Body.Close() for name, values := range resp.Header { for _, value := range values { w.Header().Add(name, value) } } w.WriteHeader(resp.StatusCode) _, _ = io.Copy(w, resp.Body) }) } func canonicalHost(raw string) string { host := raw if parsed, _, err := net.SplitHostPort(raw); err == nil { host = parsed } return strings.ToLower(strings.TrimSuffix(host, ".")) } func removeHopByHopHeaders(header http.Header) { for _, name := range []string{ "Connection", "Proxy-Connection", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization", "Te", "Trailer", "Transfer-Encoding", "Upgrade", } { header.Del(name) } }