Files
2026-08-27 12:00:08 -06:00

248 lines
6.9 KiB
Go

package main
import (
"context"
"crypto/tls"
"flag"
"fmt"
"log"
"net"
"net/http"
"os"
"os/signal"
"strings"
"syscall"
"time"
)
var (
version = "1.0.0"
)
func parseTLSVersion(v string) (uint16, error) {
switch strings.TrimSpace(v) {
case "1.0", "tls1.0", "TLS1.0":
return tls.VersionTLS10, nil
case "1.1", "tls1.1", "TLS1.1":
return tls.VersionTLS11, nil
case "1.2", "tls1.2", "TLS1.2":
return tls.VersionTLS12, nil
case "1.3", "tls1.3", "TLS1.3":
return tls.VersionTLS13, nil
default:
return 0, fmt.Errorf("unknown TLS version %q (valid options: 1.0, 1.1, 1.2, 1.3)", v)
}
}
func printUsage() {
fmt.Fprintf(os.Stderr, `smallprox v%s - Miniature TLS 1.2 Reverse Proxy
Usage:
smallprox -backend <url> [options]
Options:
-backend, -b string
Target backend URL (required, e.g. https://tls13-service.example.com)
-listen, -l string
Address and port to bind to (default ":8443")
-tls-min string
Minimum incoming TLS version: 1.0, 1.1, 1.2, 1.3 (default "1.2")
-tls-max string
Maximum incoming TLS version: 1.0, 1.1, 1.2, 1.3 (default "1.2")
-cert string
Path to TLS certificate PEM file (optional; self-signed generated if omitted)
-key string
Path to TLS private key PEM file (optional; self-signed generated if omitted)
-http
Listen in plain HTTP mode instead of HTTPS (default false)
-preserve-host
Preserve incoming Host header instead of rewriting to backend host (default false)
-version, -v
Print version information and exit
-help, -h
Show this help message
Examples:
# Basic reverse proxy to a TLS 1.3 backend (PowerShell -> TLS 1.2 proxy:8443 -> TLS 1.3 backend)
smallprox -backend https://api.example.com
# Custom listen port and specific TLS range
smallprox -listen :9443 -backend https://api.example.com -tls-min 1.2 -tls-max 1.3
# Using custom certificates
smallprox -backend https://api.example.com -cert server.crt -key server.key
# PowerShell usage example:
# [Net.ServicePointManager]::SecurityProtocol = [Net.SecurityProtocolType]::Tls12
# [System.Net.ServicePointManager]::ServerCertificateValidationCallback = {$true}
# Invoke-RestMethod -Uri https://localhost:8443/api/endpoint
`, version)
}
func main() {
var (
backendFlag string
backendShort string
listenAddr string
listenShort string
certFile string
keyFile string
tlsMinStr string
tlsMaxStr string
plainHTTP bool
preserveHost bool
showVersion bool
showVersionShort bool
showHelp bool
showHelpShort bool
)
flag.StringVar(&backendFlag, "backend", "", "Target backend URL")
flag.StringVar(&backendShort, "b", "", "Target backend URL (short)")
flag.StringVar(&listenAddr, "listen", ":8443", "Listen address")
flag.StringVar(&listenShort, "l", ":8443", "Listen address (short)")
flag.StringVar(&certFile, "cert", "", "Path to TLS cert file")
flag.StringVar(&keyFile, "key", "", "Path to TLS key file")
flag.StringVar(&tlsMinStr, "tls-min", "1.2", "Minimum incoming TLS version")
flag.StringVar(&tlsMaxStr, "tls-max", "1.2", "Maximum incoming TLS version")
flag.BoolVar(&plainHTTP, "http", false, "Listen in plain HTTP mode")
flag.BoolVar(&preserveHost, "preserve-host", false, "Preserve incoming Host header")
flag.BoolVar(&showVersion, "version", false, "Print version")
flag.BoolVar(&showVersionShort, "v", false, "Print version (short)")
flag.BoolVar(&showHelp, "help", false, "Show help")
flag.BoolVar(&showHelpShort, "h", false, "Show help (short)")
flag.Usage = printUsage
flag.Parse()
if showHelp || showHelpShort {
printUsage()
os.Exit(0)
}
if showVersion || showVersionShort {
fmt.Printf("smallprox version %s\n", version)
os.Exit(0)
}
// Resolve flags with short versions
rawBackend := backendFlag
if rawBackend == "" {
rawBackend = backendShort
}
// Also check positional argument if no flag given
if rawBackend == "" && flag.NArg() > 0 {
rawBackend = flag.Arg(0)
}
if rawBackend == "" {
fmt.Fprintln(os.Stderr, "[ERROR] Backend URL is required.")
printUsage()
os.Exit(1)
}
targetURL, err := ParseTargetURL(rawBackend)
if err != nil {
log.Fatalf("[FATAL] %v", err)
}
// Determine listen address
bindAddr := listenAddr
if bindAddr == ":8443" && listenShort != ":8443" {
bindAddr = listenShort
}
// Parse TLS versions
minTLS, err := parseTLSVersion(tlsMinStr)
if err != nil {
log.Fatalf("[FATAL] Invalid -tls-min: %v", err)
}
maxTLS, err := parseTLSVersion(tlsMaxStr)
if err != nil {
log.Fatalf("[FATAL] Invalid -tls-max: %v", err)
}
if minTLS > maxTLS {
log.Fatalf("[FATAL] -tls-min (%s) cannot be greater than -tls-max (%s)", tlsMinStr, tlsMaxStr)
}
proxyHandler := NewProxyHandler(ProxyConfig{
TargetURL: targetURL,
PreserveHost: preserveHost,
})
server := &http.Server{
Addr: bindAddr,
Handler: proxyHandler,
ReadTimeout: 60 * time.Second,
WriteTimeout: 60 * time.Second,
IdleTimeout: 120 * time.Second,
}
protocol := "HTTPS"
if plainHTTP {
protocol = "HTTP"
} else {
tlsConfig := &tls.Config{
MinVersion: minTLS,
MaxVersion: maxTLS,
}
if certFile != "" && keyFile != "" {
cert, err := tls.LoadX509KeyPair(certFile, keyFile)
if err != nil {
log.Fatalf("[FATAL] Failed to load certificate/key: %v", err)
}
tlsConfig.Certificates = []tls.Certificate{cert}
log.Printf("[INFO] Loaded TLS certificate from %s and %s", certFile, keyFile)
} else {
hostPart, _, _ := net.SplitHostPort(bindAddr)
extraHosts := []string{}
if hostPart != "" {
extraHosts = append(extraHosts, hostPart)
}
cert, err := generateSelfSignedCert(extraHosts)
if err != nil {
log.Fatalf("[FATAL] Failed to generate self-signed certificate: %v", err)
}
tlsConfig.Certificates = []tls.Certificate{cert}
log.Printf("[INFO] Generated ephemeral self-signed TLS certificate")
}
server.TLSConfig = tlsConfig
}
// Channel to listen for shutdown signals
stopChan := make(chan os.Signal, 1)
signal.Notify(stopChan, os.Interrupt, syscall.SIGTERM)
go func() {
log.Printf("[INFO] smallprox v%s starting...", version)
log.Printf("[INFO] Listening on %s://%s (incoming TLS: %s - %s)", protocol, bindAddr, tlsVersionName(minTLS), tlsVersionName(maxTLS))
log.Printf("[INFO] Proxying to %s (backend TLS verification disabled)", targetURL.String())
var err error
if plainHTTP {
err = server.ListenAndServe()
} else {
// TLSConfig already contains the certificate
err = server.ListenAndServeTLS("", "")
}
if err != nil && err != http.ErrServerClosed {
log.Fatalf("[FATAL] Server listener error: %v", err)
}
}()
<-stopChan
log.Printf("[INFO] Shutting down smallprox gracefully...")
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := server.Shutdown(shutdownCtx); err != nil {
log.Printf("[ERROR] Server shutdown error: %v", err)
}
log.Printf("[INFO] smallprox stopped.")
}