Files
ssh-tarpit/ssh-tarpit.go
T
tchivert 2cff2185d5
docker / docker (push) Successful in 2m13s
build / build (push) Successful in 43s
disable WAL
2026-07-17 22:59:33 +02:00

415 lines
10 KiB
Go

package main
import (
"context"
"crypto/rand"
"database/sql"
"fmt"
"log"
"net"
"os"
"os/signal"
"strconv"
"strings"
"sync"
"syscall"
"time"
// MySQL driver registration
_ "github.com/go-sql-driver/mysql"
// Pure-Go SQLite driver
_ "modernc.org/sqlite"
"github.com/oschwald/geoip2-golang"
"github.com/pires/go-proxyproto"
)
// Config holds all configuration sourced from environment variables.
type Config struct {
Port int
Tarpit bool
TarpitDelay time.Duration
TarpitMaxDur time.Duration // 0 = unlimited
DBPath string
GeoIPPath string
MySQLUser string
MySQLPass string
MySQLHost string
MySQLPort string
MySQLDB string
Banner string
MaxConns int
ShutdownTimeout time.Duration
ProxyProtocol bool
}
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
func loadConfig() Config {
return Config{
Port: getEnvInt("SSH_TARPIT_PORT", 2222),
Tarpit: getEnvBool("SSH_TARPIT_TARPIT", false),
TarpitDelay: getEnvDuration("SSH_TARPIT_TARPIT_DELAY", 1*time.Second),
TarpitMaxDur: getEnvDuration("SSH_TARPIT_TARPIT_MAX_DURATION", 0),
DBPath: getEnv("SSH_TARPIT_DB_PATH", "./data/logs.db"),
GeoIPPath: getEnv("SSH_TARPIT_GEOIP_PATH", ""),
MySQLUser: getEnv("SSH_TARPIT_MYSQL_USER", ""),
MySQLPass: getEnv("SSH_TARPIT_MYSQL_PASS", ""),
MySQLHost: getEnv("SSH_TARPIT_MYSQL_HOST", "localhost"),
MySQLPort: getEnv("SSH_TARPIT_MYSQL_PORT", "3306"),
MySQLDB: getEnv("SSH_TARPIT_MYSQL_DB", "sshtarpit"),
Banner: getEnv("SSH_TARPIT_BANNER", "SSH-2.0-OpenSSH_9.1p1 Debian-1"),
MaxConns: getEnvInt("SSH_TARPIT_MAX_CONNS", 100),
ShutdownTimeout: getEnvDuration("SSH_TARPIT_SHUTDOWN_TIMEOUT", 30*time.Second),
ProxyProtocol: getEnvBool("SSH_TARPIT_PROXY_PROTOCOL", false),
}
}
// Env helpers
func getEnv(key, defaultVal string) string {
if v := os.Getenv(key); v != "" {
return v
}
return defaultVal
}
func getEnvBool(key string, defaultVal bool) bool {
v := os.Getenv(key)
if v == "" {
return defaultVal
}
switch strings.ToLower(v) {
case "true", "1", "yes", "on":
return true
default:
return false
}
}
func getEnvInt(key string, defaultVal int) int {
v := os.Getenv(key)
if v == "" {
return defaultVal
}
i, err := strconv.Atoi(v)
if err != nil {
log.Printf("Invalid value for %s=%q, using default %d: %v", key, v, defaultVal, err)
return defaultVal
}
return i
}
func getEnvDuration(key string, defaultVal time.Duration) time.Duration {
v := os.Getenv(key)
if v == "" {
return defaultVal
}
d, err := time.ParseDuration(v)
if err != nil {
log.Printf("Invalid duration for %s=%q, using default %s: %v", key, v, defaultVal, err)
return defaultVal
}
return d
}
// ConnectionLogger persists connection records.
type ConnectionLogger interface {
LogConnection(ip, country, city string, latitude, longitude float64, isp string) error
Close() error
}
// dbLogger implements ConnectionLogger with a *sql.DB connection pool.
type dbLogger struct {
db *sql.DB
}
func (l *dbLogger) LogConnection(ip, country, city string, latitude, longitude float64, isp string) error {
_, err := l.db.Exec(
"INSERT INTO connections (ip, country, city, latitude, longitude, isp) VALUES (?, ?, ?, ?, ?, ?)",
ip, country, city, latitude, longitude, isp,
)
return err
}
func (l *dbLogger) Close() error {
return l.db.Close()
}
func newLogger(cfg Config) (ConnectionLogger, error) {
var driverName, dsn string
if cfg.MySQLUser != "" {
driverName = "mysql"
dsn = fmt.Sprintf("%s:%s@tcp(%s:%s)/%s?parseTime=true",
cfg.MySQLUser, cfg.MySQLPass, cfg.MySQLHost, cfg.MySQLPort, cfg.MySQLDB)
} else {
driverName = "sqlite"
dsn = cfg.DBPath
}
db, err := sql.Open(driverName, dsn)
if err != nil {
return nil, fmt.Errorf("sql.Open: %w", err)
}
// Connection pool tuning
db.SetMaxOpenConns(25)
db.SetMaxIdleConns(5)
db.SetConnMaxLifetime(5 * time.Minute)
// SQLite-specific pragma
if driverName == "sqlite" {
if _, err := db.Exec("PRAGMA busy_timeout=5000"); err != nil {
log.Printf("Warning: could not set busy_timeout: %v", err)
}
}
// Create table once at startup
createSQL := `CREATE TABLE IF NOT EXISTS connections (
ip TEXT,
country TEXT,
city TEXT,
latitude REAL,
longitude REAL,
isp TEXT,
timestamp TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)`
if _, err := db.Exec(createSQL); err != nil {
db.Close()
return nil, fmt.Errorf("create table: %w", err)
}
log.Printf("Database ready (driver=%s)", driverName)
return &dbLogger{db: db}, nil
}
// GeoIP helpers
type geoipReaders struct {
city *geoip2.Reader
asn *geoip2.Reader
}
func newGeoIPReaders(geoipPath string) (*geoipReaders, error) {
city, err := geoip2.Open(geoipPath + "/GeoLite2-City.mmdb")
if err != nil {
return nil, fmt.Errorf("open City DB: %w", err)
}
asn, err := geoip2.Open(geoipPath + "/GeoLite2-ASN.mmdb")
if err != nil {
city.Close()
return nil, fmt.Errorf("open ASN DB: %w", err)
}
return &geoipReaders{city: city, asn: asn}, nil
}
func (g *geoipReaders) Close() {
if g.city != nil {
g.city.Close()
}
if g.asn != nil {
g.asn.Close()
}
}
func (g *geoipReaders) lookup(ip string) (country, city string, latitude, longitude float64, isp string, err error) {
parsedIP := net.ParseIP(ip)
if parsedIP == nil {
return "", "", 0, 0, "", fmt.Errorf("invalid IP: %s", ip)
}
record, err := g.city.City(parsedIP)
if err != nil {
return "", "", 0, 0, "", fmt.Errorf("city lookup: %w", err)
}
latitude, longitude = record.Location.Latitude, record.Location.Longitude
country = record.Country.Names["en"]
city = record.City.Names["en"]
asnRecord, err := g.asn.ASN(parsedIP)
if err != nil {
return "", "", 0, 0, "", fmt.Errorf("ASN lookup: %w", err)
}
isp = asnRecord.AutonomousSystemOrganization
return country, city, latitude, longitude, isp, nil
}
// Main
func main() {
log.SetFlags(log.LstdFlags | log.Lmsgprefix)
log.SetPrefix("[ssh-tarpit] ")
cfg := loadConfig()
// Initialise database connection pool once
logger, err := newLogger(cfg)
if err != nil {
log.Fatalf("Failed to initialise database: %v", err)
}
defer logger.Close()
// Initialise GeoIP readers once (if configured)
var geoReaders *geoipReaders
if cfg.GeoIPPath != "" {
geoReaders, err = newGeoIPReaders(cfg.GeoIPPath)
if err != nil {
log.Printf("Warning: GeoIP initialisation failed, falling back to IP-only logging: %v", err)
geoReaders = nil
} else {
defer geoReaders.Close()
}
}
addr := fmt.Sprintf(":%d", cfg.Port)
rawLn, err := net.Listen("tcp", addr)
if err != nil {
log.Fatalf("Failed to listen on %s: %v", addr, err)
}
var ln net.Listener
if cfg.ProxyProtocol {
ln = &proxyproto.Listener{
Listener: rawLn,
ReadHeaderTimeout: 5 * time.Second,
}
log.Printf("Listening on %s (PROXY protocol enabled)", addr)
} else {
ln = rawLn
log.Printf("Listening on %s", addr)
}
// Signal handling for graceful shutdown
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
sigCh := make(chan os.Signal, 1)
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
go func() {
sig := <-sigCh
log.Printf("Received %v, shutting down...", sig)
cancel()
ln.Close() // unblock Accept
}()
// Semaphore to limit concurrent connections
sem := make(chan struct{}, cfg.MaxConns)
var wg sync.WaitGroup
acceptLoop:
for {
conn, err := ln.Accept()
if err != nil {
select {
case <-ctx.Done():
break acceptLoop
default:
log.Printf("Accept error: %v", err)
continue
}
}
select {
case sem <- struct{}{}:
wg.Add(1)
go func() {
defer wg.Done()
defer func() { <-sem }()
handleConnection(ctx, conn, logger, geoReaders, cfg)
}()
default:
log.Printf("Connection limit reached (%d), rejecting %s", cfg.MaxConns, conn.RemoteAddr())
conn.Close()
}
}
// Wait for active connections to drain
log.Println("Waiting for active connections to finish...")
done := make(chan struct{})
go func() {
wg.Wait()
close(done)
}()
select {
case <-done:
log.Println("All connections finished cleanly")
case <-time.After(cfg.ShutdownTimeout):
log.Printf("Timed out after %s waiting for connections", cfg.ShutdownTimeout)
}
}
// Connection handler
func handleConnection(ctx context.Context, conn net.Conn, logger ConnectionLogger, geoReaders *geoipReaders, cfg Config) {
defer conn.Close()
// Parse remote IP properly
remoteAddr := conn.RemoteAddr().String()
ip, _, err := net.SplitHostPort(remoteAddr)
if err != nil {
log.Printf("Failed to parse remote address %q: %v", remoteAddr, err)
return
}
// Set a deadline for GeoIP / DB operations
conn.SetDeadline(time.Now().Add(30 * time.Second))
// Look up GeoIP data if available, otherwise just log the IP
country, city, isp := "", "", ""
var latitude, longitude float64
if geoReaders != nil {
country, city, latitude, longitude, isp, err = geoReaders.lookup(ip)
if err != nil {
log.Printf("GeoIP lookup failed for %s: %v", ip, err)
// Fall back to IP-only logging
country, city, latitude, longitude, isp = "", "", 0, 0, ""
}
}
if err := logger.LogConnection(ip, country, city, latitude, longitude, isp); err != nil {
log.Printf("Failed to log connection from %s: %v", ip, err)
}
if geoReaders != nil {
log.Printf("Connection from %s (%s, %s, %s)", ip, country, city, isp)
} else {
log.Printf("Connection from %s", ip)
}
// Reset deadline for the remainder of the connection
conn.SetDeadline(time.Time{})
if !cfg.Tarpit {
// Send banner and hold briefly
conn.Write([]byte(cfg.Banner + "\r\n"))
select {
case <-time.After(2 * time.Second):
case <-ctx.Done():
}
return
}
// Tarpit mode: drip random bytes periodically
tarpitCtx, tarpitCancel := context.WithCancel(ctx)
defer tarpitCancel()
if cfg.TarpitMaxDur > 0 {
tarpitCtx, tarpitCancel = context.WithTimeout(ctx, cfg.TarpitMaxDur)
defer tarpitCancel()
}
b := make([]byte, 1)
for {
select {
case <-tarpitCtx.Done():
return
default:
}
if _, err := rand.Read(b); err != nil {
log.Printf("rand.Read error: %v", err)
return
}
if _, err := conn.Write([]byte{charset[int(b[0])%len(charset)]}); err != nil {
return
}
time.Sleep(cfg.TarpitDelay)
}
}