mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: use authenticated urls for pubsub (#14261)
This commit is contained in:
@@ -10,7 +10,10 @@ import (
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
"github.com/aws/aws-sdk-go-v2/config"
|
||||
"github.com/aws/aws-sdk-go-v2/feature/rds/auth"
|
||||
"github.com/lib/pq"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
)
|
||||
|
||||
type awsIamRdsDriver struct {
|
||||
@@ -18,7 +21,10 @@ type awsIamRdsDriver struct {
|
||||
cfg aws.Config
|
||||
}
|
||||
|
||||
var _ driver.Driver = &awsIamRdsDriver{}
|
||||
var (
|
||||
_ driver.Driver = &awsIamRdsDriver{}
|
||||
_ database.ConnectorCreator = &awsIamRdsDriver{}
|
||||
)
|
||||
|
||||
// Register initializes and registers our aws iam rds wrapped database driver.
|
||||
func Register(ctx context.Context, parentName string) (string, error) {
|
||||
@@ -65,6 +71,16 @@ func (d *awsIamRdsDriver) Open(name string) (driver.Conn, error) {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// Connector returns a driver.Connector that fetches a new authentication token for each connection.
|
||||
func (d *awsIamRdsDriver) Connector(name string) (driver.Connector, error) {
|
||||
connector := &connector{
|
||||
url: name,
|
||||
cfg: d.cfg,
|
||||
}
|
||||
|
||||
return connector, nil
|
||||
}
|
||||
|
||||
func getAuthenticatedURL(cfg aws.Config, dbURL string) (string, error) {
|
||||
nURL, err := url.Parse(dbURL)
|
||||
if err != nil {
|
||||
@@ -82,3 +98,37 @@ func getAuthenticatedURL(cfg aws.Config, dbURL string) (string, error) {
|
||||
|
||||
return nURL.String(), nil
|
||||
}
|
||||
|
||||
type connector struct {
|
||||
url string
|
||||
cfg aws.Config
|
||||
dialer pq.Dialer
|
||||
}
|
||||
|
||||
var _ database.DialerConnector = &connector{}
|
||||
|
||||
func (c *connector) Connect(ctx context.Context) (driver.Conn, error) {
|
||||
nURL, err := getAuthenticatedURL(c.cfg, c.url)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("assigning authentication token to url: %w", err)
|
||||
}
|
||||
|
||||
nc, err := pq.NewConnector(nURL)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("creating new connector: %w", err)
|
||||
}
|
||||
|
||||
if c.dialer != nil {
|
||||
nc.Dialer(c.dialer)
|
||||
}
|
||||
|
||||
return nc.Connect(ctx)
|
||||
}
|
||||
|
||||
func (*connector) Driver() driver.Driver {
|
||||
return &pq.Driver{}
|
||||
}
|
||||
|
||||
func (c *connector) Dialer(dialer pq.Dialer) {
|
||||
c.dialer = dialer
|
||||
}
|
||||
|
||||
@@ -7,10 +7,11 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
|
||||
"github.com/coder/coder/v2/cli"
|
||||
awsrdsiam "github.com/coder/coder/v2/coderd/database/awsiamrds"
|
||||
"github.com/coder/coder/v2/coderd/database/awsiamrds"
|
||||
"github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
@@ -22,13 +23,15 @@ func TestDriver(t *testing.T) {
|
||||
// export DBAWSIAMRDS_TEST_URL="postgres://user@host:5432/dbname";
|
||||
url := os.Getenv("DBAWSIAMRDS_TEST_URL")
|
||||
if url == "" {
|
||||
t.Log("skipping test; no DBAWSIAMRDS_TEST_URL set")
|
||||
t.Skip()
|
||||
}
|
||||
|
||||
logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
sqlDriver, err := awsrdsiam.Register(ctx, "postgres")
|
||||
sqlDriver, err := awsiamrds.Register(ctx, "postgres")
|
||||
require.NoError(t, err)
|
||||
|
||||
db, err := cli.ConnectToPostgres(ctx, slogtest.Make(t, nil), sqlDriver, url)
|
||||
@@ -47,4 +50,23 @@ func TestDriver(t *testing.T) {
|
||||
var one int
|
||||
require.NoError(t, i.Scan(&one))
|
||||
require.Equal(t, 1, one)
|
||||
|
||||
ps, err := pubsub.New(ctx, logger, db, url)
|
||||
require.NoError(t, err)
|
||||
|
||||
gotChan := make(chan struct{})
|
||||
subCancel, err := ps.Subscribe("test", func(_ context.Context, _ []byte) {
|
||||
close(gotChan)
|
||||
})
|
||||
defer subCancel()
|
||||
require.NoError(t, err)
|
||||
|
||||
err = ps.Publish("test", []byte("hello"))
|
||||
require.NoError(t, err)
|
||||
|
||||
select {
|
||||
case <-gotChan:
|
||||
case <-ctx.Done():
|
||||
require.Fail(t, "timed out waiting for message")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user