Skip to content
Merged
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
2 changes: 2 additions & 0 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ require (
github.com/nais/naistrix v0.35.0
github.com/pkg/errors v0.9.1
github.com/pterm/pterm v0.12.83
github.com/quic-go/quic-go v0.59.1
github.com/sethvargo/go-retry v0.3.0
github.com/stretchr/testify v1.11.1
github.com/suessflorian/gqlfetch v0.7.0
Expand Down Expand Up @@ -198,6 +199,7 @@ require (
github.com/prometheus/client_model v0.6.2 // indirect
github.com/prometheus/common v0.67.5 // indirect
github.com/prometheus/procfs v0.19.2 // indirect
github.com/quic-go/qpack v0.6.0 // indirect
github.com/rivo/uniseg v0.4.7 // indirect
github.com/russross/blackfriday/v2 v2.1.0 // indirect
github.com/sagikazarmark/locafero v0.12.0 // indirect
Expand Down
6 changes: 6 additions & 0 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -457,6 +457,12 @@ github.com/prometheus/procfs v0.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4
github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw=
github.com/pterm/pterm v0.12.83 h1:ie+YmGmA727VuhxBlyGr74Ks+7McV6kT99IB8EU80aA=
github.com/pterm/pterm v0.12.83/go.mod h1:xlgc6bFWyJIMtmLJvGim+L7jhSReilOlOnodeIYe4Tk=
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII=
github.com/quic-go/quic-go v0.57.0 h1:AsSSrrMs4qI/hLrKlTH/TGQeTMY0ib1pAOX7vA3AdqE=
github.com/quic-go/quic-go v0.57.0/go.mod h1:ly4QBAjHA2VhdnxhojRsCUOeJwKYg+taDlos92xb1+s=
github.com/quic-go/quic-go v0.59.1 h1:0Gmua0HW1Tv7ANR7hUYwRyD0MG5OJfgvYSZasGZzBic=
github.com/quic-go/quic-go v0.59.1/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ=
github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88=
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
Expand Down
2 changes: 2 additions & 0 deletions internal/alpha/command/alpha.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package command

import (
"github.com/nais/cli/internal/alpha/command/flag"
postgrescmd "github.com/nais/cli/internal/alpha/postgres/command"
"github.com/nais/cli/internal/flags"
krakend "github.com/nais/cli/internal/krakend/command"
mcpcmd "github.com/nais/cli/internal/mcp/command"
Expand All @@ -18,6 +19,7 @@ func Alpha(parentFlags *flags.GlobalFlags) *naistrix.Command {
SubCommands: []*naistrix.Command{
krakend.Krakend(flags),
mcpcmd.MCP(flags),
postgrescmd.Postgres(parentFlags),
},
}
}
134 changes: 134 additions & 0 deletions internal/alpha/postgres/access.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,134 @@
package postgres

import (
"context"
"fmt"
"time"

"github.com/Khan/genqlient/graphql"
"github.com/nais/cli/internal/naisapi"
"github.com/nais/cli/internal/naisapi/gql"
)

// Access contains the brokered connection materials. Do not log this value.
type Access struct {
State gql.PostgresAccessState
Message string
Connection *Connection
}

type Connection struct {
Username, Password, CACertificate, ServerName, RelayEndpoint, RelayAccess, RelayToken string
}

type AccessAPI interface {
ActiveBranch(context.Context, string, string, string) (string, error)
Create(context.Context, gql.CreatePostgresAccessInput) (string, error)
Get(context.Context, string, string, string) (Access, error)
}

type graphqlAccessAPI struct{ client graphql.Client }

func NewAPI(ctx context.Context) (AccessAPI, error) {
client, err := naisapi.GraphqlClient(ctx)
if err != nil {
return nil, err
}
return graphqlAccessAPI{client}, nil
}

func (a graphqlAccessAPI) ActiveBranch(ctx context.Context, team, environment, name string) (string, error) {
_ = `# @genqlient
query GetActivePostgresBranchAlpha($team: Slug!, $environment: String!, $postgres: String!) {
team(slug: $team) { environment(name: $environment) { postgres(name: $postgres) { activeBranch { name } } } }
}
`
result, err := gql.GetActivePostgresBranchAlpha(ctx, a.client, team, environment, name)
if err != nil {
return "", err
}
if result.Team.Environment.Postgres.ActiveBranch == nil {
return "", fmt.Errorf("postgres %q has no active branch; specify --branch", name)
}
return result.Team.Environment.Postgres.ActiveBranch.Name, nil
}

func (a graphqlAccessAPI) Create(ctx context.Context, input gql.CreatePostgresAccessInput) (string, error) {
_ = `# @genqlient
mutation CreatePostgresAccessAlpha($input: CreatePostgresAccessInput!) {
createPostgresAccess(input: $input) { name }
}
`
result, err := gql.CreatePostgresAccessAlpha(ctx, a.client, input)
if err != nil {
return "", err
}
return result.CreatePostgresAccess.Name, nil
}

func (a graphqlAccessAPI) Get(ctx context.Context, team, environment, name string) (Access, error) {
_ = `# @genqlient
query GetPostgresAccessAlpha($team: Slug!, $environment: String!, $name: String!) {
team(slug: $team) { environment(name: $environment) { postgresAccess(name: $name) {
state message connection { username password caCertificate serverName relayEndpoint relayAccess relayToken }
} } }
}
`
result, err := gql.GetPostgresAccessAlpha(ctx, a.client, team, environment, name)
if err != nil {
return Access{}, err
}
got := result.Team.Environment.PostgresAccess
access := Access{State: got.State}
if got.Message != nil {
access.Message = *got.Message
}
if got.Connection != nil {
c := got.Connection
access.Connection = &Connection{c.Username, c.Password, c.CaCertificate, c.ServerName, c.RelayEndpoint, c.RelayAccess, c.RelayToken}
}
return access, nil
}

func waitForAccess(ctx context.Context, api AccessAPI, team, environment, name string, interval time.Duration) (Connection, error) {
for {
access, err := api.Get(ctx, team, environment, name)
if err != nil {
return Connection{}, fmt.Errorf("retrieve postgres access: %w", err)
}
switch access.State {
case gql.PostgresAccessStateReady:
if access.Connection == nil {
return Connection{}, fmt.Errorf("postgres access %q is ready without connection materials", name)
}
return *access.Connection, nil
case gql.PostgresAccessStateFailed, gql.PostgresAccessStateExpired:
return Connection{}, fmt.Errorf("postgres access %q is %s: %s", name, access.State, access.Message)
case gql.PostgresAccessStatePending:
default:
return Connection{}, fmt.Errorf("postgres access %q has unknown state %q", name, access.State)
}
timer := time.NewTimer(interval)
select {
case <-ctx.Done():
timer.Stop()
return Connection{}, fmt.Errorf("waiting for postgres access %q: %w", name, ctx.Err())
case <-timer.C:
}
}
}

func CreateAndWait(ctx context.Context, api AccessAPI, input gql.CreatePostgresAccessInput) (Connection, error) {
// Bound creation and polling together; a stalled API must not hang the command.
setupCtx, cancel := context.WithTimeout(ctx, 60*time.Second)
defer cancel()
name, err := api.Create(setupCtx, input)
if err != nil {
return Connection{}, fmt.Errorf("create postgres access: %w", err)
}
connection, err := waitForAccess(setupCtx, api, input.TeamSlug, input.EnvironmentName, name, time.Second)
if err != nil {
return Connection{}, fmt.Errorf("access %q was created but is not ready (it expires after its requested TTL): %w", name, err)
}
return connection, nil
}
71 changes: 71 additions & 0 deletions internal/alpha/postgres/access_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
package postgres

import (
"context"
"errors"
"strings"
"testing"
"time"

"github.com/nais/cli/internal/naisapi/gql"
)

type fakeAccessAPI struct {
states []Access
calls int
}

func (f *fakeAccessAPI) ActiveBranch(context.Context, string, string, string) (string, error) {
return "main", nil
}

func (f *fakeAccessAPI) Create(context.Context, gql.CreatePostgresAccessInput) (string, error) {
return "access-1", nil
}

func (f *fakeAccessAPI) Get(context.Context, string, string, string) (Access, error) {
index := f.calls
f.calls++
if index >= len(f.states) {
return Access{}, errors.New("unexpected poll")
}
return f.states[index], nil
}

func TestWaitForAccess(t *testing.T) {
for _, tt := range []struct {
name string
states []Access
want string
}{
{"ready", []Access{{State: gql.PostgresAccessStatePending}, {State: gql.PostgresAccessStateReady, Connection: &Connection{Username: "alice"}}}, ""},
{"failed", []Access{{State: gql.PostgresAccessStateFailed, Message: "database unavailable"}}, "database unavailable"},
{"expired", []Access{{State: gql.PostgresAccessStateExpired}}, "EXPIRED"},
{"missing materials", []Access{{State: gql.PostgresAccessStateReady}}, "without connection materials"},
} {
t.Run(tt.name, func(t *testing.T) {
fake := &fakeAccessAPI{states: tt.states}
got, err := waitForAccess(context.Background(), fake, "team", "dev", "access-1", time.Millisecond)
if tt.want == "" {
if err != nil || got.Username != "alice" {
t.Fatalf("got %+v, err %v", got, err)
}
} else if err == nil || !strings.Contains(err.Error(), tt.want) {
t.Fatalf("expected %q, got %v", tt.want, err)
}
if fake.calls != len(tt.states) {
t.Fatalf("polled %d times, want %d", fake.calls, len(tt.states))
}
})
}
}

func TestWaitCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
fake := &fakeAccessAPI{states: []Access{{State: gql.PostgresAccessStatePending}}}
_, err := waitForAccess(ctx, fake, "team", "dev", "access-1", time.Hour)
if !errors.Is(err, context.Canceled) {
t.Fatalf("expected cancellation, got %v", err)
}
}
Loading
Loading