Skip to content
Open
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
103 changes: 103 additions & 0 deletions egress/egress.go
Original file line number Diff line number Diff line change
@@ -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]
Comment thread
greptile-apps[bot] marked this conversation as resolved.

// 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)
}
126 changes: 126 additions & 0 deletions egress/egress_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
33 changes: 33 additions & 0 deletions internal/cache/egress_hook_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
3 changes: 3 additions & 0 deletions internal/cache/redis.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@ import (
"time"

"github.com/redis/go-redis/v9"

"github.com/enterpilot/gomodel/egress"
)

const (
Expand Down Expand Up @@ -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()
Expand Down
37 changes: 37 additions & 0 deletions internal/responsecache/egress_hook_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
14 changes: 13 additions & 1 deletion internal/responsecache/vecstore_pgvector.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"github.com/jackc/pgx/v5/pgxpool"

"github.com/enterpilot/gomodel/config"
"github.com/enterpilot/gomodel/egress"
)

type pgVecStore struct {
Expand All @@ -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
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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)
}
Expand Down
Loading