diff --git a/agent/agent.go b/agent/agent.go index 53344cbff2..51d39719b0 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -111,6 +111,12 @@ type Client interface { ConnectRPC28(ctx context.Context) ( proto.DRPCAgentClient28, tailnetproto.DRPCTailnetClient28, error, ) + // ConnectRPC28WithRole is like ConnectRPC28 but sends an explicit + // role query parameter to the server. The workspace agent should + // use role "agent" to enable connection monitoring. + ConnectRPC28WithRole(ctx context.Context, role string) ( + proto.DRPCAgentClient28, tailnetproto.DRPCTailnetClient28, error, + ) tailnet.DERPMapRewriter agentsdk.RefreshableSessionTokenProvider } @@ -997,8 +1003,10 @@ func (a *agent) run() (retErr error) { return xerrors.Errorf("refresh token: %w", err) } - // ConnectRPC returns the dRPC connection we use for the Agent and Tailnet v2+ APIs - aAPI, tAPI, err := a.client.ConnectRPC28(a.hardCtx) + // ConnectRPC returns the dRPC connection we use for the Agent and Tailnet v2+ APIs. + // We pass role "agent" to enable connection monitoring on the server, which tracks + // the agent's connectivity state (first_connected_at, last_connected_at, disconnected_at). + aAPI, tAPI, err := a.client.ConnectRPC28WithRole(a.hardCtx, "agent") if err != nil { return err } diff --git a/agent/agenttest/client.go b/agent/agenttest/client.go index a52ce250eb..517fac6dca 100644 --- a/agent/agenttest/client.go +++ b/agent/agenttest/client.go @@ -124,6 +124,12 @@ func (c *Client) Close() { c.derpMapOnce.Do(func() { close(c.derpMapUpdates) }) } +func (c *Client) ConnectRPC28WithRole(ctx context.Context, _ string) ( + agentproto.DRPCAgentClient28, proto.DRPCTailnetClient28, error, +) { + return c.ConnectRPC28(ctx) +} + func (c *Client) ConnectRPC28(ctx context.Context) ( agentproto.DRPCAgentClient28, proto.DRPCTailnetClient28, error, ) { diff --git a/coderd/workspaceagentsrpc.go b/coderd/workspaceagentsrpc.go index 4e4cbce1ea..7272f73613 100644 --- a/coderd/workspaceagentsrpc.go +++ b/coderd/workspaceagentsrpc.go @@ -59,6 +59,17 @@ func (api *API) workspaceAgentRPC(rw http.ResponseWriter, r *http.Request) { return } + // The role parameter distinguishes the real workspace agent from + // other clients using the same agent token (e.g. coder-logstream-kube). + // Only connections with the "agent" role trigger connection monitoring + // that updates first_connected_at/last_connected_at/disconnected_at. + // For backward compatibility, we default to monitoring when the role + // is omitted, since older agents don't send this parameter. In a + // future release, once all agents include role=agent, we can change + // this default to skip monitoring for unspecified roles. + role := r.URL.Query().Get("role") + monitorConnection := role == "" || role == "agent" + api.WebsocketWaitMutex.Lock() api.WebsocketWaitGroup.Add(1) api.WebsocketWaitMutex.Unlock() @@ -121,10 +132,15 @@ func (api *API) workspaceAgentRPC(rw http.ResponseWriter, r *http.Request) { slog.F("agent_api_version", workspaceAgent.APIVersion), slog.F("agent_resource_id", workspaceAgent.ResourceID)) - closeCtx, closeCtxCancel := context.WithCancel(ctx) - defer closeCtxCancel() - monitor := api.startAgentYamuxMonitor(closeCtx, workspace, workspaceAgent, build, mux) - defer monitor.close() + if monitorConnection { + closeCtx, closeCtxCancel := context.WithCancel(ctx) + defer closeCtxCancel() + monitor := api.startAgentYamuxMonitor(closeCtx, workspace, workspaceAgent, build, mux) + defer monitor.close() + } else { + logger.Debug(ctx, "skipping agent connection monitoring", + slog.F("role", role)) + } agentAPI := agentapi.New(agentapi.Options{ AgentID: workspaceAgent.ID, diff --git a/coderd/workspaceagentsrpc_test.go b/coderd/workspaceagentsrpc_test.go index 525b8a981d..b819eaf690 100644 --- a/coderd/workspaceagentsrpc_test.go +++ b/coderd/workspaceagentsrpc_test.go @@ -11,6 +11,7 @@ import ( agentproto "github.com/coder/coder/v2/agent/proto" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/coderd/database/dbfake" "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/coder/v2/coderd/rbac" @@ -168,3 +169,85 @@ func TestAgentAPI_LargeManifest(t *testing.T) { }) } } + +func TestWorkspaceAgentRPCRole(t *testing.T) { + t.Parallel() + + t.Run("AgentRoleMonitorsConnection", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, db := coderdtest.NewWithDatabase(t, nil) + user := coderdtest.CreateFirstUser(t, client) + r := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{ + OrganizationID: user.OrganizationID, + OwnerID: user.UserID, + }).WithAgent().Do() + + // Connect with role=agent using ConnectRPCWithRole. This is + // how the real workspace agent connects. + ac := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken)) + conn, err := ac.ConnectRPCWithRole(ctx, "agent") + require.NoError(t, err) + defer func() { + _ = conn.Close() + }() + + // The connection monitor updates the database asynchronously, + // so we need to wait for first_connected_at to be set. + var agent database.WorkspaceAgent + require.Eventually(t, func() bool { + agent, err = db.GetWorkspaceAgentByID(dbauthz.AsSystemRestricted(ctx), r.Agents[0].ID) + if err != nil { + return false + } + return agent.FirstConnectedAt.Valid + }, testutil.WaitShort, testutil.IntervalFast) + assert.True(t, agent.LastConnectedAt.Valid, + "last_connected_at should be set for agent role") + }) + + t.Run("NonAgentRoleSkipsMonitoring", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, db := coderdtest.NewWithDatabase(t, nil) + user := coderdtest.CreateFirstUser(t, client) + r := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{ + OrganizationID: user.OrganizationID, + OwnerID: user.UserID, + }).WithAgent().Do() + + // Connect with a non-agent role using ConnectRPCWithRole. + // This is how coder-logstream-kube should connect. + ac := agentsdk.New(client.URL, agentsdk.WithFixedToken(r.AgentToken)) + conn, err := ac.ConnectRPCWithRole(ctx, "logstream-kube") + require.NoError(t, err) + + // Send a log to confirm the RPC connection is functional. + agentAPI := agentproto.NewDRPCAgentClient(conn) + _, err = agentAPI.BatchCreateLogs(ctx, &agentproto.BatchCreateLogsRequest{ + LogSourceId: []byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, + }) + // We don't care about the log source error, just that the + // RPC is functional. + _ = err + + // Close the connection and give the server time to process. + _ = conn.Close() + time.Sleep(100 * time.Millisecond) + + // Verify that connectivity timestamps were never set. + agent, err := db.GetWorkspaceAgentByID(dbauthz.AsSystemRestricted(ctx), r.Agents[0].ID) + require.NoError(t, err) + assert.False(t, agent.FirstConnectedAt.Valid, + "first_connected_at should NOT be set for non-agent role") + assert.False(t, agent.LastConnectedAt.Valid, + "last_connected_at should NOT be set for non-agent role") + assert.False(t, agent.DisconnectedAt.Valid, + "disconnected_at should NOT be set for non-agent role") + }) + + // NOTE: Backward compatibility (empty role) is implicitly tested by + // existing tests like TestWorkspaceAgentReportStats which use + // ConnectRPC() (no role). The server defaults to monitoring when + // the role query parameter is omitted. +} diff --git a/codersdk/agentsdk/agentsdk.go b/codersdk/agentsdk/agentsdk.go index feb8acc560..fae6148e36 100644 --- a/codersdk/agentsdk/agentsdk.go +++ b/codersdk/agentsdk/agentsdk.go @@ -152,7 +152,7 @@ func (c *Client) RewriteDERPMap(derpMap *tailcfg.DERPMap) { // Release Versions from 2.9+ // Deprecated: use ConnectRPC20WithTailnet func (c *Client) ConnectRPC20(ctx context.Context) (proto.DRPCAgentClient20, error) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 0)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 0), "") if err != nil { return nil, err } @@ -165,7 +165,7 @@ func (c *Client) ConnectRPC20(ctx context.Context) (proto.DRPCAgentClient20, err func (c *Client) ConnectRPC20WithTailnet(ctx context.Context) ( proto.DRPCAgentClient20, tailnetproto.DRPCTailnetClient20, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 0)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 0), "") if err != nil { return nil, nil, err } @@ -176,7 +176,7 @@ func (c *Client) ConnectRPC20WithTailnet(ctx context.Context) ( // maximally compatible with Coderd Release Versions from 2.12+ // Deprecated: use ConnectRPC21WithTailnet func (c *Client) ConnectRPC21(ctx context.Context) (proto.DRPCAgentClient21, error) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 1)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 1), "") if err != nil { return nil, err } @@ -188,7 +188,7 @@ func (c *Client) ConnectRPC21(ctx context.Context) (proto.DRPCAgentClient21, err func (c *Client) ConnectRPC21WithTailnet(ctx context.Context) ( proto.DRPCAgentClient21, tailnetproto.DRPCTailnetClient21, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 1)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 1), "") if err != nil { return nil, nil, err } @@ -200,7 +200,7 @@ func (c *Client) ConnectRPC21WithTailnet(ctx context.Context) ( func (c *Client) ConnectRPC22(ctx context.Context) ( proto.DRPCAgentClient22, tailnetproto.DRPCTailnetClient22, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 2)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 2), "") if err != nil { return nil, nil, err } @@ -212,7 +212,7 @@ func (c *Client) ConnectRPC22(ctx context.Context) ( func (c *Client) ConnectRPC23(ctx context.Context) ( proto.DRPCAgentClient23, tailnetproto.DRPCTailnetClient23, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 3)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 3), "") if err != nil { return nil, nil, err } @@ -224,7 +224,7 @@ func (c *Client) ConnectRPC23(ctx context.Context) ( func (c *Client) ConnectRPC24(ctx context.Context) ( proto.DRPCAgentClient24, tailnetproto.DRPCTailnetClient24, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 4)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 4), "") if err != nil { return nil, nil, err } @@ -236,7 +236,7 @@ func (c *Client) ConnectRPC24(ctx context.Context) ( func (c *Client) ConnectRPC25(ctx context.Context) ( proto.DRPCAgentClient25, tailnetproto.DRPCTailnetClient25, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 5)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 5), "") if err != nil { return nil, nil, err } @@ -248,7 +248,7 @@ func (c *Client) ConnectRPC25(ctx context.Context) ( func (c *Client) ConnectRPC26(ctx context.Context) ( proto.DRPCAgentClient26, tailnetproto.DRPCTailnetClient26, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 6)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 6), "") if err != nil { return nil, nil, err } @@ -260,7 +260,7 @@ func (c *Client) ConnectRPC26(ctx context.Context) ( func (c *Client) ConnectRPC27(ctx context.Context) ( proto.DRPCAgentClient27, tailnetproto.DRPCTailnetClient27, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 7)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 7), "") if err != nil { return nil, nil, err } @@ -272,25 +272,53 @@ func (c *Client) ConnectRPC27(ctx context.Context) ( func (c *Client) ConnectRPC28(ctx context.Context) ( proto.DRPCAgentClient28, tailnetproto.DRPCTailnetClient28, error, ) { - conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 8)) + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 8), "") if err != nil { return nil, nil, err } return proto.NewDRPCAgentClient(conn), tailnetproto.NewDRPCTailnetClient(conn), nil } -// ConnectRPC connects to the workspace agent API and tailnet API -func (c *Client) ConnectRPC(ctx context.Context) (drpc.Conn, error) { - return c.connectRPCVersion(ctx, proto.CurrentVersion) +// ConnectRPC28WithRole is like ConnectRPC28 but sends an explicit role +// query parameter to the server. Use "agent" for workspace agents to +// enable connection monitoring. +func (c *Client) ConnectRPC28WithRole(ctx context.Context, role string) ( + proto.DRPCAgentClient28, tailnetproto.DRPCTailnetClient28, error, +) { + conn, err := c.connectRPCVersion(ctx, apiversion.New(2, 8), role) + if err != nil { + return nil, nil, err + } + return proto.NewDRPCAgentClient(conn), tailnetproto.NewDRPCTailnetClient(conn), nil } -func (c *Client) connectRPCVersion(ctx context.Context, version *apiversion.APIVersion) (drpc.Conn, error) { +// ConnectRPC connects to the workspace agent API and tailnet API. +// It does not send a role query parameter, so the server will apply +// its default behavior (currently: enable connection monitoring for +// backward compatibility). Use ConnectRPCWithRole to explicitly +// identify the caller's role. +func (c *Client) ConnectRPC(ctx context.Context) (drpc.Conn, error) { + return c.connectRPCVersion(ctx, proto.CurrentVersion, "") +} + +// ConnectRPCWithRole connects to the workspace agent RPC API with an +// explicit role. The role parameter is sent to the server to identify +// the type of client. Use "agent" for workspace agents to enable +// connection monitoring. +func (c *Client) ConnectRPCWithRole(ctx context.Context, role string) (drpc.Conn, error) { + return c.connectRPCVersion(ctx, proto.CurrentVersion, role) +} + +func (c *Client) connectRPCVersion(ctx context.Context, version *apiversion.APIVersion, role string) (drpc.Conn, error) { rpcURL, err := c.SDK.URL.Parse("/api/v2/workspaceagents/me/rpc") if err != nil { return nil, xerrors.Errorf("parse url: %w", err) } q := rpcURL.Query() q.Add("version", version.String()) + if role != "" { + q.Add("role", role) + } rpcURL.RawQuery = q.Encode() jar, err := cookiejar.New(nil)