Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
121 changes: 98 additions & 23 deletions cmd/pgproxy/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,15 @@ import (
"crypto/tls"
"crypto/x509"
"database/sql"
_ "github.com/microsoft/go-mssqldb"
"log/slog"
"net/http"
"net/url"
"os"
"os/signal"
"pgproxy/internal/pipeline"
"pgproxy/internal/sqlnorm"
"pgproxy/internal/wire/tds"
"strings"
"syscall"
"time"
Expand Down Expand Up @@ -139,22 +143,22 @@ func main() {
shuntSSLMode := parseSSLMode(cfg.ShuntNodeURL)

proxyCfg := proxy.Config{
ListenAddr: cfg.ListenAddr,
BackendAddr: parseBackendAddr(cfg.BackendURL),
BackendSSLMode: backendSSLMode,
BackendUser: backendUser,
BackendPassword: backendPassword,
BackendDatabase: backendDatabase,
ShuntAddr: parseBackendAddr(cfg.ShuntNodeURL),
ShuntSSLMode: shuntSSLMode,
ShuntUser: shuntUser,
ShuntPassword: shuntPassword,
ShuntDatabase: shuntDatabase,
Mode: cfg.Mode,
Analyzer: analyzer,
Cache: approvalCache,
Reporter: rep,
MetricsReporter: metricsRep,
ListenAddr: cfg.ListenAddr,
BackendAddr: parseBackendAddr(cfg.BackendURL),
BackendSSLMode: backendSSLMode,
BackendUser: backendUser,
BackendPassword: backendPassword,
BackendDatabase: backendDatabase,
ShuntAddr: parseBackendAddr(cfg.ShuntNodeURL),
ShuntSSLMode: shuntSSLMode,
ShuntUser: shuntUser,
ShuntPassword: shuntPassword,
ShuntDatabase: shuntDatabase,
Mode: cfg.Mode,
Analyzer: analyzer,
Cache: approvalCache,
Reporter: rep,
MetricsReporter: metricsRep,
TLSConfig: tlsConfig,
BackendTLSConfig: backendTLSConfig,
RequireClientTLS: cfg.RequireClientTLS,
Expand Down Expand Up @@ -234,6 +238,77 @@ func main() {
errChan <- p.ListenAndServe()
}()

// Experimental: TDS passthrough frontend (multi-protocol milestone 2,
// see docs/multiprotocol-design.md). Enabled only when both env vars
// are set, e.g. TDS_LISTEN_ADDR=:1434 TDS_BACKEND_ADDR=localhost:1433.
if tdsListen := os.Getenv("TDS_LISTEN_ADDR"); tdsListen != "" {
tdsBackend := os.Getenv("TDS_BACKEND_ADDR")
if tdsBackend == "" {
slog.Error("TDS_LISTEN_ADDR set but TDS_BACKEND_ADDR missing; TDS frontend disabled")
} else {
tdsSrv := &tds.Server{ListenAddr: tdsListen, BackendAddr: tdsBackend}
if certPath := os.Getenv("TDS_TLS_CERT_PATH"); certPath != "" {
cert, cerr := tls.LoadX509KeyPair(certPath, os.Getenv("TDS_TLS_KEY_PATH"))
if cerr != nil {
slog.Error("TDS TLS disabled: certificate load failed", "error", cerr)
} else {
tdsSrv.TLSConfig = tds.NewTLSConfig(cert)
slog.Info("TDS strict TLS enabled (TDS 8.0)", "cert", certPath)
}
} else if os.Getenv("TDS_TLS_SELF_SIGNED") == "true" {
tlsCfg, terr := tds.SelfSignedTLSConfig("localhost", "127.0.0.1")
if terr != nil {
slog.Error("TDS TLS disabled: self-signed generation failed", "error", terr)
} else {
tdsSrv.TLSConfig = tlsCfg
slog.Warn("TDS strict TLS enabled with SELF-SIGNED certificate (dev only)")
}
}
if cfg.AnalysisEnabled {
tdsMode := cfg.Mode
if tdsMode == config.ModeShunt {
slog.Warn("shunt mode not yet supported on the TDS frontend; using blocking")
tdsMode = config.ModeBlocking
}
tdsLLM, err := agent.NewBedrockClient(ctx)
if err != nil {
slog.Error("TDS analysis disabled: Bedrock client init failed", "error", err)
} else {
var tsqlTools *agent.TSQLTools
if dsn := os.Getenv("TDS_BACKEND_DSN"); dsn != "" {
tdsDB, derr := sql.Open("sqlserver", dsn)
if derr != nil {
slog.Error("TDS tools disabled: backend DSN open failed", "error", derr)
} else {
tsqlTools = agent.NewTSQLTools(tdsDB)
slog.Info("TDS agent tools enabled (DMV + SHOWPLAN)")
}
} else {
slog.Warn("TDS_BACKEND_DSN not set; tsql agent runs without engine tools")
}
tsqlPipeline := pipeline.New(pipeline.Config{
Analyzer: agent.NewTSQLAgent(tdsLLM, tsqlTools),
Cache: approvalCache,
Mode: tdsMode,
SkipAnalysis: tds.SkipAnalysis,
Fingerprint: sqlnorm.Fingerprint,
AnalysisTimeout: 60 * time.Second,
CacheTTL: 24 * time.Hour,
})
tdsSrv.Decide = func(ctx context.Context, sql string) pipeline.Verdict {
return tsqlPipeline.Decide(pipeline.Statement{Database: "tsql", SQL: sql})
}
slog.Info("TDS analysis enabled", "mode", tdsMode, "persona", "tsql")
}
}
go func() {
if err := tdsSrv.ListenAndServe(ctx); err != nil {
errChan <- err
}
}()
}
}

// Wait for signal or error
select {
case sig := <-sigChan:
Expand Down Expand Up @@ -380,9 +455,9 @@ func initMetricsReporter(ctx context.Context, cfg *config.Config) reporter.Metri
func initCache(ctx context.Context, cfg *config.Config) (cache.ApprovalCache, masking.MetadataCache, masking.InferredMaskingCache, redis.UniversalClient) {
if cfg.ValkeyURL != "" {
vc, err := cache.NewValkeyApprovalCache(ctx, cache.ValkeyConfig{
URL: cfg.ValkeyURL,
KeyPrefix: "pgproxy:",
ClusterMode: cfg.ValkeyClusterMode,
URL: cfg.ValkeyURL,
KeyPrefix: "pgproxy:",
ClusterMode: cfg.ValkeyClusterMode,
IAMAuth: cfg.ValkeyIAMAuth,
IAMUsername: cfg.ValkeyIAMUsername,
CacheName: cfg.ValkeyCacheName,
Expand Down Expand Up @@ -510,10 +585,10 @@ func initMasking(ctx context.Context, cfg *config.Config, db *sql.DB, metadataCa
svc, err := masking.NewService(ctx, db, masking.ServiceConfig{
Enabled: true,
ConfigPath: cfg.MaskingConfigPath,
Cache: nil, // Query rewrite cache: nil triggers default in-memory cache (caches by query fingerprint only, not database-aware)
MetadataCache: metadataCache, // Cache table column info to avoid slow information_schema queries
InferenceCache: inferenceCache, // Cache LLM-inferred masking decisions
DatabaseName: dbName, // Masking operates on single database extracted from BackendURL
Cache: nil, // Query rewrite cache: nil triggers default in-memory cache (caches by query fingerprint only, not database-aware)
MetadataCache: metadataCache, // Cache table column info to avoid slow information_schema queries
InferenceCache: inferenceCache, // Cache LLM-inferred masking decisions
DatabaseName: dbName, // Masking operates on single database extracted from BackendURL
MetadataCacheTTL: cfg.MetadataCacheTTL,
}, cfg)
if err != nil {
Expand Down
Loading
Loading