404 lines
10 KiB
Go
404 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"
|
|
)
|
|
|
|
// 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
|
|
}
|
|
|
|
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", "./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),
|
|
}
|
|
}
|
|
|
|
// 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 pragmas for better concurrent write performance
|
|
if driverName == "sqlite" {
|
|
if _, err := db.Exec("PRAGMA journal_mode=WAL"); err != nil {
|
|
log.Printf("Warning: could not set WAL mode: %v", err)
|
|
}
|
|
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)
|
|
ln, err := net.Listen("tcp", addr)
|
|
if err != nil {
|
|
log.Fatalf("Failed to listen on %s: %v", addr, err)
|
|
}
|
|
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)
|
|
}
|
|
}
|