diff --git a/egress/egress.go b/egress/egress.go new file mode 100644 index 000000000..2c2d0b207 --- /dev/null +++ b/egress/egress.go @@ -0,0 +1,103 @@ +// Package egress carries the dial hook a distribution installs to police the +// connections GoModel opens directly. +// +// HTTP clients are steered by the standard proxy variables, so a +// distribution that wants to see every outbound HTTP request only has to set +// HTTP_PROXY. The database, cache, and vector store clients speak their own +// protocols over raw TCP and read no such variable: without a hook they +// connect wherever their URL points, whatever policy the distribution +// believes it is enforcing. This package is that hook. +// +// It is process-wide on purpose, matching the proxy variables it +// complements: the clients live behind several layers of construction, and a +// guarantee an operator is told covers the process cannot depend on a +// parameter each of them remembers to pass along. +package egress + +import ( + "context" + "errors" + "net" + "sync/atomic" + "time" +) + +// DialFunc opens one connection. It matches net.Dialer.DialContext, which is +// what the database and cache clients expect. +type DialFunc func(ctx context.Context, network, address string) (net.Conn, error) + +var installed atomic.Pointer[DialFunc] + +// defaultDial is what core uses until a distribution installs a hook. +var defaultDial DialFunc = (&net.Dialer{Timeout: 30 * time.Second, KeepAlive: 30 * time.Second}).DialContext + +// Hook is an installed policy, handed to whoever installed it. Removing the +// policy needs this handle, so a package that merely imports egress cannot +// take the guarantee away from the clients running under it. +type Hook struct{ dial *DialFunc } + +// Install routes every direct connection core opens from now on through +// dial, and returns the handle that removes it again. It belongs to +// whatever composes the process - a distribution's startup path - and there +// is one: installing over an existing hook is an error rather than a silent +// replacement. +// +// Connections already open are unaffected: a hook installed at startup, as +// GoModel Pro's air-gapped mode does before the application is built, sees +// every connection the gateway makes. +func Install(dial DialFunc) (*Hook, error) { + if dial == nil { + return nil, errors.New("egress: dial hook is required") + } + if !installed.CompareAndSwap(nil, &dial) { + return nil, errors.New("egress: a dial hook is already installed") + } + return &Hook{dial: &dial}, nil +} + +// Uninstall removes this hook, so core's clients dial directly again. A hook +// that is no longer the installed one leaves it alone, and calling twice is +// harmless. +func (h *Hook) Uninstall() { + if h == nil { + return + } + installed.CompareAndSwap(h.dial, nil) +} + +// Installed reports whether a distribution has installed a hook. +func Installed() bool { return installed.Load() != nil } + +// DialContext opens a connection through the installed hook, or with the +// default dialer when there is none. Clients that take a dialer are given +// this function; clients that take an interface are given Dialer. +func DialContext(ctx context.Context, network, address string) (net.Conn, error) { + if dial := installed.Load(); dial != nil { + return (*dial)(ctx, network, address) + } + return defaultDial(ctx, network, address) +} + +// Lookup resolves host for a client that resolves names itself before it +// dials. With a hook installed the name is passed through untouched, so the +// hook decides what it may resolve to and dials the address that passed; +// without one the client's own resolver answers, keeping its behaviour - +// multiple addresses and their fallbacks included - exactly as it was. +// +// It is consulted per connection, never snapshotted at client construction: +// a client built before the hook was installed would otherwise keep handing +// it addresses it had already chosen, which no hostname policy can judge. +func Lookup(ctx context.Context, host string, resolve func(context.Context, string) ([]string, error)) ([]string, error) { + if Installed() { + return []string{host}, nil + } + return resolve(ctx, host) +} + +// Dialer adapts DialContext to the interface the MongoDB driver takes. +type Dialer struct{} + +// DialContext implements the driver's dialer interface. +func (Dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + return DialContext(ctx, network, address) +} diff --git a/egress/egress_test.go b/egress/egress_test.go new file mode 100644 index 000000000..c2b86f435 --- /dev/null +++ b/egress/egress_test.go @@ -0,0 +1,126 @@ +package egress + +import ( + "context" + "errors" + "net" + "testing" +) + +func TestDialContextUsesTheInstalledHook(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + // Cleanup goes through the handles the test holds, the way any other + // owner removes a hook - never by reaching into the package's state. + var held []*Hook + t.Cleanup(func() { + for _, hook := range held { + hook.Uninstall() + } + }) + + // Without a hook the default dialer connects. + conn, err := DialContext(context.Background(), "tcp", listener.Addr().String()) + if err != nil { + t.Fatalf("default dialer: %v", err) + } + conn.Close() + if Installed() { + t.Fatal("no hook is installed yet") + } + + // With one, every connection goes through it - including the ones opened + // through the Dialer adapter. + refused := errors.New("outside the air gap") + var dialed []string + hook, err := Install(func(_ context.Context, _, address string) (net.Conn, error) { + dialed = append(dialed, address) + return nil, refused + }) + if err != nil { + t.Fatal(err) + } + held = append(held, hook) + if !Installed() { + t.Fatal("Installed must report the hook") + } + if _, err := DialContext(context.Background(), "tcp", "db.example.com:5432"); !errors.Is(err, refused) { + t.Fatalf("hook error = %v, want the hook's refusal", err) + } + if _, err := (Dialer{}).DialContext(context.Background(), "tcp", "cache.example.com:6379"); !errors.Is(err, refused) { + t.Fatalf("adapter error = %v, want the hook's refusal", err) + } + if len(dialed) != 2 || dialed[0] != "db.example.com:5432" || dialed[1] != "cache.example.com:6379" { + t.Fatalf("hook saw %v", dialed) + } + + // A second hook is refused rather than silently replacing the first: the + // policy belongs to whoever composed the process. + if _, err := Install(func(context.Context, string, string) (net.Conn, error) { return nil, nil }); err == nil { + t.Fatal("installing over an existing hook must be an error") + } + if _, err := DialContext(context.Background(), "tcp", "db.example.com:5432"); !errors.Is(err, refused) { + t.Fatalf("the first hook must still be in force, got %v", err) + } + if _, err := Install(nil); err == nil { + t.Fatal("a nil dial must be an error; the handle removes a hook") + } + + // The handle its owner holds restores the default dialer, and says so + // only once: a stale handle must not remove a policy someone else + // installed afterwards. + hook.Uninstall() + if Installed() { + t.Fatal("the handle must remove the hook") + } + replacement, err := Install(func(context.Context, string, string) (net.Conn, error) { return nil, refused }) + if err != nil { + t.Fatal(err) + } + held = append(held, replacement) + hook.Uninstall() + if !Installed() { + t.Fatal("a stale handle must not remove the current hook") + } + replacement.Uninstall() + conn, err = DialContext(context.Background(), "tcp", listener.Addr().String()) + if err != nil { + t.Fatalf("default dialer after removal: %v", err) + } + conn.Close() +} + +// Clients that resolve a name themselves before dialing ask Lookup which +// answer to use. The decision is made per connection: a client built before +// a hook was installed would otherwise keep handing it addresses it had +// already chosen, which no hostname policy can judge. +func TestLookupFollowsTheHookAtCallTime(t *testing.T) { + resolved := []string{"10.0.0.4"} + resolve := func(context.Context, string) ([]string, error) { return resolved, nil } + + // A client constructed with no hook installed. + got, err := Lookup(context.Background(), "db.internal", resolve) + if err != nil { + t.Fatal(err) + } + if len(got) != 1 || got[0] != "10.0.0.4" { + t.Fatalf("without a hook the client's own resolver answers, got %v", got) + } + + // The hook arrives afterwards; the same client must now hand it the name. + hook, err := Install(func(context.Context, string, string) (net.Conn, error) { return nil, nil }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(hook.Uninstall) + got, err = Lookup(context.Background(), "db.internal", resolve) + if err != nil { + t.Fatal(err) + } + if len(got) != 1 || got[0] != "db.internal" { + t.Fatalf("a hook installed later must still see the hostname, got %v", got) + } +} diff --git a/internal/cache/egress_hook_test.go b/internal/cache/egress_hook_test.go new file mode 100644 index 000000000..ac10548a2 --- /dev/null +++ b/internal/cache/egress_hook_test.go @@ -0,0 +1,33 @@ +package cache + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/enterpilot/gomodel/egress" +) + +// Redis speaks its own protocol over raw TCP and reads no proxy variable, so +// an egress policy reaches it only through the hook. The dial is refused on +// purpose: what matters is that the hook was asked, and for the address the +// URL names. +func TestRedisStoreDialsThroughTheEgressHook(t *testing.T) { + var address string + hook, err := egress.Install(func(_ context.Context, _, addr string) (net.Conn, error) { + address = addr + return nil, errors.New("refused by the test hook") + }) + if err != nil { + t.Fatal(err) + } + defer hook.Uninstall() + + if _, err := NewRedisStore(RedisStoreConfig{URL: "redis://cache.example.com:6379/0"}); err == nil { + t.Fatal("expected the refused dial to fail the connection") + } + if address != "cache.example.com:6379" { + t.Fatalf("hook saw %q, want the configured cache address", address) + } +} diff --git a/internal/cache/redis.go b/internal/cache/redis.go index 4a9b8270d..0a9c48d8e 100644 --- a/internal/cache/redis.go +++ b/internal/cache/redis.go @@ -7,6 +7,8 @@ import ( "time" "github.com/redis/go-redis/v9" + + "github.com/enterpilot/gomodel/egress" ) const ( @@ -34,6 +36,7 @@ func NewRedisStore(cfg RedisStoreConfig) (*RedisStore, error) { if err != nil { return nil, fmt.Errorf("invalid redis URL: %w", err) } + opts.Dialer = egress.DialContext client := redis.NewClient(opts) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() diff --git a/internal/responsecache/egress_hook_test.go b/internal/responsecache/egress_hook_test.go new file mode 100644 index 000000000..92703d8ef --- /dev/null +++ b/internal/responsecache/egress_hook_test.go @@ -0,0 +1,37 @@ +package responsecache + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/enterpilot/gomodel/config" + "github.com/enterpilot/gomodel/egress" +) + +// The pgvector store opens its own pool, so it needs the hook wired +// separately from the PostgreSQL storage backend. The dial is refused on +// purpose: what matters is that the hook was asked, and for the host the +// URL names rather than an address pgx resolved on its own. +func TestPGVectorStoreDialsThroughTheEgressHook(t *testing.T) { + var address string + hook, err := egress.Install(func(_ context.Context, _, addr string) (net.Conn, error) { + address = addr + return nil, errors.New("refused by the test hook") + }) + if err != nil { + t.Fatal(err) + } + defer hook.Uninstall() + + if _, err := newPGVectorStore(config.PGVectorConfig{ + URL: "postgres://user:pw@vectors.example.com:5432/gomodel", + Dimension: 8, + }); err == nil { + t.Fatal("expected the refused dial to fail the store") + } + if address != "vectors.example.com:5432" { + t.Fatalf("hook saw %q, want the configured vector store address", address) + } +} diff --git a/internal/responsecache/vecstore_pgvector.go b/internal/responsecache/vecstore_pgvector.go index 77e54046a..7c8e43815 100644 --- a/internal/responsecache/vecstore_pgvector.go +++ b/internal/responsecache/vecstore_pgvector.go @@ -10,6 +10,7 @@ import ( "github.com/jackc/pgx/v5/pgxpool" "github.com/enterpilot/gomodel/config" + "github.com/enterpilot/gomodel/egress" ) type pgVecStore struct { @@ -34,7 +35,18 @@ func newPGVectorStore(cfg config.PGVectorConfig) (*pgVecStore, error) { if err := validatePGIdentifier(tbl); err != nil { return nil, fmt.Errorf("vecstore pgvector: table: %w", err) } - pool, err := pgxpool.New(context.Background(), cfg.URL) + poolCfg, err := pgxpool.ParseConfig(cfg.URL) + if err != nil { + return nil, fmt.Errorf("vecstore pgvector: url: %w", err) + } + // See NewPostgreSQL: the dial and the lookup before it both go through + // the egress hook, decided per connection. + poolCfg.ConnConfig.DialFunc = egress.DialContext + resolve := poolCfg.ConnConfig.LookupFunc + poolCfg.ConnConfig.LookupFunc = func(ctx context.Context, host string) ([]string, error) { + return egress.Lookup(ctx, host, resolve) + } + pool, err := pgxpool.NewWithConfig(context.Background(), poolCfg) if err != nil { return nil, fmt.Errorf("vecstore pgvector: connect: %w", err) } diff --git a/internal/storage/egress_hook_test.go b/internal/storage/egress_hook_test.go new file mode 100644 index 000000000..792c29341 --- /dev/null +++ b/internal/storage/egress_hook_test.go @@ -0,0 +1,64 @@ +package storage + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/enterpilot/gomodel/egress" +) + +// The database clients speak their own protocols over raw TCP and read no +// proxy variable, so an egress policy reaches them only through the hook. +// These tests fail the dial on purpose: what matters is that the hook was +// asked at all, and with the address the URL names. +func TestPostgreSQLDialsThroughTheEgressHook(t *testing.T) { + address, restore := recordDials(t) + if _, err := NewPostgreSQL(context.Background(), PostgreSQLConfig{URL: "postgres://user:pw@db.example.com:5432/gomodel"}); err == nil { + t.Fatal("expected the refused dial to fail the connection") + } + restore() + if *address != "db.example.com:5432" { + t.Fatalf("hook saw %q, want the configured database address", *address) + } +} + +func TestMongoDBDialsThroughTheEgressHook(t *testing.T) { + address, restore := recordDials(t) + _, err := NewMongoDB(context.Background(), MongoDBConfig{ + URL: "mongodb://db.example.com:27017/gomodel?serverSelectionTimeoutMS=200&connectTimeoutMS=200", + }) + restore() + if err == nil { + t.Fatal("expected the refused dial to fail the connection") + } + if *address != "db.example.com:27017" { + t.Fatalf("hook saw %q, want the configured database address", *address) + } +} + +// recordDials installs a hook that refuses every connection and records the +// last address it was asked for. The returned func removes it again; call it +// before asserting, so a failure never leaves the hook installed for the +// rest of the package's tests. +func recordDials(t *testing.T) (*string, func()) { + t.Helper() + var address string + hook, err := egress.Install(func(_ context.Context, _, addr string) (net.Conn, error) { + address = addr + return nil, errors.New("refused by the test hook") + }) + if err != nil { + t.Fatal(err) + } + removed := false + remove := func() { + if !removed { + removed = true + hook.Uninstall() + } + } + t.Cleanup(remove) + return &address, remove +} diff --git a/internal/storage/mongodb.go b/internal/storage/mongodb.go index 8963161ee..e8ce33af7 100644 --- a/internal/storage/mongodb.go +++ b/internal/storage/mongodb.go @@ -8,6 +8,8 @@ import ( "go.mongodb.org/mongo-driver/v2/mongo" "go.mongodb.org/mongo-driver/v2/mongo/options" + + "github.com/enterpilot/gomodel/egress" ) // DefaultMongoDatabase is the database used when neither the explicit Database @@ -29,7 +31,7 @@ func NewMongoDB(ctx context.Context, cfg MongoDBConfig) (MongoDBStorage, error) dbName := resolveMongoDatabase(cfg) // Create client options - clientOpts := options.Client().ApplyURI(cfg.URL) + clientOpts := options.Client().ApplyURI(cfg.URL).SetDialer(egress.Dialer{}) // Connect to MongoDB client, err := mongo.Connect(clientOpts) diff --git a/internal/storage/postgresql.go b/internal/storage/postgresql.go index 3dfd6228a..13fe10252 100644 --- a/internal/storage/postgresql.go +++ b/internal/storage/postgresql.go @@ -6,6 +6,8 @@ import ( "math" "github.com/jackc/pgx/v5/pgxpool" + + "github.com/enterpilot/gomodel/egress" ) // postgresStorage implements Storage for PostgreSQL @@ -34,6 +36,19 @@ func NewPostgreSQL(ctx context.Context, cfg PostgreSQLConfig) (PostgreSQLStorage poolCfg.MaxConns = 10 // default } + // Every connection the pool opens goes through the egress hook, so a + // distribution that polices outbound traffic sees the database too. pgx + // resolves the host before it dials, so the lookup is deferred to the + // hook as well: with one installed the name reaches it unresolved and + // the hook dials the address that passed, and without one pgx's own + // resolver answers. Both are decided per connection, so a hook installed + // after this pool was built still sees hostnames. + poolCfg.ConnConfig.DialFunc = egress.DialContext + resolve := poolCfg.ConnConfig.LookupFunc + poolCfg.ConnConfig.LookupFunc = func(ctx context.Context, host string) ([]string, error) { + return egress.Lookup(ctx, host, resolve) + } + // Create the connection pool pool, err := pgxpool.NewWithConfig(ctx, poolCfg) if err != nil {