chore: refactor instance identity to be a SessionTokenProvider (#19566)

Refactors Agent instance identity to be a SessionTokenProvider.

Refactors the CLI to create Agent clients via a centralized function, rather than add-hoc via individual command handlers and their flags.

This allows commands besides `coder agent`, but which still use the agent identity, to support instance identity authentication.

Fixes #19111 by unifying all API requests to go thru the SessionTokenProvider for auth credentials.
This commit is contained in:
Spike Curtis
2025-09-03 10:38:42 +04:00
committed by GitHub
parent ee35ad3a57
commit 1354d84eb4
35 changed files with 640 additions and 478 deletions
+6 -12
View File
@@ -432,8 +432,7 @@ func TestExternalAuthCallback(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, template.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
_, err := agentClient.ExternalAuth(context.Background(), agentsdk.ExternalAuthRequest{
Match: "github.com",
})
@@ -464,8 +463,7 @@ func TestExternalAuthCallback(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, template.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
token, err := agentClient.ExternalAuth(context.Background(), agentsdk.ExternalAuthRequest{
Match: "github.com/asd/asd",
})
@@ -565,8 +563,7 @@ func TestExternalAuthCallback(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, template.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
resp := coderdtest.RequestExternalAuthCallback(t, "github", client)
require.Equal(t, http.StatusTemporaryRedirect, resp.StatusCode)
@@ -627,8 +624,7 @@ func TestExternalAuthCallback(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, template.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
token, err := agentClient.ExternalAuth(context.Background(), agentsdk.ExternalAuthRequest{
Match: "github.com/asd/asd",
@@ -674,8 +670,7 @@ func TestExternalAuthCallback(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, template.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
token, err := agentClient.ExternalAuth(context.Background(), agentsdk.ExternalAuthRequest{
Match: "github.com/asd/asd",
@@ -740,8 +735,7 @@ func TestExternalAuthCallback(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, template.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
token, err := agentClient.ExternalAuth(t.Context(), agentsdk.ExternalAuthRequest{
Match: "github.com/asd/asd",
+2 -4
View File
@@ -118,8 +118,7 @@ func TestAgentGitSSHKey(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, project.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
@@ -157,8 +156,7 @@ func TestAgentGitSSHKey_APIKeyScopes(t *testing.T) {
workspace := coderdtest.CreateWorkspace(t, client, project.ID)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
+2 -4
View File
@@ -585,8 +585,7 @@ func TestTemplateInsights_Golden(t *testing.T) {
continue
}
authToken := uuid.New()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken.String())
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken.String()))
workspace.agentClient = agentClient
var apps []*proto.App
@@ -1494,8 +1493,7 @@ func TestUserActivityInsights_Golden(t *testing.T) {
continue
}
authToken := uuid.New()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken.String())
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken.String()))
workspace.agentClient = agentClient
var apps []*proto.App
@@ -90,8 +90,7 @@ func TestCollectInsights(t *testing.T) {
// Start an agent so that we can generate stats.
var agentClients []agentproto.DRPCAgentClient
for i, agent := range []database.WorkspaceAgent{agent1, agent2} {
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(agent.AuthToken.String())
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(agent.AuthToken.String()))
agentClient.SDK.SetLogger(logger.Leveled(slog.LevelDebug).Named(fmt.Sprintf("agent%d", i+1)))
conn, err := agentClient.ConnectRPC(context.Background())
require.NoError(t, err)
@@ -875,8 +875,7 @@ func prepareWorkspaceAndAgent(ctx context.Context, t *testing.T, client *codersd
})
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, workspace.LatestBuild.ID)
ac := agentsdk.New(client.URL)
ac.SetSessionToken(authToken)
ac := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
conn, err := ac.ConnectRPC(ctx)
require.NoError(t, err)
agentAPI := agentproto.NewDRPCAgentClient(conn)
+16 -32
View File
@@ -228,8 +228,7 @@ func TestWorkspaceAgentLogs(t *testing.T) {
OwnerID: user.UserID,
}).WithAgent().Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
err := agentClient.PatchLogs(ctx, agentsdk.PatchLogs{
Logs: []agentsdk.Log{
{
@@ -269,8 +268,7 @@ func TestWorkspaceAgentLogs(t *testing.T) {
OrganizationID: user.OrganizationID,
OwnerID: user.UserID,
}).WithAgent().Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
err := agentClient.PatchLogs(ctx, agentsdk.PatchLogs{
Logs: []agentsdk.Log{
{
@@ -314,8 +312,7 @@ func TestWorkspaceAgentLogs(t *testing.T) {
updates, err := client.WatchWorkspace(ctx, r.Workspace.ID)
require.NoError(t, err)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
err = agentClient.PatchLogs(ctx, agentsdk.PatchLogs{
Logs: []agentsdk.Log{{
CreatedAt: dbtime.Now(),
@@ -360,8 +357,7 @@ func TestWorkspaceAgentAppStatus(t *testing.T) {
return a
}).Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
t.Run("Success", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
@@ -542,8 +538,7 @@ func TestWorkspaceAgentConnectRPC(t *testing.T) {
require.NoError(t, err)
coderdtest.AwaitWorkspaceBuildJobCompleted(t, client, stopBuild.ID)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
_, err = agentClient.ConnectRPC(ctx)
require.Error(t, err)
@@ -568,8 +563,7 @@ func TestWorkspaceAgentConnectRPC(t *testing.T) {
)
require.NoError(t, err)
// Then: the agent token should no longer be valid
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(wsb.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken((wsb.AgentToken)))
_, err = agentClient.ConnectRPC(ctx)
require.Error(t, err)
var sdkErr *codersdk.Error
@@ -890,8 +884,7 @@ func TestWorkspaceAgentTailnetDirectDisabled(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
// Verify that the manifest has DisableDirectConnections set to true.
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
rpc, err := agentClient.ConnectRPC(ctx)
require.NoError(t, err)
defer func() {
@@ -1742,8 +1735,7 @@ func TestWorkspaceAgentAppHealth(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
conn, err := agentClient.ConnectRPC(ctx)
require.NoError(t, err)
defer func() {
@@ -1818,8 +1810,7 @@ func TestWorkspaceAgentPostLogSource(t *testing.T) {
OwnerID: user.UserID,
}).WithAgent().Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
req := agentsdk.PostLogSourceRequest{
ID: uuid.New(),
@@ -1867,8 +1858,7 @@ func TestWorkspaceAgent_LifecycleState(t *testing.T) {
}
}
ac := agentsdk.New(client.URL)
ac.SetSessionToken(r.AgentToken)
ac := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
conn, err := ac.ConnectRPC(ctx)
require.NoError(t, err)
defer func() {
@@ -1965,8 +1955,7 @@ func TestWorkspaceAgent_Metadata(t *testing.T) {
}
}
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
ctx := testutil.Context(t, testutil.WaitMedium)
conn, err := agentClient.ConnectRPC(ctx)
@@ -2229,8 +2218,7 @@ func TestWorkspaceAgent_Metadata_CatchMemoryLeak(t *testing.T) {
}
}
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
ctx := testutil.Context(t, testutil.WaitSuperLong)
conn, err := agentClient.ConnectRPC(ctx)
@@ -2335,8 +2323,7 @@ func TestWorkspaceAgent_Startup(t *testing.T) {
OrganizationID: user.OrganizationID,
OwnerID: user.UserID,
}).WithAgent().Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
ctx := testutil.Context(t, testutil.WaitMedium)
@@ -2382,8 +2369,7 @@ func TestWorkspaceAgent_Startup(t *testing.T) {
OwnerID: user.UserID,
}).WithAgent().Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
ctx := testutil.Context(t, testutil.WaitMedium)
@@ -2547,8 +2533,7 @@ func TestWorkspaceAgentExternalAuthListen(t *testing.T) {
return agents
}).Do()
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
// We need to include an invalid oauth token that is not expired.
dbgen.ExternalAuthLink(t, db, database.ExternalAuthLink{
@@ -3028,8 +3013,7 @@ func TestReinit(t *testing.T) {
pubsubSpy.Unlock()
agentCtx := testutil.Context(t, testutil.WaitShort)
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(r.AgentToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
agentReinitializedCh := make(chan *agentsdk.ReinitializationEvent)
go func() {
+2 -4
View File
@@ -68,8 +68,7 @@ func TestWorkspaceAgentReportStats(t *testing.T) {
},
).Do()
ac := agentsdk.New(client.URL)
ac.SetSessionToken(r.AgentToken)
ac := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
conn, err := ac.ConnectRPC(context.Background())
require.NoError(t, err)
defer func() {
@@ -155,8 +154,7 @@ func TestAgentAPI_LargeManifest(t *testing.T) {
agents[0].ApiKeyScope = string(tc.apiKeyScope)
return agents
}).Do()
ac := agentsdk.New(client.URL)
ac.SetSessionToken(r.AgentToken)
ac := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken))
conn, err := ac.ConnectRPC(ctx)
defer func() {
_ = conn.Close()
+1 -2
View File
@@ -482,8 +482,7 @@ func createWorkspaceWithApps(t *testing.T, client *codersdk.Client, orgID uuid.U
require.Equal(t, appURL.String(), app.SubdomainName)
}
agentClient := agentsdk.New(client.URL)
agentClient.SetSessionToken(authToken)
agentClient := agentsdk.New(client.URL, agentsdk.WithFixedToken(authToken))
// TODO (@dean): currently, the primary app host is used when generating
// the port URL we tell the agent to use. We don't have any plans to change
+12 -22
View File
@@ -51,11 +51,9 @@ func TestPostWorkspaceAuthAzureInstanceIdentity(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
client.HTTPClient = metadataClient
agentClient := &agentsdk.Client{
SDK: client,
}
_, err := agentClient.AuthAzureInstanceIdentity(ctx)
agentClient := agentsdk.New(client.URL, agentsdk.WithAzureInstanceIdentity())
agentClient.SDK.HTTPClient = metadataClient
err := agentClient.RefreshToken(ctx)
require.NoError(t, err)
}
@@ -97,11 +95,9 @@ func TestPostWorkspaceAuthAWSInstanceIdentity(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
client.HTTPClient = metadataClient
agentClient := &agentsdk.Client{
SDK: client,
}
_, err := agentClient.AuthAWSInstanceIdentity(ctx)
agentClient := agentsdk.New(client.URL, agentsdk.WithAWSInstanceIdentity())
agentClient.SDK.HTTPClient = metadataClient
err := agentClient.RefreshToken(ctx)
require.NoError(t, err)
})
}
@@ -119,10 +115,8 @@ func TestPostWorkspaceAuthGoogleInstanceIdentity(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
agentClient := &agentsdk.Client{
SDK: client,
}
_, err := agentClient.AuthGoogleInstanceIdentity(ctx, "", metadata)
agentClient := agentsdk.New(client.URL, agentsdk.WithGoogleInstanceIdentity("", metadata))
err := agentClient.RefreshToken(ctx)
var apiErr *codersdk.Error
require.ErrorAs(t, err, &apiErr)
require.Equal(t, http.StatusUnauthorized, apiErr.StatusCode())
@@ -139,10 +133,8 @@ func TestPostWorkspaceAuthGoogleInstanceIdentity(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
agentClient := &agentsdk.Client{
SDK: client,
}
_, err := agentClient.AuthGoogleInstanceIdentity(ctx, "", metadata)
agentClient := agentsdk.New(client.URL, agentsdk.WithGoogleInstanceIdentity("", metadata))
err := agentClient.RefreshToken(ctx)
var apiErr *codersdk.Error
require.ErrorAs(t, err, &apiErr)
require.Equal(t, http.StatusNotFound, apiErr.StatusCode())
@@ -184,10 +176,8 @@ func TestPostWorkspaceAuthGoogleInstanceIdentity(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
defer cancel()
agentClient := &agentsdk.Client{
SDK: client,
}
_, err := agentClient.AuthGoogleInstanceIdentity(ctx, "", metadata)
agentClient := agentsdk.New(client.URL, agentsdk.WithGoogleInstanceIdentity("", metadata))
err := agentClient.RefreshToken(ctx)
require.NoError(t, err)
})
}