feat: add agent-side workspace context sources and Agent API v2.10 PushContextState (#25983)

Adds the agent half of the workspace context sources RFC. The agent now
resolves instruction files, skills, and MCP configs into a typed
`Snapshot`, watches the relevant paths recursively, exposes the source
list over a workspace-agent HTTP API, and pushes each `Snapshot` to
coderd over a new `PushContextState` RPC on Agent API v2.10.

The coderd-side handler is a stub returning `Unimplemented` for now.
Real persistence to `workspace_agent_context`, chatd hydration on dirty
events, and the `KindMCPServer` MCP provider are tracked by
[CODAGT-569](https://linear.app/codercom/issue/CODAGT-569/enable-agent-api-v210-pushcontextstate-bump-currentminor-wire-coderd).
This matches the pattern used for v2.7 `ReportBoundaryLogs` in
[#21293](https://github.com/coder/coder/pull/21293), which bumped the
version and shipped a stub server so the wire and client could iterate
before the persistence layer landed.

## What ships

### agent/agentcontext (new package)

- `Source`, `Resource` (kinds `instruction_file`, `skill`, `mcp_config`,
`mcp_server` plus reserved `plugin`/`hook`/`subagent`/`command`),
`ResourceStatus`, `Snapshot`, `ComputeAggregateHash`.
- `Manager` owns the in-memory source list, performs the initial resolve
synchronously in `NewManager`, runs a re-resolve/watcher loop in `Run`,
exposes
`AddSource`/`RemoveSource`/`Sources`/`HasSource`/`Snapshot`/`SubscribeChanges`/`Resync`/`SeedSources`/`Close`.
- `Resolver` walks scan roots, classifies recognized files, enforces 64
KiB per-resource, 2 MiB aggregate, and 500-resource caps with
`StatusOversize`/`StatusExcluded`/`StatusUnreadable`/`StatusInvalid`
outcomes, skips `node_modules`/`vendor`/etc., validates symlink targets
stay inside the scan root, stamps `SourcePath` on user-derived
resources, and optionally pulls MCP server tool lists via an
`MCPProvider` interface. MCP config resources ship metadata only (size,
hash) so secrets in env blocks never leave the agent.
- `Watcher` is a recursive `fsnotify` wrapper with a 250 ms debounce,
dynamic arming of newly created directories, and an ENOSPC-tolerant
degraded mode that no-ops further syncs until the manager resyncs
explicitly.
- HTTP API for `GET/POST /sources`, `GET/DELETE /sources/{path}`, `POST
/resync` mounted at `/api/v0/context`.
- `Pusher` interface plus `RunPush` goroutine with exponential backoff
capped at 30 s. `DRPCPusher` adapts the generated `DRPCAgentClient210`
to `Pusher` and translates `drpcerr.Unimplemented` to
`ErrPushUnimplemented` so the push loop exits cleanly when talking to
coderd deployments that have not enabled the real handler.

### agent/proto (v2.10)

- New messages `ContextResource`, `PushContextStateRequest`,
`PushContextStateResponse` and the `PushContextState` RPC on `service
Agent`.
- Generated `DRPCAgentClient210` interface and
`codersdk/agentsdk.Client.ConnectRPC210` / `ConnectRPC210WithRole`.
- `tailnet/proto.CurrentMinor` bumped from `9` to `10`.

### Agent wiring

- `agent.Options.Client` declares both v2.9 and v2.10 connectors;
`run()` dials with `ConnectRPC210WithRole`.
- `apiConnRoutineManager` holds a `DRPCAgentClient210`. Existing v2.8
routines keep their narrower `DRPCAgentClient28` signature thanks to
interface embedding.
- `startAgentAPI210` is the v2.10 counterpart to `startAgentAPI` for
routines that need the new client. The push context state routine uses
it.
- A `contextManager` is constructed in `agent.init()`, seeded from the
existing `CODER_AGENT_EXP_*_DIRS` env vars, started in its own goroutine
under `gracefulCtx`, and closed in `agent.Close`.
- `handleManifest` calls `Manager.SeedSources` for sources rooted at the
manifest directory, then `Resync` after `manifest.Swap`, so the snapshot
reflects the workspace working directory immediately instead of waiting
for the next filesystem event.
- HTTP routes mounted at `/api/v0/context` when the manager is up.

### Coderd stub

`coderd/agentapi/context.go` returns `drpcerr.Unimplemented` for
`PushContextState`. The real handler that persists
`workspace_agent_context` rows, hydrates chats, and emits dirty events
lives in CODAGT-569.

## Tests

24 tests across `agent/agentcontext` cover types, paths, resolver
behavior with file caps, skill containers, MCP secret omission, symlink
target validation, the recursive watcher firing on real fsnotify events,
manager source CRUD / `Resync` / `SeedSources` / `Run` lifetime, the
HTTP API, the DRPC adapter, and the push retry / initial-flag /
unimplemented paths. Passes `go test -race -count=2`.

`TestAgent_ContextStatePushed` boots a full agent against
`agenttest.FakeAgentAPI` (which now records `PushContextState` traffic)
and asserts the seeded `AGENTS.md` appears in a snapshot push with
`schema_version = 1`.

<details>
<summary>Notes for reviewers</summary>

- Source CRUD is workspace-agent-token only; coderd is not in the path
for source mutation.
- Per-resource cap 64 KiB, aggregate 2 MiB, count cap 500; resources
past the cap ship with `StatusExcluded` and an empty payload so the
aggregate hash still detects content edits. MCP-emitted resources
enforce both a per-provider count cap and the aggregate byte cap.
- Symlinks inside the scan root are followed; symlinks pointing outside
(or broken) are rejected with `StatusExcluded` so credentials reachable
via a stray symlink stay off the wire.
- The initial push gates `lifecycle = ready` in the eventual full
design. For this PR the `SeedSources` plus `handleManifest`-driven
`Resync` keeps the snapshot fresh; the live push loop ships now and
DRPCPusher translates the coderd `Unimplemented` stub into a clean exit.
- The `PLUGIN`/`HOOK`/`SUBAGENT`/`COMMAND` kinds are reserved in proto
and Go enums but unused; the Claude Code plugin resolver ships in a
follow-up that does not need a schema migration.
- Two follow-ups remain, both tracked by
[CODAGT-569](https://linear.app/codercom/issue/CODAGT-569/enable-agent-api-v210-pushcontextstate-bump-currentminor-wire-coderd):
(1) the chatd-side handler that persists snapshots and dirties chats;
(2) the `coder exp chat context` CLI command set for
`list`/`show`/`add`/`remove`/`refresh`.

</details>

_This PR was authored by Coder Agents on Kyle Carberry's behalf._
This commit is contained in:
Kyle Carberry
2026-06-08 12:08:40 -07:00
committed by GitHub
parent d85211685b
commit cd3692c0c2
30 changed files with 7148 additions and 1046 deletions
+204
View File
@@ -0,0 +1,204 @@
package agentcontext
import (
"context"
"encoding/hex"
"errors"
"net/http"
"net/url"
"strconv"
"github.com/go-chi/chi/v5"
"github.com/coder/coder/v2/coderd/httpapi"
"github.com/coder/coder/v2/codersdk"
)
// SourceResponse is the on-wire representation of a Source.
// Matches the path-only RFC schema; future additions (tags,
// labels) can land additively without breaking clients.
type SourceResponse struct {
Path string `json:"path"`
}
// SourceRequest is the request body for POST /sources.
type SourceRequest struct {
Path string `json:"path"`
}
// SnapshotResource is the on-wire representation of a Resource.
// Payloads are omitted; clients that need the bytes go through
// the drpc PushContextState path.
type SnapshotResource struct {
ID string `json:"id"`
Kind string `json:"kind"`
Source string `json:"source"`
SourcePath string `json:"source_path,omitempty"`
ContentHash string `json:"content_hash"`
SizeBytes uint64 `json:"size_bytes"`
Status string `json:"status"`
Error string `json:"error,omitempty"`
Description string `json:"description,omitempty"`
}
// SnapshotResponse is the on-wire representation of a Snapshot
// returned by the resync endpoint.
type SnapshotResponse struct {
Version uint64 `json:"version"`
SchemaVersion uint64 `json:"schema_version"`
AggregateHash string `json:"aggregate_hash"`
Resources []SnapshotResource `json:"resources"`
PayloadBytes uint64 `json:"payload_bytes"`
SnapshotError string `json:"snapshot_error,omitempty"`
}
// API exposes the Manager over HTTP. The routes match the RFC:
//
// GET /api/v0/context/sources
// POST /api/v0/context/sources { path }
// GET /api/v0/context/sources/{path}
// DELETE /api/v0/context/sources/{path}
// POST /api/v0/context/resync
//
// {path} is URL-encoded canonical path. Callers pass either the
// canonical or original path; the handler canonicalizes before
// matching.
type API struct {
manager *Manager
}
// NewAPI wraps the supplied Manager.
func NewAPI(m *Manager) *API {
return &API{manager: m}
}
// Routes returns the chi handler for /api/v0/context/*. Mount
// it at "/api/v0/context".
func (a *API) Routes() http.Handler {
r := chi.NewRouter()
r.Route("/sources", func(r chi.Router) {
r.Get("/", a.handleListSources)
r.Post("/", a.handleAddSource)
r.Get("/{path}", a.handleGetSource)
r.Delete("/{path}", a.handleRemoveSource)
})
r.Post("/resync", a.handleResync)
return r
}
func (a *API) handleListSources(rw http.ResponseWriter, r *http.Request) {
sources := a.manager.Sources()
out := make([]SourceResponse, 0, len(sources))
for _, s := range sources {
out = append(out, SourceResponse(s))
}
httpapi.Write(r.Context(), rw, http.StatusOK, out)
}
func (a *API) handleAddSource(rw http.ResponseWriter, r *http.Request) {
var req SourceRequest
if !httpapi.Read(r.Context(), rw, r, &req) {
return
}
s, err := a.manager.AddSource(Source(req))
if err != nil {
httpapi.Write(r.Context(), rw, http.StatusBadRequest, codersdk.Response{
Message: "Could not add context source.",
Detail: err.Error(),
})
return
}
httpapi.Write(r.Context(), rw, http.StatusCreated, SourceResponse(s))
}
func (a *API) handleGetSource(rw http.ResponseWriter, r *http.Request) {
raw := chi.URLParam(r, "path")
decoded, err := url.PathUnescape(raw)
if err != nil {
httpapi.Write(r.Context(), rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid context source path.",
Detail: err.Error(),
})
return
}
canonical, ok := a.manager.HasSource(decoded)
if !ok {
httpapi.Write(r.Context(), rw, http.StatusNotFound, codersdk.Response{
Message: "Context source not found.",
Detail: "No source registered for path " + strconv.Quote(decoded) + ".",
})
return
}
httpapi.Write(r.Context(), rw, http.StatusOK, SourceResponse{Path: canonical})
}
func (a *API) handleRemoveSource(rw http.ResponseWriter, r *http.Request) {
raw := chi.URLParam(r, "path")
decoded, err := url.PathUnescape(raw)
if err != nil {
httpapi.Write(r.Context(), rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid context source path.",
Detail: err.Error(),
})
return
}
if err := a.manager.RemoveSource(decoded); err != nil {
if errors.Is(err, ErrSourceNotFound) {
httpapi.Write(r.Context(), rw, http.StatusNotFound, codersdk.Response{
Message: "Context source not found.",
Detail: err.Error(),
})
return
}
httpapi.Write(r.Context(), rw, http.StatusBadRequest, codersdk.Response{
Message: "Could not remove context source.",
Detail: err.Error(),
})
return
}
rw.WriteHeader(http.StatusNoContent)
}
func (a *API) handleResync(rw http.ResponseWriter, r *http.Request) {
snap, err := a.manager.Resync(r.Context())
if err != nil {
status := http.StatusInternalServerError
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
status = http.StatusGatewayTimeout
}
httpapi.Write(r.Context(), rw, status, codersdk.Response{
Message: "Resync failed.",
Detail: err.Error(),
})
return
}
httpapi.Write(r.Context(), rw, http.StatusOK, snapshotResponse(snap))
}
// snapshotResponse converts a Snapshot to its on-wire form for
// the resync endpoint. Payloads are omitted; the per-resource
// payload bytes ship via the drpc PushContextState path.
func snapshotResponse(s Snapshot) SnapshotResponse {
out := SnapshotResponse{
Version: s.Version,
SchemaVersion: s.SchemaVersion,
AggregateHash: hex.EncodeToString(s.AggregateHash[:]),
Resources: make([]SnapshotResource, 0, len(s.Resources)),
PayloadBytes: s.PayloadBytes,
SnapshotError: s.SnapshotError,
}
for _, r := range s.Resources {
out.Resources = append(out.Resources, SnapshotResource{
ID: r.ID,
Kind: r.Kind.String(),
Source: r.Source,
SourcePath: r.SourcePath,
ContentHash: hex.EncodeToString(r.ContentHash[:]),
SizeBytes: r.SizeBytes,
Status: r.Status.String(),
Error: r.Error,
Description: r.Description,
})
}
return out
}
+177
View File
@@ -0,0 +1,177 @@
package agentcontext_test
import (
"bytes"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"net/url"
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/agent/agentcontext"
"github.com/coder/coder/v2/testutil"
)
func newAPITestServer(t *testing.T, opts agentcontext.ManagerOptions) (*httptest.Server, *agentcontext.Manager) {
t.Helper()
m := newTestManager(t, opts)
api := agentcontext.NewAPI(m)
srv := httptest.NewServer(api.Routes())
t.Cleanup(srv.Close)
return srv, m
}
// doRequest issues an HTTP request bounded by testutil.WaitShort
// and returns the status code and response body. The response
// body is closed before doRequest returns.
func doRequest(t *testing.T, method, requrl string, body io.Reader) (int, []byte) {
t.Helper()
ctx := testutil.Context(t, testutil.WaitShort)
req, err := http.NewRequestWithContext(ctx, method, requrl, body)
require.NoError(t, err)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
res, err := http.DefaultClient.Do(req) //nolint:bodyclose // closed below.
require.NoError(t, err)
defer res.Body.Close()
bodyBytes, err := io.ReadAll(res.Body)
require.NoError(t, err)
return res.StatusCode, bodyBytes
}
func TestAPI_ListSourcesEmpty(t *testing.T) {
t.Parallel()
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
status, body := doRequest(t, http.MethodGet, srv.URL+"/sources", nil)
require.Equal(t, http.StatusOK, status)
var got []agentcontext.SourceResponse
require.NoError(t, json.Unmarshal(body, &got))
require.Empty(t, got)
}
func TestAPI_AddAndListSource(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
body, _ := json.Marshal(agentcontext.SourceRequest{Path: src})
status, addBody := doRequest(t, http.MethodPost, srv.URL+"/sources", bytes.NewReader(body))
require.Equal(t, http.StatusCreated, status)
var created agentcontext.SourceResponse
require.NoError(t, json.Unmarshal(addBody, &created))
require.Equal(t, src, created.Path)
// List should show the new source.
listStatus, listBody := doRequest(t, http.MethodGet, srv.URL+"/sources", nil)
require.Equal(t, http.StatusOK, listStatus)
var list []agentcontext.SourceResponse
require.NoError(t, json.Unmarshal(listBody, &list))
require.Len(t, list, 1)
require.Equal(t, src, list[0].Path)
}
func TestAPI_AddSourceRejected(t *testing.T) {
t.Parallel()
wd := t.TempDir()
outside := t.TempDir()
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd},
})
body, _ := json.Marshal(agentcontext.SourceRequest{Path: outside})
status, _ := doRequest(t, http.MethodPost, srv.URL+"/sources", bytes.NewReader(body))
require.Equal(t, http.StatusBadRequest, status)
}
func TestAPI_GetAndDeleteSource(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
srv, m := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
_, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
status, body := doRequest(t, http.MethodGet, srv.URL+"/sources/"+url.PathEscape(src), nil)
require.Equal(t, http.StatusOK, status)
var got agentcontext.SourceResponse
require.NoError(t, json.Unmarshal(body, &got))
require.Equal(t, src, got.Path)
delStatus, _ := doRequest(t, http.MethodDelete, srv.URL+"/sources/"+url.PathEscape(src), nil)
require.Equal(t, http.StatusNoContent, delStatus)
require.Empty(t, m.Sources())
}
func TestAPI_GetSourceNotFound(t *testing.T) {
t.Parallel()
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
status, _ := doRequest(t, http.MethodGet, srv.URL+"/sources/"+url.PathEscape("/never-added"), nil)
require.Equal(t, http.StatusNotFound, status)
}
func TestAPI_DeleteSourceNotFound(t *testing.T) {
t.Parallel()
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
status, _ := doRequest(t, http.MethodDelete, srv.URL+"/sources/"+url.PathEscape("/never-added"), nil)
require.Equal(t, http.StatusNotFound, status)
}
func TestAPI_Resync(t *testing.T) {
t.Parallel()
wd := t.TempDir()
mustWriteFile(t, filepath.Join(wd, "AGENTS.md"), "hello")
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
})
status, body := doRequest(t, http.MethodPost, srv.URL+"/resync", nil)
require.Equal(t, http.StatusOK, status)
var snap agentcontext.SnapshotResponse
require.NoError(t, json.Unmarshal(body, &snap))
require.Equal(t, uint64(1), snap.SchemaVersion)
require.NotEmpty(t, snap.AggregateHash)
require.Len(t, snap.Resources, 1)
require.Equal(t, "instruction_file", snap.Resources[0].Kind)
require.Equal(t, "ok", snap.Resources[0].Status)
}
func TestAPI_AddSourceMalformedBody(t *testing.T) {
t.Parallel()
srv, _ := newAPITestServer(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
status, _ := doRequest(t, http.MethodPost, srv.URL+"/sources", bytes.NewReader([]byte("{not json")))
require.Equal(t, http.StatusBadRequest, status)
}
+32
View File
@@ -0,0 +1,32 @@
package agentcontext
// defaultBuiltinRoots returns the scan roots layered in front
// of any user-added sources. These mirror the paths the legacy
// agentcontextconfig API resolves at every chat hydrate. The
// list is intentionally tolerant of missing entries; the
// resolver silently skips canonicalization failures and
// non-existent paths.
func defaultBuiltinRoots() []string {
return []string{
// User-level Coder config.
"~/.coder",
"~/.coder/skills",
// Claude Code plugin cache, picked up by the plugin
// RFC follow-up. v1 ignores plugin manifests, but
// watching the directory now prevents a surprise
// dirty bit when the resolver eventually classifies
// them.
"~/.claude/plugins/cache",
}
}
// defaultAllowedRoots returns the allow-list applied to runtime
// AddSource calls when ManagerOptions.AllowedRoots is empty.
// The set matches the RFC's authorization section: the home
// directory's Coder and Claude config trees. The Manager
// appends the working directory lazily on every check, which
// picks up the workspace's resolved path even when the manifest
// is loaded after agent init.
func defaultAllowedRoots() []string {
return []string{"~", "~/.coder", "~/.claude"}
}
+24
View File
@@ -0,0 +1,24 @@
// Package agentcontext consolidates the agent-side plumbing that
// resolves, watches, and pushes workspace context (instruction
// files, skills, and MCP configuration) to coderd.
//
// This is the agent half of the design described in
// "RFC: Workspace Context Sources for Coder Agents". It owns:
//
// - User-declared scan roots (Sources) layered on top of
// built-in defaults.
// - A resolver that classifies files under each scan root into
// typed Resources (instruction files, skills, MCP configs,
// MCP servers).
// - A unified recursive fsnotify watcher that signals a
// re-resolve when any recognized file changes.
// - An HTTP API at /api/v0/context/sources for source CRUD
// and /api/v0/context/resync for synchronous push barriers.
// - A Pusher abstraction so the latest Snapshot can be shipped
// to coderd without coupling this package to any particular
// drpc client version.
//
// The package is purely additive: existing agent code paths
// (agent/agentcontextconfig and agent/x/agentmcp) continue to
// operate unchanged.
package agentcontext
+178
View File
@@ -0,0 +1,178 @@
package agentcontext
import (
"context"
"golang.org/x/xerrors"
"google.golang.org/protobuf/types/known/structpb"
"storj.io/drpc/drpcerr"
agentproto "github.com/coder/coder/v2/agent/proto"
)
// DRPCPusher adapts a generated DRPCAgentClient to the
// agentcontext.Pusher interface. The adapter is the only place
// that knows about the wire protobuf types; the rest of the
// package operates on the Go Snapshot/Resource value types.
//
// Use NewDRPCPusher to construct an instance. The pusher's
// behavior is identical to invoking PushContextState directly:
// per-request retries are handled by Manager.RunPush.
type DRPCPusher struct {
client agentproto.DRPCAgentClient210
}
// NewDRPCPusher wraps the supplied drpc client. The client must
// implement the v2.10 Agent API.
func NewDRPCPusher(client agentproto.DRPCAgentClient210) *DRPCPusher {
return &DRPCPusher{client: client}
}
// PushContextState satisfies the Pusher interface.
//
// drpc returns an Unimplemented error when the peer's service
// definition does not include the RPC. The adapter translates
// that into ErrPushUnimplemented so RunPush stops gracefully
// when an old coderd is on the other end.
func (p *DRPCPusher) PushContextState(ctx context.Context, req *PushRequest) (*PushResponse, error) {
if p == nil || p.client == nil {
return nil, xerrors.New("agentcontext: DRPCPusher has no client")
}
resp, err := p.client.PushContextState(ctx, pushRequestToProto(req))
if err != nil {
if drpcerr.Code(err) == drpcerr.Unimplemented {
return nil, ErrPushUnimplemented
}
return nil, err
}
return &PushResponse{Accepted: resp.GetAccepted()}, nil
}
// pushRequestToProto converts the Go push payload to its
// generated protobuf equivalent. The Kind on each Resource
// selects which body variant of the proto oneof is set; a body
// is always set (zero-valued if necessary) so coderd can tell
// the kind even when Status != OK.
func pushRequestToProto(req *PushRequest) *agentproto.PushContextStateRequest {
pb := &agentproto.PushContextStateRequest{
Version: req.Version,
AggregateHash: append([]byte(nil), req.AggregateHash[:]...),
Initial: req.Initial,
SchemaVersion: req.SchemaVersion,
SnapshotError: req.SnapshotError,
Resources: make([]*agentproto.ContextResource, 0, len(req.Resources)),
}
for i := range req.Resources {
r := req.Resources[i]
entry := &agentproto.ContextResource{
Source: r.Source,
ContentHash: append([]byte(nil), r.ContentHash[:]...),
Status: resourceStatusToProto(r.Status),
SizeBytes: r.SizeBytes,
Error: r.Error,
}
setResourceBody(entry, r)
if r.SourcePath != "" {
sp := r.SourcePath
entry.SourcePath = &sp
}
pb.Resources = append(pb.Resources, entry)
}
return pb
}
// setResourceBody picks the proto oneof variant for r's Kind and
// populates the kind-specific fields from r. A body is set even
// when status is not OK so coderd can attribute the failure to a
// known kind. Unknown kinds leave the body unset; the recipient
// can surface that as "kind not recognized".
func setResourceBody(entry *agentproto.ContextResource, r Resource) {
switch r.Kind {
case KindInstructionFile:
entry.Body = &agentproto.ContextResource_InstructionFile{
InstructionFile: &agentproto.InstructionFileBody{
Content: append([]byte(nil), r.Payload...),
},
}
case KindSkill:
entry.Body = &agentproto.ContextResource_Skill{
Skill: &agentproto.SkillMetaBody{
Meta: append([]byte(nil), r.Payload...),
Name: r.Name,
Description: r.Description,
},
}
case KindMCPConfig:
// MCPConfigBody is intentionally empty: secrets in env
// blocks must not leave the agent.
entry.Body = &agentproto.ContextResource_McpConfig{
McpConfig: &agentproto.MCPConfigBody{},
}
case KindMCPServer:
entry.Body = &agentproto.ContextResource_McpServer{
McpServer: &agentproto.MCPServerBody{
ServerName: serverNameOrSource(r),
Description: r.Description,
Tools: mcpToolsToProto(r.Tools),
},
}
}
}
// serverNameOrSource returns r.Name when populated and falls
// back to r.Source so providers that have not yet adopted the
// Name field still produce a usable wire value.
func serverNameOrSource(r Resource) string {
if r.Name != "" {
return r.Name
}
return r.Source
}
// mcpToolsToProto converts the Go MCPTool slice to its wire
// representation. InputSchema is marshaled via structpb.NewStruct;
// schemas that fail to convert are dropped from the wire copy
// (the resource ContentHash still detects the change) and the
// tool ships with InputSchema unset rather than failing the
// whole push.
func mcpToolsToProto(in []MCPTool) []*agentproto.MCPTool {
if len(in) == 0 {
return nil
}
out := make([]*agentproto.MCPTool, 0, len(in))
for _, t := range in {
entry := &agentproto.MCPTool{
Name: t.Name,
Description: t.Description,
}
if len(t.InputSchema) > 0 {
if s, err := structpb.NewStruct(t.InputSchema); err == nil {
entry.InputSchema = s
}
}
out = append(out, entry)
}
return out
}
// resourceStatusToProto maps a ResourceStatus to its proto enum.
func resourceStatusToProto(s ResourceStatus) agentproto.ContextResource_Status {
switch s {
case StatusOK:
return agentproto.ContextResource_OK
case StatusOversize:
return agentproto.ContextResource_OVERSIZE
case StatusUnreadable:
return agentproto.ContextResource_UNREADABLE
case StatusInvalid:
return agentproto.ContextResource_INVALID
case StatusExcluded:
return agentproto.ContextResource_EXCLUDED
default:
return agentproto.ContextResource_STATUS_UNSPECIFIED
}
}
// Ensure DRPCPusher continues to satisfy the Pusher interface
// even if the interface gains methods in the future.
var _ Pusher = (*DRPCPusher)(nil)
+197
View File
@@ -0,0 +1,197 @@
package agentcontext_test
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"storj.io/drpc/drpcerr"
"github.com/coder/coder/v2/agent/agentcontext"
agentproto "github.com/coder/coder/v2/agent/proto"
)
// fakeDRPCClient stubs out the DRPCAgentClient210 surface for
// the parts of the interface the adapter exercises. Only
// PushContextState is implemented; every other method panics
// because the adapter never calls them.
type fakeDRPCClient struct {
agentproto.DRPCAgentClient210
lastReq *agentproto.PushContextStateRequest
resp *agentproto.PushContextStateResponse
err error
}
func (f *fakeDRPCClient) PushContextState(_ context.Context, req *agentproto.PushContextStateRequest) (*agentproto.PushContextStateResponse, error) {
f.lastReq = req
if f.err != nil {
return nil, f.err
}
if f.resp == nil {
return &agentproto.PushContextStateResponse{Accepted: true}, nil
}
return f.resp, nil
}
func TestDRPCPusher_HappyPathSerializesAllFields(t *testing.T) {
t.Parallel()
client := &fakeDRPCClient{}
pusher := agentcontext.NewDRPCPusher(client)
req := &agentcontext.PushRequest{
Version: 7,
AggregateHash: [32]byte{0xaa, 0xbb, 0xcc},
Initial: true,
SchemaVersion: 1,
SnapshotError: "watcher degraded",
Resources: []agentcontext.Resource{
{
ID: "instruction_file:/tmp/AGENTS.md",
Kind: agentcontext.KindInstructionFile,
Source: "/tmp/AGENTS.md",
ContentHash: [32]byte{0x01, 0x02},
Payload: []byte("body"),
SizeBytes: 4,
Status: agentcontext.StatusOK,
Description: "tagline",
SourcePath: "/tmp",
},
{
ID: "skill:/tmp/.agents/skills/foo",
Kind: agentcontext.KindSkill,
Source: "/tmp/.agents/skills/foo",
Status: agentcontext.StatusInvalid,
Error: "bad frontmatter",
SizeBytes: 99,
},
{
ID: "skill:/tmp/.agents/skills/code-review",
Kind: agentcontext.KindSkill,
Source: "/tmp/.agents/skills/code-review",
ContentHash: [32]byte{0x03},
Payload: []byte("---\nname: code-review\n---\nbody\n"),
SizeBytes: 31,
Status: agentcontext.StatusOK,
Name: "code-review",
Description: "Critical review for Go PRs.",
SourcePath: "/tmp",
},
{
ID: "mcp_config:/tmp/.mcp.json",
Kind: agentcontext.KindMCPConfig,
Source: "/tmp/.mcp.json",
ContentHash: [32]byte{0x04},
SizeBytes: 412,
Status: agentcontext.StatusOK,
SourcePath: "/tmp",
},
{
ID: "mcp_server:github",
Kind: agentcontext.KindMCPServer,
Source: "github",
Name: "github",
ContentHash: [32]byte{0x05},
SizeBytes: 138,
Status: agentcontext.StatusOK,
Description: "GitHub MCP server (1 tool)",
SourcePath: "/tmp/.mcp.json",
Tools: []agentcontext.MCPTool{{
Name: "create_issue",
Description: "Create a GitHub issue",
InputSchema: map[string]any{
"type": "object",
"required": []any{"title"},
},
}},
},
},
}
resp, err := pusher.PushContextState(context.Background(), req)
require.NoError(t, err)
require.True(t, resp.Accepted)
pb := client.lastReq
require.NotNil(t, pb)
require.Equal(t, uint64(7), pb.Version)
require.Equal(t, []byte{0xaa, 0xbb, 0xcc, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, pb.AggregateHash)
require.True(t, pb.Initial)
require.Equal(t, uint64(1), pb.SchemaVersion)
require.Equal(t, "watcher degraded", pb.SnapshotError)
require.Len(t, pb.Resources, 5)
// Instruction file: wire-flat fields plus typed body.
instr := pb.Resources[0]
require.Equal(t, "/tmp/AGENTS.md", instr.Source)
require.Equal(t, agentproto.ContextResource_OK, instr.Status)
require.NotNil(t, instr.SourcePath)
require.Equal(t, "/tmp", *instr.SourcePath)
instrBody := instr.GetInstructionFile()
require.NotNil(t, instrBody, "instruction_file body must be set")
require.Equal(t, []byte("body"), instrBody.GetContent())
require.Nil(t, instr.GetSkill())
require.Nil(t, instr.GetMcpConfig())
require.Nil(t, instr.GetMcpServer())
// Skill with INVALID status still has the skill body set so
// coderd can attribute the failure to the correct kind.
invalidSkill := pb.Resources[1]
require.Equal(t, agentproto.ContextResource_INVALID, invalidSkill.Status)
require.Equal(t, "bad frontmatter", invalidSkill.Error)
require.NotNil(t, invalidSkill.GetSkill(), "skill body must be set even when status != OK")
require.Nil(t, invalidSkill.SourcePath, "empty user source must remain optional/nil")
// OK skill: meta + name + description populated.
skill := pb.Resources[2]
skillBody := skill.GetSkill()
require.NotNil(t, skillBody)
require.Equal(t, []byte("---\nname: code-review\n---\nbody\n"), skillBody.GetMeta())
require.Equal(t, "code-review", skillBody.GetName())
require.Equal(t, "Critical review for Go PRs.", skillBody.GetDescription())
// MCP config: body present but empty. SizeBytes / ContentHash
// on the outer resource still detect changes.
mcpCfg := pb.Resources[3]
require.Equal(t, uint64(412), mcpCfg.SizeBytes)
require.NotNil(t, mcpCfg.GetMcpConfig(), "mcp_config body must be set")
// MCP server: structured tool list with input schema.
mcpSrv := pb.Resources[4]
srvBody := mcpSrv.GetMcpServer()
require.NotNil(t, srvBody)
require.Equal(t, "github", srvBody.GetServerName())
require.Equal(t, "GitHub MCP server (1 tool)", srvBody.GetDescription())
require.Len(t, srvBody.GetTools(), 1)
tool := srvBody.GetTools()[0]
require.Equal(t, "create_issue", tool.GetName())
require.Equal(t, "Create a GitHub issue", tool.GetDescription())
require.NotNil(t, tool.GetInputSchema(), "input_schema must be set when supplied")
require.Equal(t, "object", tool.GetInputSchema().GetFields()["type"].GetStringValue())
}
func TestDRPCPusher_UnimplementedTranslated(t *testing.T) {
t.Parallel()
client := &fakeDRPCClient{err: drpcerr.WithCode(drpcerr.WithCode(context.Canceled, 0), drpcerr.Unimplemented)}
pusher := agentcontext.NewDRPCPusher(client)
_, err := pusher.PushContextState(context.Background(), &agentcontext.PushRequest{})
require.ErrorIs(t, err, agentcontext.ErrPushUnimplemented)
}
func TestDRPCPusher_PropagatesOtherErrors(t *testing.T) {
t.Parallel()
want := drpcerr.WithCode(context.DeadlineExceeded, 42)
client := &fakeDRPCClient{err: want}
pusher := agentcontext.NewDRPCPusher(client)
_, err := pusher.PushContextState(context.Background(), &agentcontext.PushRequest{})
require.ErrorIs(t, err, context.DeadlineExceeded)
}
func TestDRPCPusher_NilClientErrors(t *testing.T) {
t.Parallel()
pusher := agentcontext.NewDRPCPusher(nil)
_, err := pusher.PushContextState(context.Background(), &agentcontext.PushRequest{})
require.Error(t, err)
}
+7
View File
@@ -0,0 +1,7 @@
package agentcontext
// ManagerStarted exposes the unexported started() channel for
// use by external _test packages. Production code does not need
// this signal; the agent calls Run synchronously after wiring
// the Manager. Tests use it to coordinate without polling.
func ManagerStarted(m *Manager) <-chan struct{} { return m.started() }
+675
View File
@@ -0,0 +1,675 @@
package agentcontext
import (
"context"
"strings"
"sync"
"time"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/quartz"
)
// CurrentSchemaVersion is the on-wire shape version. Bump
// whenever the resource format changes in a way that requires
// coderd-side awareness.
const CurrentSchemaVersion uint64 = 1
// ManagerOptions configures a Manager. Zero values get sensible
// defaults.
type ManagerOptions struct {
// Logger receives diagnostic messages. Required.
Logger slog.Logger
// Clock is the time source used for the watcher's
// debounce timer. Optional; defaults to quartz.NewReal().
Clock quartz.Clock
// WorkingDir is evaluated on every resolve, mirroring the
// existing agent convention. The result is used as a
// scan root.
WorkingDir func() string
// InitialSources seeds the Manager's source list at boot
// time. Sources from CODER_AGENT_EXP_*_DIRS env vars or
// startup scripts are layered here.
InitialSources []Source
// AllowedRoots restricts which paths may be added as
// sources at runtime. When empty the package falls back
// to [~, ~/.coder, ~/.claude] plus the working directory.
// Tests override this to exercise the validation logic
// directly; production callers leave it unset.
AllowedRoots []string
// Resolver, when non-nil, replaces the default resolver.
// Tests use this to inject MCP providers and tighten
// caps.
Resolver *Resolver
// Debounce overrides the watcher's debounce window.
Debounce time.Duration
// SchemaVersion is the version stamped on each Snapshot.
// Use CurrentSchemaVersion (the default) unless rolling
// out a schema change.
SchemaVersion uint64
}
// Source is a user-declared scan root added to the agent's
// in-memory list via the HTTP API or boot-time env seeding.
// Identity is the canonical absolute path.
type Source struct {
// Path is the canonical absolute path (symlinks resolved,
// ~ expanded). Empty means the zero value.
Path string
}
// Manager orchestrates source CRUD, resolution, watching, and
// Pusher fan-out. Construct with NewManager; start its lifecycle
// goroutines with Run; tear down with Close.
type Manager struct {
logger slog.Logger
clock quartz.Clock
workingDir func() string
allowedRoots []string
resolver *Resolver
debounce time.Duration
schemaVersion uint64
mu sync.Mutex
sources []Source
// sourceIndex maps canonical path -> position in sources
// for O(1) lookups during AddSource / RemoveSource.
sourceIndex map[string]int
// snapshot is the latest result of a resolver pass. It is
// replaced atomically under mu.
snapshot Snapshot
// version monotonically increases per resolve pass.
version uint64
// resolveEpoch increments at the start of every resolver
// pass that drops m.mu around the filesystem walk. Each
// pass captures the epoch it claimed; at publish time it
// compares its captured epoch against the current epoch and
// skips the publish if a newer pass has started, preventing
// an old walk's stale result from overwriting a newer one's
// fresh result at a higher version number.
resolveEpoch uint64
// subscribers receive a non-blocking signal whenever the
// snapshot changes. Subscribers must drain their channel
// promptly; the Manager drops sends to full channels.
subscribers map[chan struct{}]struct{}
// trigger fires when AddSource / RemoveSource / watcher
// observe a change.
trigger chan struct{}
// running tracks Run lifetime.
running bool
closed bool
closedCh chan struct{}
runDoneCh chan struct{}
runStartedCh chan struct{}
watcher *Watcher
}
// NewManager validates options, canonicalizes initial sources,
// performs the first resolver pass synchronously, and returns
// the resulting Manager. Run must be called separately to start
// the watcher and re-resolve goroutine.
func NewManager(opts ManagerOptions) *Manager {
clock := opts.Clock
if clock == nil {
clock = quartz.NewReal()
}
debounce := opts.Debounce
if debounce <= 0 {
debounce = DefaultWatchDebounce
}
schemaVersion := opts.SchemaVersion
if schemaVersion == 0 {
schemaVersion = CurrentSchemaVersion
}
resolver := opts.Resolver
if resolver == nil {
resolver = &Resolver{}
}
m := &Manager{
logger: opts.Logger,
clock: clock,
workingDir: opts.WorkingDir,
allowedRoots: append([]string(nil), opts.AllowedRoots...),
resolver: resolver,
debounce: debounce,
schemaVersion: schemaVersion,
sources: make([]Source, 0),
sourceIndex: make(map[string]int),
subscribers: make(map[chan struct{}]struct{}),
trigger: make(chan struct{}, 1),
closedCh: make(chan struct{}),
runDoneCh: make(chan struct{}),
runStartedCh: make(chan struct{}),
}
for _, s := range opts.InitialSources {
canonical, err := CanonicalizePath(s.Path)
if err != nil {
// Initial sources may not exist yet at boot
// time; log and skip rather than abort the
// agent.
m.logger.Warn(context.Background(),
"skipping invalid initial source",
slog.F("path", s.Path),
slog.Error(err))
continue
}
if _, ok := m.sourceIndex[canonical]; ok {
continue
}
m.sourceIndex[canonical] = len(m.sources)
m.sources = append(m.sources, Source{Path: canonical})
}
// First snapshot is computed eagerly. The push protocol
// requires a snapshot to be present before the agent signals
// lifecycle = ready, so callers can rely on Snapshot() being
// populated immediately after NewManager returns.
m.resolveLocked()
return m
}
// Run starts the watcher and the re-resolve goroutine. Run
// blocks until ctx is canceled or Close is called. It is safe
// to call Run at most once per Manager.
func (m *Manager) Run(ctx context.Context) error {
m.mu.Lock()
if m.running {
m.mu.Unlock()
return xerrors.New("agentcontext: Manager.Run called more than once")
}
if m.closed {
m.mu.Unlock()
return xerrors.New("agentcontext: Manager already closed")
}
m.running = true
close(m.runStartedCh)
m.mu.Unlock()
// Close any early-exit path so Close does not block on
// runDoneCh after Run already set running=true. The deferred
// close runs even when NewWatcher fails.
defer close(m.runDoneCh)
watcher, err := NewWatcher(WatcherOptions{
Logger: m.logger.Named("watcher"),
Clock: m.clock,
Debounce: m.debounce,
MaxDepth: m.resolver.MaxDepth,
OnChange: m.signal,
})
if err != nil {
// NewWatcher already falls back to degraded mode on
// init failure, so an actual error here is
// exceptional.
return xerrors.Errorf("create watcher: %w", err)
}
m.mu.Lock()
m.watcher = watcher
roots := m.scanRootsLocked()
m.mu.Unlock()
watcher.Sync(ctx, roots)
defer watcher.Close()
for {
select {
case <-ctx.Done():
return ctx.Err()
case <-m.closedCh:
return nil
case <-m.trigger:
m.mu.Lock()
roots := m.scanRootsLocked()
m.mu.Unlock()
watcher.Sync(ctx, roots)
m.resolveAndBroadcast(ctx)
}
}
}
// started returns a channel that is closed once Run has
// claimed the running flag. Tests use it to coordinate with
// the watcher loop without polling; a closed channel never
// blocks, so this is safe to call repeatedly.
func (m *Manager) started() <-chan struct{} {
return m.runStartedCh
}
// Close stops the Manager. Close is idempotent; subsequent
// calls block until Run exits.
func (m *Manager) Close() error {
m.mu.Lock()
if m.closed {
running := m.running
m.mu.Unlock()
if running {
<-m.runDoneCh
}
return nil
}
m.closed = true
running := m.running
close(m.closedCh)
m.mu.Unlock()
if running {
<-m.runDoneCh
}
return nil
}
// Sources returns a defensive copy of the current source list.
// The returned slice is safe to mutate.
func (m *Manager) Sources() []Source {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]Source, len(m.sources))
copy(out, m.sources)
return out
}
// HasSource reports whether path matches an existing source
// after canonicalization. Returns the canonical path on
// success.
func (m *Manager) HasSource(path string) (canonical string, ok bool) {
c, err := CanonicalizePath(path)
if err != nil {
return "", false
}
m.mu.Lock()
defer m.mu.Unlock()
_, ok = m.sourceIndex[c]
return c, ok
}
// AddSource adds a new source. The path is canonicalized and
// validated against the AllowedRoots set. AddSource is
// idempotent.
func (m *Manager) AddSource(s Source) (Source, error) {
canonical, err := CanonicalizePath(s.Path)
if err != nil {
return Source{}, xerrors.Errorf("canonicalize: %w", err)
}
if err := ValidateSourcePath(canonical, m.effectiveAllowedRoots()); err != nil {
return Source{}, err
}
m.mu.Lock()
if _, ok := m.sourceIndex[canonical]; ok {
out := m.sources[m.sourceIndex[canonical]]
m.mu.Unlock()
return out, nil
}
m.sourceIndex[canonical] = len(m.sources)
m.sources = append(m.sources, Source{Path: canonical})
m.mu.Unlock()
m.signal()
return Source{Path: canonical}, nil
}
// SeedSources canonicalizes and inserts a batch of trusted
// sources without applying AllowedRoots validation. It is the
// late-binding equivalent of ManagerOptions.InitialSources for
// callers that need the working directory to resolve relative
// paths but only learn the working directory after Run has
// started. Paths that fail canonicalization are silently
// skipped, matching the boot-time seeding contract. SeedSources
// is idempotent: previously seeded canonical paths are
// deduplicated via the existing source index.
//
// AddSource is the correct entry point for untrusted HTTP
// callers; this method exists only for the agent's manifest-
// triggered seeding from CODER_AGENT_EXP_*_DIRS, where the
// template author already authorized the paths.
func (m *Manager) SeedSources(sources []Source) {
if len(sources) == 0 {
return
}
m.mu.Lock()
changed := false
for _, s := range sources {
canonical, err := CanonicalizePath(s.Path)
if err != nil {
m.logger.Warn(context.Background(),
"skipping invalid seeded source",
slog.F("path", s.Path),
slog.Error(err))
continue
}
if _, ok := m.sourceIndex[canonical]; ok {
continue
}
m.sourceIndex[canonical] = len(m.sources)
m.sources = append(m.sources, Source{Path: canonical})
changed = true
}
m.mu.Unlock()
if changed {
m.signal()
}
}
// RemoveSource removes the source matching path. Path is
// canonicalized before matching. Returns ErrSourceNotFound when
// no such source exists or when the path cannot be canonicalized.
func (m *Manager) RemoveSource(path string) error {
canonical, err := CanonicalizePath(path)
if err != nil {
// A path that does not canonicalize cannot match any
// existing source. Mirror HasSource semantics by
// reporting not-found rather than leaking the
// canonicalize error to API callers.
return ErrSourceNotFound
}
m.mu.Lock()
idx, ok := m.sourceIndex[canonical]
if !ok {
m.mu.Unlock()
return ErrSourceNotFound
}
// O(n) compaction is fine for the typical handful of
// user-added sources.
m.sources = append(m.sources[:idx], m.sources[idx+1:]...)
delete(m.sourceIndex, canonical)
for i := idx; i < len(m.sources); i++ {
m.sourceIndex[m.sources[i].Path] = i
}
m.mu.Unlock()
m.signal()
return nil
}
// Snapshot returns the latest Snapshot. The returned value is
// safe to share but shares the same Resources slice as the
// internal state; callers must not mutate it.
func (m *Manager) Snapshot() Snapshot {
m.mu.Lock()
defer m.mu.Unlock()
return m.snapshot
}
// SubscribeChanges returns a buffered channel that receives a
// signal whenever the snapshot changes. The unsubscribe
// callback is safe to call from any goroutine and is
// idempotent.
func (m *Manager) SubscribeChanges() (<-chan struct{}, func()) {
ch := make(chan struct{}, 1)
m.mu.Lock()
m.subscribers[ch] = struct{}{}
m.mu.Unlock()
// OnceFunc returns a closure that runs the underlying
// function at most once. Subsequent invocations are no-ops,
// matching the idempotency contract callers rely on.
unsub := sync.OnceFunc(func() {
m.mu.Lock()
delete(m.subscribers, ch)
m.mu.Unlock()
// Don't close ch: readers may still be in flight.
})
return ch, unsub
}
// Resync forces an immediate re-resolve and returns the new
// Snapshot. Resync is safe to call regardless of whether Run is
// active. Like resolveAndBroadcast, Resync drops the Manager's
// mutex around the resolver pass so concurrent Sources,
// AddSource, RemoveSource, and Snapshot calls do not block on
// filesystem I/O. When the watcher is active, Resync also
// re-arms it so newly added scan roots are observed for
// subsequent edits.
func (m *Manager) Resync(ctx context.Context) (Snapshot, error) {
if ctxErr := ctx.Err(); ctxErr != nil {
return m.Snapshot(), ctxErr
}
m.mu.Lock()
if m.closed {
m.mu.Unlock()
return m.Snapshot(), ErrManagerClosed
}
roots := m.scanRootsLocked()
resolver := m.resolver
watcher := m.watcher
schemaVersion := m.schemaVersion
m.resolveEpoch++
myEpoch := m.resolveEpoch
m.mu.Unlock()
if ctxErr := ctx.Err(); ctxErr != nil {
return m.Snapshot(), ctxErr
}
snap := resolver.ResolveContext(ctx, roots)
if ctxErr := ctx.Err(); ctxErr != nil {
// Cancellation mid-walk yields a partial or empty
// Snapshot whose SnapshotError is set to
// "context canceled". Publishing it would replace
// the live Snapshot with empty resources until the
// next trigger, so bail without touching state.
return m.Snapshot(), ctxErr
}
if snap.SnapshotError == "" && watcher != nil {
if d := watcher.Degraded(); d != "" {
snap.SnapshotError = d
}
}
snap.SchemaVersion = schemaVersion
m.mu.Lock()
if m.closed {
m.mu.Unlock()
return m.Snapshot(), ErrManagerClosed
}
if m.resolveEpoch != myEpoch {
// A newer resolve pass started while this one was
// walking the filesystem. The newer pass's data
// strictly supersedes ours, so skip the publish to
// avoid overwriting a fresher Snapshot at a higher
// version. Return the currently published Snapshot,
// which is at least as fresh as ours. The watcher
// is NOT re-armed: the winning pass already synced
// with the current roots, and replaying our stale
// root set here would drop watches on sources that
// only the newer pass knows about.
published := m.snapshot
m.mu.Unlock()
return published, nil
}
m.version++
snap.Version = m.version
m.snapshot = snap
subs := make([]chan struct{}, 0, len(m.subscribers))
for ch := range m.subscribers {
subs = append(subs, ch)
}
m.mu.Unlock()
if watcher != nil {
watcher.Sync(ctx, roots)
}
// The broadcast is unconditional: Resync waiters that
// triggered the pass without an actual content change
// still need to wake up. Subscribers compare snapshots via
// AggregateHash if they want to filter.
for _, ch := range subs {
select {
case ch <- struct{}{}:
default:
}
}
return snap, nil
}
// signal triggers a re-resolve. Sends are non-blocking; the
// trigger channel has a depth of 1, which coalesces bursts.
func (m *Manager) signal() {
select {
case m.trigger <- struct{}{}:
default:
}
}
// Trigger queues an asynchronous re-resolve. Trigger returns
// immediately; the Run goroutine performs the filesystem walk
// in the background and broadcasts when it finishes. Use
// Trigger when the caller wants the watcher to pick up an
// updated working directory or scan-root set but does not need
// the new Snapshot synchronously. Trigger is a no-op when Run
// has not started or the Manager is closed.
func (m *Manager) Trigger() {
m.signal()
}
// scanRootsLocked returns the list of ScanRoots to feed the
// resolver and watcher. The Manager's mutex must be held.
func (m *Manager) scanRootsLocked() []ScanRoot {
builtinRoots := defaultBuiltinRoots()
out := make([]ScanRoot, 0, 1+len(builtinRoots)+len(m.sources))
if m.workingDir != nil {
if wd := strings.TrimSpace(m.workingDir()); wd != "" {
out = append(out, ScanRoot{Path: wd})
}
}
for _, r := range builtinRoots {
canonical, err := CanonicalizePath(r)
if err != nil {
continue
}
out = append(out, ScanRoot{Path: canonical})
}
for _, s := range m.sources {
out = append(out, ScanRoot{Path: s.Path, UserSource: s.Path})
}
return out
}
// effectiveAllowedRoots returns the AllowedRoots augmented
// with the current working directory. The working directory is
// evaluated on every call so it picks up the workspace's
// resolved path after the agent's manifest finishes loading.
// When AllowedRoots is empty the package falls back to its
// default policy ([~, ~/.coder, ~/.claude]).
func (m *Manager) effectiveAllowedRoots() []string {
var roots []string
if len(m.allowedRoots) > 0 {
roots = append(roots, m.allowedRoots...)
} else {
roots = append(roots, defaultAllowedRoots()...)
}
if m.workingDir != nil {
if wd := strings.TrimSpace(m.workingDir()); wd != "" {
roots = append(roots, wd)
}
}
return roots
}
// resolveAndBroadcast computes a fresh snapshot and notifies
// every subscriber. The broadcast is unconditional: Resync
// waiters that triggered the pass without an actual content
// change still need to wake up. Subscribers compare snapshots
// via AggregateHash if they want to filter.
func (m *Manager) resolveAndBroadcast(ctx context.Context) {
// Snapshot the inputs under the lock, then release it
// before running the resolver. The resolver walks the
// filesystem, reads files, and hashes them; holding
// m.mu across that would block Sources, AddSource,
// RemoveSource, Snapshot, and SubscribeChanges for the
// duration of the pass.
m.mu.Lock()
roots := m.scanRootsLocked()
resolver := m.resolver
watcher := m.watcher
schemaVersion := m.schemaVersion
m.resolveEpoch++
myEpoch := m.resolveEpoch
m.mu.Unlock()
if err := ctx.Err(); err != nil {
return
}
snap := resolver.ResolveContext(ctx, roots)
if err := ctx.Err(); err != nil {
// Cancellation mid-walk yields a partial or empty
// Snapshot. Publishing it would replace the live
// Snapshot with empty resources, so bail without
// touching state. The Run loop's gracefulCtx is
// canceled only at shutdown, but defensive checks
// keep the publish contract uniform with Resync.
return
}
// Surface watcher degradation as a snapshot-level error
// when the resolver did not already emit one.
if snap.SnapshotError == "" && watcher != nil {
if d := watcher.Degraded(); d != "" {
snap.SnapshotError = d
}
}
snap.SchemaVersion = schemaVersion
m.mu.Lock()
if m.resolveEpoch != myEpoch {
// A newer resolve pass started while this one was
// walking the filesystem. Skip the publish so a
// stale-epoch result does not overwrite a fresher
// Snapshot at a higher version number. The newer
// pass will broadcast its own result.
m.mu.Unlock()
return
}
m.version++
snap.Version = m.version
m.snapshot = snap
subs := make([]chan struct{}, 0, len(m.subscribers))
for ch := range m.subscribers {
subs = append(subs, ch)
}
m.mu.Unlock()
for _, ch := range subs {
select {
case ch <- struct{}{}:
default:
}
}
}
// resolveLocked runs the resolver inline while m.mu is held.
// It is used by the synchronous initial resolve in NewManager,
// where there is no concurrent reader. Background re-resolves
// must use resolveAndBroadcast, which drops the lock around
// filesystem I/O.
func (m *Manager) resolveLocked() {
roots := m.scanRootsLocked()
snap := m.resolver.Resolve(roots)
m.version++
snap.Version = m.version
snap.SchemaVersion = m.schemaVersion
// Surface watcher degradation as a snapshot-level error
// when the resolver did not already emit one.
if snap.SnapshotError == "" && m.watcher != nil {
if d := m.watcher.Degraded(); d != "" {
snap.SnapshotError = d
}
}
m.snapshot = snap
}
// ErrSourceNotFound is returned by RemoveSource when the
// requested path is not in the source list.
var ErrSourceNotFound = xerrors.New("source not found")
// ErrManagerClosed is returned by methods called after Close.
var ErrManagerClosed = xerrors.New("agentcontext: manager closed")
+398
View File
@@ -0,0 +1,398 @@
package agentcontext_test
import (
"context"
"os"
"path/filepath"
"runtime"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/agent/agentcontext"
"github.com/coder/coder/v2/testutil"
)
// TestMain points the test binary's HOME (and USERPROFILE on
// Windows) at a fresh empty directory before any test runs.
// The package's built-in scan roots (~/.coder,
// ~/.coder/skills, ~/.claude/plugins/cache) canonicalize
// against this directory, so they resolve to non-existent
// paths and the resolver silently skips them. Without this,
// running the tests on a developer host pulls real Coder and
// Claude config files into snapshots and breaks every
// Len(Resources, N) assertion.
func TestMain(m *testing.M) {
home, err := os.MkdirTemp("", "agentcontext-test-home-")
if err != nil {
panic(err)
}
if err := os.Setenv("HOME", home); err != nil {
panic(err)
}
if runtime.GOOS == "windows" {
if err := os.Setenv("USERPROFILE", home); err != nil {
panic(err)
}
}
code := m.Run()
_ = os.RemoveAll(home)
os.Exit(code)
}
func newTestManager(t *testing.T, opts agentcontext.ManagerOptions) *agentcontext.Manager {
t.Helper()
opts.Logger = testutil.Logger(t).Named("agentcontext-test")
m := agentcontext.NewManager(opts)
t.Cleanup(func() { _ = m.Close() })
return m
}
func TestManager_InitialSnapshotIsPopulated(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), "boot snapshot")
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return dir },
})
snap := m.Snapshot()
require.Equal(t, uint64(1), snap.Version)
require.Equal(t, agentcontext.CurrentSchemaVersion, snap.SchemaVersion)
require.Len(t, snap.Resources, 1)
}
func TestManager_AddSourceTriggersResolve(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
mustWriteFile(t, filepath.Join(src, "AGENTS.md"), "from source")
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
ctx := testutil.Context(t, testutil.WaitLong)
go func() { _ = m.Run(ctx) }()
t.Cleanup(func() { _ = m.Close() })
// Subscribe before mutating so we observe the broadcast.
ch, unsub := m.SubscribeChanges()
defer unsub()
added, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
require.Equal(t, src, added.Path)
select {
case <-ch:
case <-time.After(testutil.WaitShort):
t.Fatalf("expected a change broadcast after AddSource")
}
snap := m.Snapshot()
require.Greater(t, snap.Version, uint64(1))
found := false
for _, r := range snap.Resources {
if r.Kind == agentcontext.KindInstructionFile && r.SourcePath == src {
found = true
}
}
require.True(t, found, "expected AGENTS.md attributed to the user source")
}
func TestManager_AddSourceRejectsOutsideAllowedRoots(t *testing.T) {
t.Parallel()
wd := t.TempDir()
outside := t.TempDir()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd},
})
_, err := m.AddSource(agentcontext.Source{Path: outside})
require.Error(t, err)
}
// TestManager_AddSourceAcceptsLateWorkingDir mirrors the agent's
// real boot order: AllowedRoots is configured before the
// manifest provides the workspace working directory. The Manager
// must consult WorkingDir on every check so paths under the
// resolved working dir validate once the manifest lands.
func TestManager_AddSourceAcceptsLateWorkingDir(t *testing.T) {
t.Parallel()
wd := t.TempDir()
var resolved atomic.Pointer[string]
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string {
if p := resolved.Load(); p != nil {
return *p
}
return ""
},
AllowedRoots: []string{"/never-used-home"},
})
// Before the manifest "loads", workingDir is empty; sources
// under wd must be rejected.
_, err := m.AddSource(agentcontext.Source{Path: wd})
require.Error(t, err)
// After the manifest "loads", workingDir resolves and the
// same path validates without restarting the Manager.
resolved.Store(&wd)
_, err = m.AddSource(agentcontext.Source{Path: wd})
require.NoError(t, err)
}
func TestManager_AddSourceIsIdempotent(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
added1, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
added2, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
require.Equal(t, added1.Path, added2.Path)
sources := m.Sources()
require.Len(t, sources, 1)
}
func TestManager_RemoveSource(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
_, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
require.NoError(t, m.RemoveSource(src))
require.Empty(t, m.Sources())
err = m.RemoveSource(src)
require.ErrorIs(t, err, agentcontext.ErrSourceNotFound)
}
func TestManager_HasSource(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
canonical, ok := m.HasSource(src)
require.False(t, ok)
require.Equal(t, src, canonical)
_, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
canonical, ok = m.HasSource(src)
require.True(t, ok)
require.Equal(t, src, canonical)
}
func TestManager_ResyncReturnsLatestSnapshot(t *testing.T) {
t.Parallel()
wd := t.TempDir()
mustWriteFile(t, filepath.Join(wd, "AGENTS.md"), "first")
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
})
ctx := testutil.Context(t, testutil.WaitLong)
runDone := make(chan struct{})
go func() {
defer close(runDone)
_ = m.Run(ctx)
}()
t.Cleanup(func() {
_ = m.Close()
<-runDone
})
// Mutate AGENTS.md and call Resync. The returned
// snapshot must reflect the new content.
require.NoError(t, os.WriteFile(filepath.Join(wd, "AGENTS.md"), []byte("second content edit"), 0o600))
snap, err := m.Resync(ctx)
require.NoError(t, err)
require.Len(t, snap.Resources, 1)
require.Equal(t, "second content edit", string(snap.Resources[0].Payload))
}
// TestManager_ResyncCanceledKeepsLiveSnapshot guards CRF-44:
// a context cancellation mid-walk must not replace the live
// Snapshot with an empty one. Resync returns the existing
// Snapshot and ctx.Err() instead of publishing a stub.
func TestManager_ResyncCanceledKeepsLiveSnapshot(t *testing.T) {
t.Parallel()
wd := t.TempDir()
mustWriteFile(t, filepath.Join(wd, "AGENTS.md"), "live content")
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
})
// Capture the live snapshot the Manager populated at
// construction time.
live := m.Snapshot()
require.Len(t, live.Resources, 1)
require.Equal(t, "live content", string(live.Resources[0].Payload))
// Cancel the context before calling Resync so
// ResolveContext observes the cancellation.
ctx, cancel := context.WithCancel(context.Background())
cancel()
snap, err := m.Resync(ctx)
require.ErrorIs(t, err, context.Canceled)
// The returned snapshot must still expose the live
// resources, not an empty result from the canceled walk.
require.Len(t, snap.Resources, 1)
require.Equal(t, "live content", string(snap.Resources[0].Payload))
// The next Snapshot call must also return live content;
// no stub was published.
after := m.Snapshot()
require.Equal(t, live.Version, after.Version)
require.Len(t, after.Resources, 1)
}
func TestManager_InitialSourcesSeeded(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
mustWriteFile(t, filepath.Join(src, "AGENTS.md"), "from initial")
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
InitialSources: []agentcontext.Source{{Path: src}},
})
sources := m.Sources()
require.Len(t, sources, 1)
require.Equal(t, src, sources[0].Path)
snap := m.Snapshot()
require.Len(t, snap.Resources, 1)
require.Equal(t, src, snap.Resources[0].SourcePath)
}
// TestManager_SeedSourcesLateBindsAfterManifest models the
// agent's behavior when CODER_AGENT_EXP_*_DIRS contains a
// relative path that cannot resolve until the manifest's
// working directory lands. SeedSources must adopt the
// previously-unresolvable path, bypass AllowedRoots
// validation, and trigger a re-resolve.
func TestManager_SeedSourcesLateBindsAfterManifest(t *testing.T) {
t.Parallel()
wd := t.TempDir()
late := t.TempDir()
mustWriteFile(t, filepath.Join(late, "AGENTS.md"), "late binding")
// AllowedRoots intentionally omits `late` so AddSource
// would reject it. SeedSources must accept it anyway,
// since the path comes from the trusted template config.
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd},
})
require.Empty(t, m.Sources())
m.SeedSources([]agentcontext.Source{{Path: late}})
sources := m.Sources()
require.Len(t, sources, 1)
require.Equal(t, late, sources[0].Path)
snap, err := m.Resync(testutil.Context(t, testutil.WaitShort))
require.NoError(t, err)
require.Len(t, snap.Resources, 1)
require.Equal(t, late, snap.Resources[0].SourcePath)
}
func TestManager_CloseIsIdempotent(t *testing.T) {
t.Parallel()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
require.NoError(t, m.Close())
require.NoError(t, m.Close())
}
func TestManager_RunOnce(t *testing.T) {
t.Parallel()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
defer cancel()
go func() { _ = m.Run(ctx) }()
// Wait for Run to claim the running flag, then verify the
// second call rejects with a deterministic error rather than
// racing the scheduler.
select {
case <-agentcontext.ManagerStarted(m):
case <-ctx.Done():
t.Fatalf("manager never started: %v", ctx.Err())
}
err := m.Run(ctx)
require.Error(t, err)
require.Contains(t, err.Error(), "more than once")
cancel()
_ = m.Close()
}
func TestManager_SubscribeBroadcastOnChange(t *testing.T) {
t.Parallel()
wd := t.TempDir()
src := t.TempDir()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return wd },
AllowedRoots: []string{wd, src},
})
ctx := testutil.Context(t, testutil.WaitLong)
go func() { _ = m.Run(ctx) }()
ch, unsub := m.SubscribeChanges()
defer unsub()
_, err := m.AddSource(agentcontext.Source{Path: src})
require.NoError(t, err)
select {
case <-ch:
case <-time.After(testutil.WaitShort):
t.Fatal("expected subscriber to be notified")
}
}
+30
View File
@@ -0,0 +1,30 @@
package agentcontext
// MCPProvider supplies the live MCP server portion of a
// snapshot. Implementations typically wrap an existing MCP
// manager (e.g. agent/x/agentmcp.Manager) and translate each
// server's tool list into a KindMCPServer resource.
//
// The interface is intentionally minimal so the existing MCP
// lifecycle code can be reused without refactoring; a follow-up
// change absorbs the lifecycle into this package.
type MCPProvider interface {
// MCPResources returns one Resource per MCP server known
// to the provider. Each Resource must:
//
// - Have Kind == KindMCPServer.
// - Use the server name as Source.
// - Set Name to the server name (matches Source today;
// reserved for the case where a future provider scheme
// decouples them).
// - Populate ContentHash over a canonical encoding of the
// server name plus the tool list (proto Tools field)
// so any tool-set change flips the dirty bit.
// - Carry a Description summarizing the server.
// - Populate Tools with the structured tool list; Payload
// is unused for this kind and should be left empty.
//
// Implementations should never block; the resolver calls
// this on every re-resolve.
MCPResources() []Resource
}
+121
View File
@@ -0,0 +1,121 @@
package agentcontext
import (
"os"
"path/filepath"
"strings"
"golang.org/x/xerrors"
)
// CanonicalizePath produces the canonical form of a user-
// supplied path. The result is absolute, has ~ expanded, has
// path-traversal segments collapsed, and has symlinks resolved
// when the target exists. The path is left lexically clean if
// it does not yet exist (so adding a not-yet-created directory
// remains possible).
//
// CanonicalizePath returns the original input when it is empty.
func CanonicalizePath(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return "", xerrors.New("path is empty")
}
// Expand ~ and ~/ prefixes against the current user's home
// directory. Other ~user forms are not supported on
// purpose; the agent runs as a known user.
if raw == "~" || strings.HasPrefix(raw, "~/") {
home, err := os.UserHomeDir()
if err != nil {
return "", xerrors.Errorf("expand home dir: %w", err)
}
if raw == "~" {
raw = home
} else {
raw = filepath.Join(home, raw[2:])
}
}
if !filepath.IsAbs(raw) {
// Fail closed: relative paths could mean different
// things depending on the agent's working directory at
// add-time, so require the caller to absolutize first.
return "", xerrors.Errorf("path %q is not absolute", raw)
}
cleaned := filepath.Clean(raw)
if resolved, err := filepath.EvalSymlinks(cleaned); err == nil {
return resolved, nil
}
return cleaned, nil
}
// ValidateSourcePath enforces the path-validation rules from
// the RFC's Authorization section. It rejects:
//
// - Paths containing ".." segments after expansion.
// - Paths resolving outside the supplied allowedRoots, unless
// allowedRoots is empty (which disables the check).
//
// allowedRoots are canonicalized lazily; missing roots are
// silently skipped so a workspace with no $HOME does not break
// validation for project-relative roots.
func ValidateSourcePath(canonical string, allowedRoots []string) error {
if canonical == "" {
return xerrors.New("path is empty")
}
// filepath.Clean drops "." but leaves ".." when no parent
// is available. Reject defensively.
for _, part := range strings.Split(canonical, string(os.PathSeparator)) {
if part == ".." {
return xerrors.Errorf("path %q contains parent traversal segments", canonical)
}
}
if len(allowedRoots) == 0 {
return nil
}
// Build canonical, deduplicated allowed roots. Missing
// roots (e.g. an unconfigured ~/.claude/) are skipped.
roots := make([]string, 0, len(allowedRoots))
seen := make(map[string]struct{}, len(allowedRoots))
for _, raw := range allowedRoots {
c, err := CanonicalizePath(raw)
if err != nil {
continue
}
if _, ok := seen[c]; ok {
continue
}
seen[c] = struct{}{}
roots = append(roots, c)
}
if len(roots) == 0 {
// All configured roots were invalid; treat as "deny
// everything" so misconfiguration fails closed.
return xerrors.Errorf("path %q is not inside any allowed root", canonical)
}
for _, root := range roots {
if pathHasPrefix(canonical, root) {
return nil
}
}
return xerrors.Errorf("path %q is not inside any allowed root", canonical)
}
// pathHasPrefix reports whether path is equal to or a
// descendant of prefix. Both arguments must already be clean,
// absolute paths.
func pathHasPrefix(path, prefix string) bool {
if path == prefix {
return true
}
withSep := prefix
if !strings.HasSuffix(withSep, string(os.PathSeparator)) {
withSep += string(os.PathSeparator)
}
return strings.HasPrefix(path, withSep)
}
+142
View File
@@ -0,0 +1,142 @@
package agentcontext_test
import (
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/agent/agentcontext"
)
// switchHomeEnv overrides the platform-specific environment
// variable consulted by os.UserHomeDir for the duration of the
// test. Windows reads USERPROFILE; Linux and macOS read HOME.
func switchHomeEnv(t *testing.T, dir string) {
t.Helper()
switch runtime.GOOS {
case "windows":
t.Setenv("USERPROFILE", dir)
default:
t.Setenv("HOME", dir)
}
}
func TestCanonicalizePath_AbsoluteCleansAndResolves(t *testing.T) {
t.Parallel()
dir := t.TempDir()
got, err := agentcontext.CanonicalizePath(filepath.Join(dir, "a", "..", "b"))
require.NoError(t, err)
// Path does not exist; EvalSymlinks fails. Result is
// lexically cleaned: filepath.Clean drops the "..".
require.Equal(t, filepath.Join(dir, "b"), got)
}
func TestCanonicalizePath_RelativeRejected(t *testing.T) {
t.Parallel()
_, err := agentcontext.CanonicalizePath("relative/path")
require.Error(t, err)
}
//nolint:paralleltest,tparallel // Uses t.Setenv.
func TestCanonicalizePath_TildeExpansion(t *testing.T) {
home := t.TempDir()
switchHomeEnv(t, home)
got, err := agentcontext.CanonicalizePath("~/.coder")
require.NoError(t, err)
require.Equal(t, filepath.Join(home, ".coder"), got)
}
//nolint:paralleltest,tparallel // Uses t.Setenv.
func TestCanonicalizePath_BareTildeExpandsToHome(t *testing.T) {
home := t.TempDir()
switchHomeEnv(t, home)
got, err := agentcontext.CanonicalizePath("~")
require.NoError(t, err)
// Canonicalize the same home path through the function under
// test so the comparison handles platform-specific behavior of
// EvalSymlinks (Windows can fail to resolve directories that
// Linux/macOS resolve cleanly).
want, err := agentcontext.CanonicalizePath(home)
require.NoError(t, err)
require.Equal(t, want, got)
}
func TestCanonicalizePath_FollowsSymlinks(t *testing.T) {
t.Parallel()
if runtime.GOOS == "windows" {
t.Skip("os.Symlink requires developer mode or admin on Windows")
}
dir := t.TempDir()
realDir := filepath.Join(dir, "real")
link := filepath.Join(dir, "link")
require.NoError(t, os.MkdirAll(realDir, 0o755))
require.NoError(t, os.Symlink(realDir, link))
got, err := agentcontext.CanonicalizePath(link)
require.NoError(t, err)
// On macOS the temp dir is itself symlinked; both realDir and got
// pass through the same EvalSymlinks so they line up.
want, err := filepath.EvalSymlinks(realDir)
require.NoError(t, err)
require.Equal(t, want, got)
}
func TestValidateSourcePath_RejectsParentSegments(t *testing.T) {
t.Parallel()
root := t.TempDir()
// Build /a/../b underneath a real allowed root so the path is
// absolute on every platform. Validation must still reject the
// embedded ".." segment before it ever touches allowedRoots.
bad := filepath.Join(root, "a") + string(os.PathSeparator) + ".." + string(os.PathSeparator) + "b"
err := agentcontext.ValidateSourcePath(bad, []string{root})
require.Error(t, err)
require.Contains(t, err.Error(), "parent traversal")
}
func TestValidateSourcePath_AllowsInsideRoot(t *testing.T) {
t.Parallel()
dir := t.TempDir()
child := filepath.Join(dir, "child")
require.NoError(t, os.MkdirAll(child, 0o755))
require.NoError(t, agentcontext.ValidateSourcePath(child, []string{dir}))
require.NoError(t, agentcontext.ValidateSourcePath(dir, []string{dir}))
}
func TestValidateSourcePath_RejectsOutsideRoot(t *testing.T) {
t.Parallel()
root := t.TempDir()
other := t.TempDir()
err := agentcontext.ValidateSourcePath(other, []string{root})
require.Error(t, err)
require.Contains(t, err.Error(), "not inside any allowed root")
}
func TestValidateSourcePath_EmptyAllowedRootsBypass(t *testing.T) {
t.Parallel()
require.NoError(t, agentcontext.ValidateSourcePath("/anywhere", nil))
}
func TestValidateSourcePath_InvalidRootsFailClosed(t *testing.T) {
t.Parallel()
// All allowed roots are relative and therefore invalid;
// validation must fail closed.
err := agentcontext.ValidateSourcePath("/anywhere", []string{"relative-only"})
require.Error(t, err)
}
func TestValidateSourcePath_PathPrefixIsPathAware(t *testing.T) {
t.Parallel()
// "/a-prefix" is not inside "/a", even though it starts
// with the same bytes.
dir := t.TempDir()
sibling := strings.TrimRight(dir, string(os.PathSeparator)) + "-sibling"
require.NoError(t, os.MkdirAll(sibling, 0o755))
t.Cleanup(func() { _ = os.RemoveAll(sibling) })
err := agentcontext.ValidateSourcePath(sibling, []string{dir})
require.Error(t, err)
}
+202
View File
@@ -0,0 +1,202 @@
package agentcontext
import (
"context"
"errors"
"time"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/quartz"
)
// PushRequest is the wire-format-independent payload the
// Manager hands to a Pusher. It mirrors the protobuf
// PushContextStateRequest message reserved in the RFC.
//
// Keeping the shape in plain Go lets this package compile
// without bumping the drpc proto version. The follow-up
// integration change can add a thin adapter that converts
// PushRequest to proto and back.
type PushRequest struct {
Version uint64
AggregateHash [32]byte
Resources []Resource
Initial bool
SchemaVersion uint64
SnapshotError string
}
// PushResponse is the wire-format-independent return value of
// a push.
type PushResponse struct {
Accepted bool
}
// Pusher delivers snapshots to coderd. Concrete implementations
// wrap a drpc client (Agent API v2.10 and later) or, in tests,
// a recording in-memory fake.
//
// PushContextState must respect ctx cancellation; the Manager
// retries on transient errors with backoff but stops on
// ErrPushUnimplemented.
type Pusher interface {
PushContextState(ctx context.Context, req *PushRequest) (*PushResponse, error)
}
// ErrPushUnimplemented signals that the coderd peer does not
// implement PushContextState. RunPush stops pushing for the
// remainder of the connection.
var ErrPushUnimplemented = xerrors.New("agentcontext: PushContextState unimplemented")
// Default backoff timings for pushWithRetry. Exposed as named
// constants (rather than inline literals) so godoc shows them
// and a second push loop, if it ever appears, can reuse them.
const (
DefaultPushInitialBackoff = 250 * time.Millisecond
DefaultPushMaxBackoff = 30 * time.Second
)
// PushOptions parameterizes RunPush.
type PushOptions struct {
// Logger receives push success/failure diagnostics.
Logger slog.Logger
// InitialBackoff is the wait before the first retry.
// Default 250ms.
InitialBackoff time.Duration
// MaxBackoff caps the retry wait. Default 30s.
MaxBackoff time.Duration
// Clock is the time source for retry backoffs. Optional;
// defaults to the Manager's clock so tests can trap waits
// with quartz instead of real sleeps.
Clock quartz.Clock
}
// RunPush ships the current snapshot to the Pusher, then ships
// every subsequent snapshot whenever the Manager broadcasts a
// change. RunPush returns when ctx is canceled, when the
// Manager is closed, or when the Pusher signals
// ErrPushUnimplemented.
//
// The first push is always sent with Initial=true so coderd can
// distinguish a fresh boot from a drift event.
func (m *Manager) RunPush(ctx context.Context, p Pusher, opts PushOptions) error {
if p == nil {
return xerrors.New("agentcontext: Pusher is required")
}
logger := opts.Logger
initialBackoff := opts.InitialBackoff
if initialBackoff <= 0 {
initialBackoff = DefaultPushInitialBackoff
}
maxBackoff := opts.MaxBackoff
if maxBackoff <= 0 {
maxBackoff = DefaultPushMaxBackoff
}
clock := opts.Clock
if clock == nil {
clock = m.clock
}
changes, unsub := m.SubscribeChanges()
defer unsub()
// First push uses the snapshot computed by NewManager.
initial := true
for {
snap := m.Snapshot()
req := snapshotToPushRequest(snap, initial)
err := pushWithRetry(ctx, p, req, initialBackoff, maxBackoff, clock, logger)
switch {
case err == nil:
initial = false
case errors.Is(err, ErrPushUnimplemented):
logger.Warn(ctx, "coderd peer does not implement PushContextState; stopping")
return nil
case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded):
return ctx.Err()
default:
// Should be unreachable: pushWithRetry only
// returns terminal errors. Log and continue.
logger.Warn(ctx, "push terminated with non-retried error", slog.Error(err))
}
select {
case <-ctx.Done():
return ctx.Err()
case <-m.closedCh:
return nil
case <-changes:
// Shutdown comes from closedCh or ctx; the
// subscriber channel is never closed by
// SubscribeChanges.
}
}
}
// pushWithRetry retries transient errors with exponential
// backoff capped at maxBackoff. The retry loop exits when:
//
// - ctx is canceled (returns ctx.Err()).
// - The Pusher returns nil (success).
// - The Pusher returns ErrPushUnimplemented (propagated).
func pushWithRetry(
ctx context.Context,
p Pusher,
req *PushRequest,
initialBackoff, maxBackoff time.Duration,
clock quartz.Clock,
logger slog.Logger,
) error {
backoff := initialBackoff
for {
resp, err := p.PushContextState(ctx, req)
if err == nil {
if resp != nil && !resp.Accepted {
// Out-of-order or replayed push. Do not
// retry; the next change will redeliver
// the snapshot with a higher version.
logger.Debug(ctx, "push rejected, awaiting next change",
slog.F("version", req.Version))
}
return nil
}
if errors.Is(err, ErrPushUnimplemented) {
return err
}
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return err
}
logger.Warn(ctx, "push failed, retrying",
slog.F("version", req.Version),
slog.F("backoff", backoff),
slog.Error(err))
timer := clock.NewTimer(backoff)
select {
case <-ctx.Done():
timer.Stop()
return ctx.Err()
case <-timer.C:
}
backoff *= 2
if backoff > maxBackoff {
backoff = maxBackoff
}
}
}
// snapshotToPushRequest copies the Snapshot into the wire
// representation. The Resources slice is reused; callers must
// not mutate it.
func snapshotToPushRequest(s Snapshot, initial bool) *PushRequest {
return &PushRequest{
Version: s.Version,
AggregateHash: s.AggregateHash,
Resources: s.Resources,
Initial: initial,
SchemaVersion: s.SchemaVersion,
SnapshotError: s.SnapshotError,
}
}
+305
View File
@@ -0,0 +1,305 @@
package agentcontext_test
import (
"context"
"os"
"path/filepath"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"github.com/coder/coder/v2/agent/agentcontext"
"github.com/coder/coder/v2/testutil"
"github.com/coder/quartz"
)
// fakePusher records every push and lets the test control the
// returned response and error.
type fakePusher struct {
mu sync.Mutex
requests []*agentcontext.PushRequest
resp *agentcontext.PushResponse
err error
// errOnce is non-nil to simulate a single transient
// failure followed by success.
errOnce error
signal chan struct{}
}
func newFakePusher() *fakePusher {
return &fakePusher{
resp: &agentcontext.PushResponse{Accepted: true},
signal: make(chan struct{}, 16),
}
}
func (p *fakePusher) PushContextState(_ context.Context, req *agentcontext.PushRequest) (*agentcontext.PushResponse, error) {
p.mu.Lock()
defer p.mu.Unlock()
p.requests = append(p.requests, req)
if p.errOnce != nil {
err := p.errOnce
p.errOnce = nil
return nil, err
}
select {
case p.signal <- struct{}{}:
default:
}
return p.resp, p.err
}
func (p *fakePusher) snapshot() []*agentcontext.PushRequest {
p.mu.Lock()
defer p.mu.Unlock()
out := make([]*agentcontext.PushRequest, len(p.requests))
copy(out, p.requests)
return out
}
func TestRunPush_FirstPushIsInitial(t *testing.T) {
t.Parallel()
dir := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v1"), 0o600))
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return dir },
})
p := newFakePusher()
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
defer cancel()
pushDone := make(chan error, 1)
go func() {
pushDone <- m.RunPush(ctx, p, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
})
}()
// Wait for the first push.
select {
case <-p.signal:
case <-time.After(testutil.WaitShort):
t.Fatalf("expected initial push")
}
requests := p.snapshot()
require.Len(t, requests, 1)
require.True(t, requests[0].Initial, "first push must be initial")
require.Equal(t, uint64(1), requests[0].Version)
cancel()
require.ErrorIs(t, <-pushDone, context.Canceled)
}
func TestRunPush_SubsequentPushOnChange(t *testing.T) {
t.Parallel()
dir := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v1"), 0o600))
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return dir },
})
p := newFakePusher()
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
defer cancel()
pushDone := make(chan error, 1)
go func() {
pushDone <- m.RunPush(ctx, p, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
})
}()
// Initial push.
<-p.signal
// Trigger a resync via Resync.
require.NoError(t, os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v2"), 0o600))
_, err := m.Resync(ctx)
require.NoError(t, err)
// Second push.
select {
case <-p.signal:
case <-time.After(testutil.WaitShort):
t.Fatalf("expected second push after resync")
}
requests := p.snapshot()
require.GreaterOrEqual(t, len(requests), 2)
require.False(t, requests[1].Initial, "subsequent pushes must not be Initial")
require.NotEqual(t, requests[0].AggregateHash, requests[1].AggregateHash,
"second push must reflect the v2 content, not a duplicate of the first snapshot")
require.Greater(t, requests[1].Version, requests[0].Version,
"version must advance between snapshots")
cancel()
require.ErrorIs(t, <-pushDone, context.Canceled)
}
func TestRunPush_StopsOnUnimplemented(t *testing.T) {
t.Parallel()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
p := newFakePusher()
p.err = agentcontext.ErrPushUnimplemented
ctx := testutil.Context(t, testutil.WaitShort)
err := m.RunPush(ctx, p, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
})
require.NoError(t, err, "Unimplemented must stop the loop cleanly")
}
func TestRunPush_RetriesTransientError(t *testing.T) {
t.Parallel()
mClock := quartz.NewMock(t)
trap := mClock.Trap().NewTimer()
defer trap.Close()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
p := newFakePusher()
p.errOnce = xerrors.New("transient")
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
defer cancel()
pushDone := make(chan error, 1)
go func() {
pushDone <- m.RunPush(ctx, p, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
InitialBackoff: time.Second,
Clock: mClock,
})
}()
// First push hits transient and arms the retry timer. Wait for
// the timer creation, then advance the clock past the backoff.
call := trap.MustWait(ctx)
call.MustRelease(ctx)
mClock.Advance(time.Second).MustWait(ctx)
select {
case <-p.signal:
case <-time.After(testutil.WaitShort):
t.Fatalf("expected push after transient error")
}
require.GreaterOrEqual(t, len(p.snapshot()), 2)
cancel()
<-pushDone
}
// TestRunPush_ClosesOnManagerClose verifies that calling
// Manager.Close terminates an in-flight RunPush even when the
// caller's context is still live. Without this guarantee the
// agent shutdown would leak a push goroutine until the
// surrounding ctx expired.
func TestRunPush_ClosesOnManagerClose(t *testing.T) {
t.Parallel()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
p := newFakePusher()
ctx := testutil.Context(t, testutil.WaitShort)
done := make(chan error, 1)
go func() {
done <- m.RunPush(ctx, p, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
})
}()
// Wait for the initial push so the loop is parked on the
// change channel, then close the Manager and assert that
// RunPush returns promptly with a nil error.
select {
case <-p.signal:
case <-ctx.Done():
t.Fatalf("initial push never landed: %v", ctx.Err())
}
require.NoError(t, m.Close())
select {
case err := <-done:
require.NoError(t, err)
case <-ctx.Done():
t.Fatalf("RunPush did not return after Manager.Close: %v", ctx.Err())
}
}
// TestRunPush_RejectedResponseProceeds verifies the contract
// that an Accepted=false response is not retried: pushWithRetry
// returns success and RunPush parks on the next change instead
// of re-sending the same snapshot. A regression that added
// retry-on-reject logic would loop here and fail the test.
func TestRunPush_RejectedResponseProceeds(t *testing.T) {
t.Parallel()
dir := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v1"), 0o600))
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return dir },
})
p := newFakePusher()
p.resp = &agentcontext.PushResponse{Accepted: false}
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
defer cancel()
pushDone := make(chan error, 1)
go func() {
pushDone <- m.RunPush(ctx, p, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
})
}()
// Initial push delivered and accepted=false; loop must park
// on changes, not retry the same payload.
select {
case <-p.signal:
case <-ctx.Done():
t.Fatalf("initial push never landed: %v", ctx.Err())
}
// Trigger a content change so a second push lands. Without
// the change, the loop should remain parked.
require.NoError(t, os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v2"), 0o600))
_, err := m.Resync(ctx)
require.NoError(t, err)
select {
case <-p.signal:
case <-ctx.Done():
t.Fatalf("second push never landed after change: %v", ctx.Err())
}
requests := p.snapshot()
require.GreaterOrEqual(t, len(requests), 2,
"exactly one push per snapshot; rejection must not double-fire")
require.NotEqual(t, requests[0].AggregateHash, requests[1].AggregateHash)
cancel()
require.ErrorIs(t, <-pushDone, context.Canceled)
}
func TestRunPush_NilPusherErrors(t *testing.T) {
t.Parallel()
m := newTestManager(t, agentcontext.ManagerOptions{
WorkingDir: func() string { return t.TempDir() },
})
err := m.RunPush(context.Background(), nil, agentcontext.PushOptions{
Logger: testutil.Logger(t).Named("push"),
})
require.Error(t, err)
}
+990
View File
@@ -0,0 +1,990 @@
package agentcontext
import (
"context"
"crypto/sha256"
"errors"
"fmt"
"io"
"io/fs"
"math"
"os"
"path/filepath"
"regexp"
"slices"
"strconv"
"strings"
"github.com/coder/coder/v2/codersdk/workspacesdk"
)
// Default caps. Copied from the RFC. The Manager exposes
// overrides via Options.
const (
// DefaultMaxResourceBytes is the per-resource payload cap.
// Resources whose payload exceeds this size are emitted
// with Status == StatusOversize and an empty Payload.
DefaultMaxResourceBytes = 64 * 1024
// DefaultMaxSnapshotBytes is the aggregate payload cap.
// Resources past this cap are emitted with Status ==
// StatusExcluded.
DefaultMaxSnapshotBytes = 2 * 1024 * 1024
// DefaultMaxResources is the resource count cap. Resources
// past this cap are emitted with Status == StatusExcluded.
DefaultMaxResources = 500
// DefaultMaxScanDepth bounds how deep the recursive walk
// descends from each scan root. The default avoids runaway
// scans in node_modules / vendor / .git trees while still
// covering realistic monorepo layouts.
DefaultMaxScanDepth = 8
)
// File-name conventions recognized by the v1 resolver.
var (
// instructionFileNames are picked up from any scan root.
// Matching is case-insensitive on the basename.
instructionFileNames = []string{
"AGENTS.md",
"CLAUDE.md",
".cursorrules",
}
// mcpConfigFileName is recognized at any depth under a
// scan root.
mcpConfigFileName = ".mcp.json"
// skillMetaFileName is the file inside a skill directory
// that carries the skill front-matter.
skillMetaFileName = "SKILL.md"
)
// skipDirNames are directory basenames that the recursive walk
// never descends into. The list mirrors what most language
// tool-chains treat as opaque.
var skipDirNames = map[string]struct{}{
".git": {},
".hg": {},
".svn": {},
"node_modules": {},
"vendor": {},
"target": {},
"dist": {},
"build": {},
".venv": {},
"__pycache__": {},
}
// recognizedInstructionFile reports whether name is one of the
// instruction-file conventions, case-insensitively.
func recognizedInstructionFile(name string) bool {
for _, candidate := range instructionFileNames {
if strings.EqualFold(name, candidate) {
return true
}
}
return false
}
// Resolver walks one or more scan roots and produces a snapshot
// of every recognized resource it finds. The Resolver is
// stateless; the Manager owns the scan-root list and orchestrates
// successive resolves.
type Resolver struct {
// MaxResourceBytes caps the per-resource payload size. Use
// DefaultMaxResourceBytes if zero.
MaxResourceBytes uint64
// MaxSnapshotBytes caps the aggregate payload size. Use
// DefaultMaxSnapshotBytes if zero.
MaxSnapshotBytes uint64
// MaxResources caps the resource count. Use
// DefaultMaxResources if zero.
MaxResources int
// MaxDepth caps the directory walk depth. Use
// DefaultMaxScanDepth if zero.
MaxDepth int
// MCP, when non-nil, is consulted after the filesystem
// pass and contributes any KindMCPServer resources for
// live MCP servers.
MCP MCPProvider
}
// ScanRoot describes a single directory or file the resolver
// should examine.
type ScanRoot struct {
// Path is the absolute path. Symlinks should already be
// resolved.
Path string
// UserSource is the canonical source path the user
// declared, when this root came from a user-added Source.
// Empty for built-in roots.
UserSource string
}
// Resolve walks the supplied scan roots and returns a Snapshot.
// The version and schemaVersion fields are stamped by the
// caller; Resolve fills everything else. Resolve is the
// non-cancellable convenience wrapper around ResolveContext
// using context.Background.
func (r *Resolver) Resolve(roots []ScanRoot) Snapshot {
return r.ResolveContext(context.Background(), roots)
}
// ResolveContext is the cancellable variant of Resolve. The
// context is checked between scan roots so callers can bail out
// of a long pass without waiting for the current root's walk to
// finish. Cancellation never partially populates the returned
// Snapshot: a canceled context returns an empty Snapshot with
// SnapshotError set to the context error.
func (r *Resolver) ResolveContext(ctx context.Context, roots []ScanRoot) Snapshot {
res := r.normalize()
resources, snapErrs := res.walk(ctx, roots)
if err := ctx.Err(); err != nil {
return Snapshot{SnapshotError: err.Error()}
}
resources, totalBytes := res.applyCaps(resources)
// Append MCP server resources after the filesystem caps
// are applied so a runaway MCP server cannot crowd out
// instruction files.
if r.MCP != nil {
mcp := r.MCP.MCPResources()
startIdx := len(resources)
resources = append(resources, mcp...)
// MCP resources may push the aggregate over the
// count or byte cap. Apply both, picking up
// where applyCaps left off.
resources, snapErrs = res.applyMCPCaps(resources, startIdx, totalBytes, snapErrs)
}
// Deterministic order by ID for stable IDs and hashes.
slices.SortFunc(resources, func(a, b Resource) int {
return strings.Compare(a.ID, b.ID)
})
var payloadBytes uint64
for _, r := range resources {
payloadBytes += uint64(len(r.Payload))
}
hash := ComputeAggregateHash(resources)
snap := Snapshot{
Resources: resources,
AggregateHash: hash,
PayloadBytes: payloadBytes,
}
if len(snapErrs) > 0 {
// Pick the most severe single error. Today every
// snapshot-level problem is "warning equivalent" so
// the first one wins; the design reserves the field
// for a singular message.
snap.SnapshotError = snapErrs[0]
}
return snap
}
func (r *Resolver) normalize() *Resolver {
out := *r
if out.MaxResourceBytes == 0 {
out.MaxResourceBytes = DefaultMaxResourceBytes
}
if out.MaxSnapshotBytes == 0 {
out.MaxSnapshotBytes = DefaultMaxSnapshotBytes
}
if out.MaxResources == 0 {
out.MaxResources = DefaultMaxResources
}
if out.MaxDepth == 0 {
out.MaxDepth = DefaultMaxScanDepth
}
return &out
}
// walk traverses every scan root and produces an unordered
// resource list. Aggregate caps are applied separately. The ctx
// is checked between roots so callers can bail out promptly.
func (r *Resolver) walk(ctx context.Context, roots []ScanRoot) (resources []Resource, snapErrs []string) {
// Dedup roots by canonical path. The first occurrence
// wins so user-added roots that overlap with a built-in
// root attribute resources to the built-in.
seenRoot := make(map[string]struct{}, len(roots))
dedup := make([]ScanRoot, 0, len(roots))
for _, root := range roots {
if root.Path == "" {
continue
}
if _, ok := seenRoot[root.Path]; ok {
continue
}
seenRoot[root.Path] = struct{}{}
dedup = append(dedup, root)
}
// Deduplicate resources across roots by ID. Without this,
// a built-in root and a user root that both cover the
// same project tree would double-count AGENTS.md.
seenID := make(map[string]struct{})
for _, root := range dedup {
if err := ctx.Err(); err != nil {
return nil, []string{err.Error()}
}
info, err := os.Stat(root.Path)
if err != nil {
// Missing roots silently fall through. The user
// either added a path that does not exist yet or
// removed it later. The watcher will surface
// re-creation as a change event.
continue
}
if !info.IsDir() {
// Single-file roots are classified directly.
if res, ok := r.classifyFile(root.Path, root.Path, info, root.UserSource); ok {
if _, dup := seenID[res.ID]; !dup {
seenID[res.ID] = struct{}{}
resources = append(resources, res)
}
}
continue
}
walkErr := r.walkDir(ctx, root, &resources, seenID)
if walkErr != nil {
snapErrs = append(snapErrs, fmt.Sprintf("walk %q: %s", root.Path, walkErr))
}
}
return resources, snapErrs
}
// walkDir performs the recursive descent for a single scan
// directory. It honors r.MaxDepth and skipDirNames. The ctx is
// checked inside the WalkDir callback so cancellation
// terminates the walk even mid-root.
func (r *Resolver) walkDir(ctx context.Context, root ScanRoot, out *[]Resource, seenID map[string]struct{}) error {
rootDepth := strings.Count(filepath.Clean(root.Path), string(os.PathSeparator))
maxDepth := rootDepth + r.MaxDepth
return filepath.WalkDir(root.Path, func(path string, d fs.DirEntry, err error) error {
if ctxErr := ctx.Err(); ctxErr != nil {
return ctxErr
}
if err != nil {
// Surface the error as Unreadable when we can
// associate it with a single recognized file;
// otherwise let the walk continue.
if d != nil && !d.IsDir() {
kind, recognized := kindFromFilename(d.Name())
if recognized {
res := Resource{
ID: resourceID(kind, path),
Kind: kind,
Source: path,
SizeBytes: 0,
Status: StatusUnreadable,
Error: err.Error(),
SourcePath: root.UserSource,
}
if _, dup := seenID[res.ID]; !dup {
seenID[res.ID] = struct{}{}
*out = append(*out, res)
}
}
}
if errors.Is(err, fs.ErrPermission) {
// Permission errors on a directory: skip the
// subtree but continue walking siblings.
if d != nil && d.IsDir() {
return fs.SkipDir
}
}
return nil
}
if d.IsDir() {
if strings.Count(path, string(os.PathSeparator)) > maxDepth {
return fs.SkipDir
}
if _, skip := skipDirNames[d.Name()]; skip && path != root.Path {
return fs.SkipDir
}
// If we are entering a "skills container"
// directory (".agents/skills", "~/.coder/skills",
// "plugins/<plugin>/skills"), eagerly emit skill
// resources for its immediate subdirectories.
if isSkillsContainer(path) {
r.emitSkillsFromContainer(path, root, out, seenID)
}
return nil
}
// Regular file.
info, statErr := d.Info()
if statErr != nil {
return nil
}
if res, ok := r.classifyFile(root.Path, path, info, root.UserSource); ok {
if _, dup := seenID[res.ID]; dup {
return nil
}
seenID[res.ID] = struct{}{}
*out = append(*out, res)
}
return nil
})
}
// kindFromFilename maps a file basename to its ResourceKind.
// recognized=false when the name matches no convention.
func kindFromFilename(name string) (kind ResourceKind, recognized bool) {
switch {
case recognizedInstructionFile(name):
return KindInstructionFile, true
case name == mcpConfigFileName:
return KindMCPConfig, true
case name == skillMetaFileName:
return KindSkill, true
default:
return 0, false
}
}
// resolveReadTarget produces the path and FileInfo that should
// be used to read the resource. When the input is not a
// symlink the original path and info are returned unchanged.
// When it is a symlink the target is resolved and validated
// against scanRoot so a malicious AGENTS.md ->
// ~/.ssh/id_rsa cannot exfiltrate files outside the
// contributing scan root.
//
// codex follows symlinks unconditionally because it trusts the
// local user's filesystem. Coder workspaces may execute
// templates and repositories that the agent operator did not
// author, so the resolver follows symlinks only within the
// scan-root boundary. Symlinks whose targets escape the
// boundary are emitted as StatusInvalid; broken symlinks and
// non-regular targets are emitted as StatusUnreadable.
func resolveReadTarget(path string, info fs.FileInfo, scanRoot string) (readPath string, readInfo fs.FileInfo, ok bool, status ResourceStatus, errMsg string) {
if info.Mode()&fs.ModeSymlink == 0 {
return path, info, true, StatusOK, ""
}
target, err := filepath.EvalSymlinks(path)
if err != nil {
return "", nil, false, StatusUnreadable, fmt.Sprintf("symlink resolve: %v", err)
}
rootClean := filepath.Clean(scanRoot)
if !pathHasPrefix(target, rootClean) {
return "", nil, false, StatusInvalid, fmt.Sprintf("symlink target %q escapes scan root %q", target, scanRoot)
}
tgtInfo, err := os.Stat(target)
if err != nil {
return "", nil, false, StatusUnreadable, err.Error()
}
if !tgtInfo.Mode().IsRegular() {
return "", nil, false, StatusInvalid, fmt.Sprintf("symlink target %q is not a regular file", target)
}
return target, tgtInfo, true, StatusOK, ""
}
// classifyFile inspects a single file path and produces a
// Resource when the basename matches a recognized convention.
func (r *Resolver) classifyFile(scanRoot, path string, info fs.FileInfo, userSource string) (Resource, bool) {
name := info.Name()
switch {
case recognizedInstructionFile(name):
return r.readInstructionFile(scanRoot, path, info, userSource), true
case name == mcpConfigFileName:
return r.readMCPConfig(scanRoot, path, info, userSource), true
case name == skillMetaFileName:
// SKILL.md outside a skills container is still a
// valid skill if its parent directory name matches
// the front-matter name. emitSkillsFromContainer
// already handles the common case; here we cover
// "user adds a single SKILL.md file as a source".
res, ok := r.readSkillMeta(scanRoot, path, info, userSource)
return res, ok
default:
return Resource{}, false
}
}
// readInstructionFile reads an instruction file and produces a
// KindInstructionFile resource. The file is read into memory
// with the per-resource cap applied.
//
// The bytes are returned verbatim. The legacy code path in
// agentcontextconfig/api.go strips HTML comments and invisible
// Unicode before serving instruction-file contents to chat; the
// equivalent sanitization for this pipeline lives in the
// follow-up chatd integration that consumes Snapshot.Resources.
// Until that lands, downstream consumers that render these
// payloads must sanitize themselves.
func (r *Resolver) readInstructionFile(scanRoot, path string, info fs.FileInfo, userSource string) Resource {
res := r.readFileResource(KindInstructionFile, scanRoot, path, info, userSource)
if res.Status == StatusOK {
res.Description = firstLine(string(res.Payload))
}
return res
}
// readMCPConfig reads a .mcp.json file and produces a
// KindMCPConfig resource carrying only path metadata and a
// content hash.
//
// .mcp.json fragments frequently embed secret-bearing fields
// (Env tokens, Authorization headers). The resolver hashes the
// file for change detection but intentionally does not ship
// the bytes; the live MCP server's tool list arrives via the
// MCPProvider as a KindMCPServer resource, which is what
// downstream consumers actually need.
func (r *Resolver) readMCPConfig(scanRoot, path string, info fs.FileInfo, userSource string) Resource {
res := Resource{
ID: resourceID(KindMCPConfig, path),
Kind: KindMCPConfig,
Source: path,
SizeBytes: safeUint64(info.Size()),
SourcePath: userSource,
}
readPath, readInfo, ok, status, errMsg := resolveReadTarget(path, info, scanRoot)
if !ok {
res.Status = status
res.Error = errMsg
return res
}
res.SizeBytes = safeUint64(readInfo.Size())
if safeUint64(readInfo.Size()) > r.MaxResourceBytes {
res.Status = StatusOversize
res.Error = fmt.Sprintf("file size %d exceeds per-resource cap of %d bytes", readInfo.Size(), r.MaxResourceBytes)
if data, err := readFileCapped(readPath, safeInt64(r.MaxResourceBytes)); err == nil {
res.ContentHash = sha256.Sum256(data)
}
return res
}
data, err := os.ReadFile(readPath)
if err != nil {
res.Status = StatusUnreadable
res.Error = err.Error()
return res
}
res.ContentHash = sha256.Sum256(data)
return res
}
// readFileResource is the shared plumbing for kinds whose only
// difference is the enum stamped on the Resource: build the
// Resource header, enforce the per-resource size cap, read the
// file, hash it, attach the bytes. Callers add kind-specific
// post-processing (e.g. firstLine for instruction files) by
// inspecting Status==StatusOK.
func (r *Resolver) readFileResource(kind ResourceKind, scanRoot, path string, info fs.FileInfo, userSource string) Resource {
res := Resource{
ID: resourceID(kind, path),
Kind: kind,
Source: path,
SizeBytes: safeUint64(info.Size()),
SourcePath: userSource,
}
readPath, readInfo, ok, status, errMsg := resolveReadTarget(path, info, scanRoot)
if !ok {
res.Status = status
res.Error = errMsg
return res
}
res.SizeBytes = safeUint64(readInfo.Size())
if safeUint64(readInfo.Size()) > r.MaxResourceBytes {
res.Status = StatusOversize
res.Error = fmt.Sprintf("file size %d exceeds per-resource cap of %d bytes", readInfo.Size(), r.MaxResourceBytes)
// Still hash the (capped) content so a fix is
// detectable.
if data, err := readFileCapped(readPath, safeInt64(r.MaxResourceBytes)); err == nil {
res.ContentHash = sha256.Sum256(data)
}
return res
}
data, err := os.ReadFile(readPath)
if err != nil {
res.Status = StatusUnreadable
res.Error = err.Error()
return res
}
res.Payload = data
res.ContentHash = sha256.Sum256(data)
return res
}
// readSkillMeta reads a SKILL.md file, parses its front-matter,
// and emits a KindSkill resource. The name encoded in the
// front-matter must match the parent directory's basename to
// be considered valid; otherwise Status is StatusInvalid.
func (r *Resolver) readSkillMeta(scanRoot, path string, info fs.FileInfo, userSource string) (Resource, bool) {
parent := filepath.Base(filepath.Dir(path))
res := Resource{
ID: resourceID(KindSkill, filepath.Dir(path)),
Kind: KindSkill,
Source: filepath.Dir(path),
SizeBytes: safeUint64(info.Size()),
SourcePath: userSource,
}
readPath, readInfo, ok, status, errMsg := resolveReadTarget(path, info, scanRoot)
if !ok {
res.Status = status
res.Error = errMsg
return res, true
}
res.SizeBytes = safeUint64(readInfo.Size())
if safeUint64(readInfo.Size()) > r.MaxResourceBytes {
res.Status = StatusOversize
res.Error = fmt.Sprintf("file size %d exceeds per-resource cap of %d bytes", readInfo.Size(), r.MaxResourceBytes)
// Hash the (capped) prefix so an edit that keeps
// the file oversize still shifts the aggregate
// hash and triggers a re-broadcast. Mirrors the
// behavior in readFileResource.
if data, err := readFileCapped(readPath, safeInt64(r.MaxResourceBytes)); err == nil {
res.ContentHash = sha256.Sum256(data)
}
return res, true
}
data, err := os.ReadFile(readPath)
if err != nil {
res.Status = StatusUnreadable
res.Error = err.Error()
return res, true
}
res.ContentHash = sha256.Sum256(data)
name, description, _, err := workspacesdk.ParseSkillFrontmatter(string(data))
if err != nil {
res.Status = StatusInvalid
res.Error = err.Error()
return res, true
}
if name != parent {
res.Status = StatusInvalid
res.Error = fmt.Sprintf("front-matter name %q does not match directory %q", name, parent)
return res, true
}
if !workspacesdk.SkillNamePattern.MatchString(name) {
res.Status = StatusInvalid
res.Error = fmt.Sprintf("skill name %q is not kebab-case", name)
return res, true
}
res.Description = description
res.Name = name
res.Payload = data
return res, true
}
// emitSkillsFromContainer scans the immediate children of a
// recognized skills-container directory and emits one Skill
// resource per subdirectory whose SKILL.md parses cleanly.
func (r *Resolver) emitSkillsFromContainer(container string, root ScanRoot, out *[]Resource, seenID map[string]struct{}) {
entries, err := os.ReadDir(container)
if err != nil {
return
}
for _, e := range entries {
if !e.IsDir() {
continue
}
meta := filepath.Join(container, e.Name(), skillMetaFileName)
// Lstat (not Stat) so a symlinked SKILL.md is
// detected and routed through resolveReadTarget,
// which enforces the scan-root boundary.
info, err := os.Lstat(meta)
if err != nil {
continue
}
res, ok := r.readSkillMeta(root.Path, meta, info, root.UserSource)
if !ok {
continue
}
if _, dup := seenID[res.ID]; dup {
continue
}
seenID[res.ID] = struct{}{}
*out = append(*out, res)
}
}
// applyCaps enforces the resource-count cap and aggregate
// payload cap. Resources past either cap have their Status set
// to StatusExcluded and their Payload cleared. The returned
// byte total is the sum of surviving payloads, so callers that
// append additional resources (e.g. MCP server tool lists) can
// apply the same byte cap to the appended slice.
func (r *Resolver) applyCaps(resources []Resource) ([]Resource, uint64) {
// Stable sort by (Kind asc, Source asc) so excluded
// resources are deterministic.
slices.SortStableFunc(resources, func(a, b Resource) int {
if a.Kind != b.Kind {
return int(a.Kind) - int(b.Kind)
}
return strings.Compare(a.Source, b.Source)
})
var total uint64
for i := range resources {
if i >= r.MaxResources {
resources[i] = excluded(resources[i],
fmt.Sprintf("dropped to fit %d-resource snapshot count cap", r.MaxResources))
continue
}
if resources[i].Status != StatusOK {
continue
}
size := uint64(len(resources[i].Payload))
if total+size > r.MaxSnapshotBytes {
resources[i] = excluded(resources[i],
fmt.Sprintf("dropped to fit %d-byte aggregate cap", r.MaxSnapshotBytes))
continue
}
total += size
}
return resources, total
}
// applyMCPCaps enforces both the count cap and the remaining
// aggregate byte cap on MCP resources appended after
// applyCaps. startIdx is the first index of the appended tail.
// priorBytes is the sum of payload bytes already committed by
// the filesystem pass; MCP resources whose payloads would push
// the running total past MaxSnapshotBytes are stamped
// StatusExcluded. Without this guard a provider returning one
// large KindMCPServer payload would exceed the aggregate cap
// with StatusOK, breaking the contract in
// DefaultMaxSnapshotBytes.
func (r *Resolver) applyMCPCaps(resources []Resource, startIdx int, priorBytes uint64, snapErrs []string) ([]Resource, []string) {
total := priorBytes
countCapHit := false
byteCapHit := false
for i := startIdx; i < len(resources); i++ {
if i >= r.MaxResources {
resources[i] = excluded(resources[i],
fmt.Sprintf("dropped to fit %d-resource snapshot count cap", r.MaxResources))
countCapHit = true
continue
}
if resources[i].Status != StatusOK {
continue
}
size := uint64(len(resources[i].Payload))
if total+size > r.MaxSnapshotBytes {
resources[i] = excluded(resources[i],
fmt.Sprintf("dropped to fit %d-byte aggregate cap", r.MaxSnapshotBytes))
byteCapHit = true
continue
}
total += size
}
if countCapHit {
snapErrs = append(snapErrs, fmt.Sprintf("snapshot exceeds %d-resource count cap", r.MaxResources))
}
if byteCapHit {
snapErrs = append(snapErrs, fmt.Sprintf("snapshot exceeds %d-byte aggregate cap", r.MaxSnapshotBytes))
}
return resources, snapErrs
}
// excluded mutates and returns the supplied resource with the
// StatusExcluded outcome.
func excluded(r Resource, reason string) Resource {
r.Status = StatusExcluded
r.Error = reason
r.Payload = nil
return r
}
// isSkillsContainer reports whether dir is a recognized skills
// container directory whose immediate children carry SKILL.md
// files. Both bare "skills" and nested "<parent>/skills"
// directories qualify (e.g. ".agents/skills",
// "plugins/foo/skills").
func isSkillsContainer(dir string) bool {
return filepath.Base(dir) == "skills"
}
// resourceID builds a stable resource ID. Kind plus canonical
// source path is enough; sources never collide across kinds for
// v1 because each kind owns a distinct file-name pattern.
func resourceID(kind ResourceKind, source string) string {
return kind.String() + ":" + source
}
// readFileCapped reads up to maxBytes from path. It returns the
// truncated payload on success.
func readFileCapped(path string, maxBytes int64) ([]byte, error) {
f, err := os.Open(path)
if err != nil {
return nil, err
}
defer f.Close()
return io.ReadAll(io.LimitReader(f, maxBytes))
}
// firstLine returns the first non-empty trimmed line of s, used
// as a short description fallback.
func firstLine(s string) string {
for line := range strings.SplitSeq(s, "\n") {
line = strings.TrimSpace(line)
if line == "" {
continue
}
// Strip leading markdown heading markers for prettier
// descriptions.
return strings.TrimSpace(headingPrefixRegex.ReplaceAllString(line, ""))
}
return ""
}
var headingPrefixRegex = regexp.MustCompile(`^#+\s*`)
// safeUint64 converts a non-negative int64 to uint64. Negative
// inputs are clamped to 0, which is safe for the size-tracking
// fields that use it; a negative os.FileInfo size is pathological
// and never indicates real content.
func safeUint64(n int64) uint64 {
if n < 0 {
return 0
}
return uint64(n)
}
// safeInt64 converts a uint64 to int64, clamping to math.MaxInt64
// when the input would overflow. The caps configured on the
// resolver never approach 2^63 bytes, so the clamp only guards
// against pathological caller input.
func safeInt64(n uint64) int64 {
if n > math.MaxInt64 {
return math.MaxInt64
}
return int64(n)
}
// ResourceKind describes the category of a resolved context
// resource. The values mirror the proto ContextResource.Kind
// enum reserved in the RFC; future kinds (PLUGIN, HOOK,
// SUBAGENT, COMMAND) are defined here so callers can switch
// exhaustively, but no v1 resolver emits them.
type ResourceKind int
const (
KindUnspecified ResourceKind = iota
// KindInstructionFile covers AGENTS.md, CLAUDE.md,
// .cursorrules, and similar plain-text rule files that
// inject content into the model prompt.
KindInstructionFile
// KindSkill is a directory containing SKILL.md and any
// supporting files. Only the meta file is read at
// resolve time; bodies are fetched on demand.
KindSkill
// KindMCPConfig is a .mcp.json fragment declaring one or
// more MCP servers.
KindMCPConfig
// KindMCPServer is a live MCP server's resolved tool list,
// populated by an MCPProvider after the server has been
// connected.
KindMCPServer
// KindPlugin is reserved for Claude Code plugin manifests.
// Not emitted by v1.
KindPlugin
// KindHook is reserved for plugin hooks. Not emitted by v1.
KindHook
// KindSubagent is reserved for plugin-declared subagents.
// Not emitted by v1.
KindSubagent
// KindCommand is reserved for plugin slash commands.
// Not emitted by v1.
KindCommand
)
// String returns the lower-snake-case name used in IDs and
// metrics. Unknown values stringify to "unknown".
func (k ResourceKind) String() string {
switch k {
case KindInstructionFile:
return "instruction_file"
case KindSkill:
return "skill"
case KindMCPConfig:
return "mcp_config"
case KindMCPServer:
return "mcp_server"
case KindPlugin:
return "plugin"
case KindHook:
return "hook"
case KindSubagent:
return "subagent"
case KindCommand:
return "command"
default:
return "unknown"
}
}
// ResourceStatus describes whether a resource was successfully
// read and whether its payload survived the per-resource and
// aggregate caps.
//
// Note: these iota ordinals do NOT match the proto
// ContextResource.Status ordinals one-to-one. The proto enum
// reserves 0 for STATUS_UNSPECIFIED and shifts every value by
// one, so the conversion in resourceStatusToProto cannot be
// replaced with a direct int cast. ResourceKind, by contrast,
// does align with its proto counterpart.
type ResourceStatus int
const (
// StatusOK indicates the payload was populated.
StatusOK ResourceStatus = iota
// StatusOversize indicates the resource exceeded the
// per-resource size cap; payload is omitted.
StatusOversize
// StatusUnreadable indicates an IO error reading the
// resource (permission denied, broken symlink, etc.).
StatusUnreadable
// StatusInvalid indicates the resource was structurally
// malformed (bad JSON, missing front-matter, etc.).
StatusInvalid
// StatusExcluded indicates the resource was dropped to fit
// the aggregate snapshot or count cap.
StatusExcluded
)
// String returns the lower-snake-case name used in IDs and
// metrics. Unknown values stringify to "unknown".
func (s ResourceStatus) String() string {
switch s {
case StatusOK:
return "ok"
case StatusOversize:
return "oversize"
case StatusUnreadable:
return "unreadable"
case StatusInvalid:
return "invalid"
case StatusExcluded:
return "excluded"
default:
return "unknown"
}
}
// Resource is what the resolver emits for each recognized file
// or live server it discovers under a scan root. The struct is
// intentionally flat; the typed wire mapping happens in
// drpc.go where Kind selects the proto oneof variant.
type Resource struct {
// ID is stable across pushes for the same logical
// resource. The current scheme is "<kind>:<source>". It is
// used for in-snapshot dedup and as part of the aggregate
// hash; it is not transmitted on the wire.
ID string
// Kind classifies the resource. Drives which proto oneof
// variant the DRPC adapter sets.
Kind ResourceKind
// Source is the file path or MCP server name.
Source string
// ContentHash is sha256 over the resource's original
// bytes (or transport-encoded server tool list).
ContentHash [32]byte
// Payload is the full bytes when Status == StatusOK; the
// per-resource and aggregate caps may leave it empty.
// Unused for KindMCPServer (Tools is used instead).
Payload []byte
// SizeBytes is the original payload size, populated
// regardless of Status.
SizeBytes uint64
// Status records OK or a reason the payload is absent.
Status ResourceStatus
// Error is populated whenever Status != StatusOK; may
// also carry a non-fatal warning when Status == StatusOK.
Error string
// Name is the resource's own short identifier. Currently
// populated for KindSkill (from front-matter) and
// KindMCPServer (server name); empty for other kinds.
Name string
// Description is a short human-readable summary (skill
// front-matter description, MCP server description,
// instruction-file first line). Shipped on the wire only
// for kinds whose body type carries a description field.
Description string
// SourcePath is the user-declared source that contributed
// the resource; empty for built-in scan roots.
SourcePath string
// Tools is populated for KindMCPServer with the live
// server's tool list; empty otherwise.
Tools []MCPTool
}
// MCPTool mirrors the wire MCPTool message. InputSchema is the
// JSON-Schema-shaped object the MCP server reported for the
// tool's arguments.
type MCPTool struct {
Name string
Description string
InputSchema map[string]any
}
// Snapshot is the immutable bundle of resources produced by a
// single resolver pass.
type Snapshot struct {
// Version is monotonically increasing per Manager
// instance; resets when the agent process restarts.
Version uint64
// SchemaVersion is bumped if the resource shape on the
// wire changes.
SchemaVersion uint64
// AggregateHash is sha256 over a canonical encoding of
// (ID, Kind, Source, ContentHash, Status) for every
// resource. Identical inputs always produce identical
// hashes; see ComputeAggregateHash.
AggregateHash [32]byte
// Resources is sorted by ID for deterministic encoding.
Resources []Resource
// PayloadBytes is the sum of len(Resource.Payload) across
// emitted resources after caps were applied.
PayloadBytes uint64
// SnapshotError carries a single snapshot-level error
// string when present (count cap exceeded, watcher
// degraded, ENOSPC, etc.). Empty when healthy.
SnapshotError string
}
// ComputeAggregateHash produces the deterministic snapshot
// aggregate hash for the supplied resources. The caller does
// not need to pre-sort; the function sorts a copy of the slice
// to keep its inputs side-effect free.
//
// The encoding is a Netstring-style stream. Each string field
// is written as the decimal-ASCII length, the literal ':', and
// the raw UTF-8 bytes. ContentHash is written as 32 raw bytes
// without a length prefix because it is a fixed-size SHA-256
// digest. Resources are separated by a single NUL byte. The
// scheme is internal to the agent and coderd, but it is stable
// across platforms because every field has an unambiguous
// length.
func ComputeAggregateHash(resources []Resource) [32]byte {
indexed := make([]Resource, len(resources))
copy(indexed, resources)
slices.SortFunc(indexed, func(a, b Resource) int {
return strings.Compare(a.ID, b.ID)
})
h := sha256.New()
for _, r := range indexed {
writeLengthPrefixed(h, r.ID)
writeLengthPrefixed(h, r.Kind.String())
writeLengthPrefixed(h, r.Source)
_, _ = h.Write(r.ContentHash[:])
writeLengthPrefixed(h, r.Status.String())
_, _ = h.Write([]byte{0})
}
var out [32]byte
copy(out[:], h.Sum(nil))
return out
}
// writeLengthPrefixed writes a decimal-ASCII length prefix, a
// literal ':' separator, and the raw bytes of s. This matches
// the Netstring framing used by ComputeAggregateHash.
func writeLengthPrefixed(h interface{ Write([]byte) (int, error) }, s string) {
_, _ = h.Write([]byte(strconv.Itoa(len(s))))
_, _ = h.Write([]byte{':'})
_, _ = h.Write([]byte(s))
}
+554
View File
@@ -0,0 +1,554 @@
package agentcontext_test
import (
"crypto/sha256"
"os"
"path/filepath"
"runtime"
"testing"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/agent/agentcontext"
)
func mustWriteFile(t *testing.T, path, content string) {
t.Helper()
require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755))
require.NoError(t, os.WriteFile(path, []byte(content), 0o600))
}
func mustWriteSkill(t *testing.T, dir, name, description string) {
t.Helper()
require.NoError(t, os.MkdirAll(filepath.Join(dir, name), 0o755))
mustWriteFile(t, filepath.Join(dir, name, "SKILL.md"),
"---\nname: "+name+"\ndescription: "+description+"\n---\nSkill body for "+name)
}
func findResource(t *testing.T, resources []agentcontext.Resource, kind agentcontext.ResourceKind, source string) agentcontext.Resource {
t.Helper()
for _, r := range resources {
if r.Kind == kind && r.Source == source {
return r
}
}
t.Fatalf("resource not found: kind=%s source=%s", kind, source)
return agentcontext.Resource{}
}
func TestResolver_ProjectAGENTSFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), "# Project rules\n\nDo the thing.")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.KindInstructionFile, got.Kind)
require.Equal(t, agentcontext.StatusOK, got.Status)
require.Equal(t, filepath.Join(dir, "AGENTS.md"), got.Source)
require.Contains(t, string(got.Payload), "Do the thing.")
require.Equal(t, "Project rules", got.Description)
require.NotEqual(t, [32]byte{}, got.ContentHash)
}
func TestResolver_CaseInsensitiveInstructionNames(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "agents.md"), "lower\n")
mustWriteFile(t, filepath.Join(dir, "CLAUDE.md"), "claude\n")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 2)
}
func TestResolver_SkillsContainerEmitsEachSubdir(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteSkill(t, filepath.Join(dir, ".agents", "skills"), "make-coffee", "Coffee skill")
mustWriteSkill(t, filepath.Join(dir, ".agents", "skills"), "fold-laundry", "Laundry skill")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
var kinds []string
for _, res := range snap.Resources {
kinds = append(kinds, res.Kind.String()+":"+filepath.Base(res.Source))
}
require.ElementsMatch(t, []string{
"skill:make-coffee",
"skill:fold-laundry",
}, kinds)
}
func TestResolver_SkillNameMismatchInvalid(t *testing.T) {
t.Parallel()
dir := t.TempDir()
skillsDir := filepath.Join(dir, ".agents", "skills", "make-coffee")
require.NoError(t, os.MkdirAll(skillsDir, 0o755))
mustWriteFile(t, filepath.Join(skillsDir, "SKILL.md"),
"---\nname: drink-tea\ndescription: oops\n---\nBody")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.KindSkill, got.Kind)
require.Equal(t, agentcontext.StatusInvalid, got.Status)
require.Contains(t, got.Error, "does not match directory")
}
// TestResolver_SkillNameNonKebabInvalid exercises the kebab-case
// validation branch in readSkillMeta. The skill name matches the
// parent directory (so the mismatch check passes) but contains
// characters that SkillNamePattern rejects. Without this test
// the kebab branch could be deleted and the suite would still
// pass.
func TestResolver_SkillNameNonKebabInvalid(t *testing.T) {
t.Parallel()
dir := t.TempDir()
skillDir := filepath.Join(dir, ".agents", "skills", "Make_Coffee")
require.NoError(t, os.MkdirAll(skillDir, 0o755))
mustWriteFile(t, filepath.Join(skillDir, "SKILL.md"),
"---\nname: Make_Coffee\ndescription: oops\n---\nBody")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.KindSkill, got.Kind)
require.Equal(t, agentcontext.StatusInvalid, got.Status)
require.Contains(t, got.Error, "kebab-case")
}
func TestResolver_MCPConfigEmitted(t *testing.T) {
t.Parallel()
dir := t.TempDir()
contents := `{"mcpServers": {"github": {"env": {"GITHUB_TOKEN": "secret-token"}}}}`
mustWriteFile(t, filepath.Join(dir, ".mcp.json"), contents)
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.KindMCPConfig, got.Kind)
require.Equal(t, agentcontext.StatusOK, got.Status)
// The .mcp.json payload is intentionally not shipped:
// the file can contain secret-bearing Env/Headers values.
// Only the path + ContentHash are exposed, so consumers
// can detect changes without ever seeing the bytes.
require.Empty(t, got.Payload, "readMCPConfig must not include the file payload")
require.NotEqual(t, [32]byte{}, got.ContentHash, "readMCPConfig must populate ContentHash for change detection")
require.Equal(t, uint64(len(contents)), got.SizeBytes)
}
// TestResolver_SymlinkInsideScanRootAllowed exercises the
// monorepo case where AGENTS.md is symlinked to shared content
// inside the same workspace tree. The target lives under the
// scan root, so the resolver follows the symlink and emits the
// target bytes as if the symlink were a regular file.
func TestResolver_SymlinkInsideScanRootAllowed(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("symlinks require admin privileges on Windows runners")
}
t.Parallel()
dir := t.TempDir()
target := filepath.Join(dir, "docs", "AGENTS.md")
require.NoError(t, os.MkdirAll(filepath.Dir(target), 0o755))
mustWriteFile(t, target, "shared monorepo guidance")
link := filepath.Join(dir, "AGENTS.md")
require.NoError(t, os.Symlink(target, link))
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 2)
var linked agentcontext.Resource
for _, res := range snap.Resources {
if res.Source == link {
linked = res
}
}
require.Equal(t, agentcontext.StatusOK, linked.Status)
require.Equal(t, "shared monorepo guidance", string(linked.Payload))
}
// TestResolver_SymlinkOutsideScanRootRejected guards the
// security boundary. A malicious workspace cannot ship a
// snapshot containing ~/.ssh/id_rsa or /etc/passwd by placing a
// symlink with that target at AGENTS.md, .mcp.json, or
// SKILL.md inside the scan root.
func TestResolver_SymlinkOutsideScanRootRejected(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("symlinks require admin privileges on Windows runners")
}
t.Parallel()
dir := t.TempDir()
secretDir := t.TempDir()
secret := filepath.Join(secretDir, "id_rsa")
mustWriteFile(t, secret, "-----BEGIN OPENSSH PRIVATE KEY-----")
link := filepath.Join(dir, "AGENTS.md")
require.NoError(t, os.Symlink(secret, link))
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.StatusInvalid, got.Status)
require.Empty(t, got.Payload, "escaping symlink target must not be shipped")
require.Contains(t, got.Error, "escapes scan root")
}
// TestResolver_BrokenSymlink emits Unreadable for a dangling
// link rather than crashing the walk.
func TestResolver_BrokenSymlink(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("symlinks require admin privileges on Windows runners")
}
t.Parallel()
dir := t.TempDir()
link := filepath.Join(dir, "AGENTS.md")
require.NoError(t, os.Symlink(filepath.Join(dir, "does-not-exist"), link))
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
require.Equal(t, agentcontext.StatusUnreadable, snap.Resources[0].Status)
}
func TestResolver_OversizeInstructionFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
// Write a file larger than the per-resource cap.
big := make([]byte, 200)
for i := range big {
big[i] = 'a'
}
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), string(big))
r := &agentcontext.Resolver{MaxResourceBytes: 100}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.StatusOversize, got.Status)
require.Empty(t, got.Payload)
require.Equal(t, uint64(200), got.SizeBytes)
// Hash over capped slice is still populated so callers
// can detect "still oversize but content changed".
require.NotEqual(t, [32]byte{}, got.ContentHash)
}
func TestResolver_AggregateCapExcludes(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), "small")
subA := filepath.Join(dir, "a")
subB := filepath.Join(dir, "b")
mustWriteFile(t, filepath.Join(subA, "AGENTS.md"), "AAAA")
mustWriteFile(t, filepath.Join(subB, "AGENTS.md"), "BBBB")
// Aggregate cap of 9 bytes lets the first two through but
// excludes the third regardless of which order they
// appear.
r := &agentcontext.Resolver{MaxSnapshotBytes: 9}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
var excluded int
for _, res := range snap.Resources {
if res.Status == agentcontext.StatusExcluded {
excluded++
}
}
require.Equal(t, 1, excluded)
}
func TestResolver_CountCapExcludes(t *testing.T) {
t.Parallel()
dir := t.TempDir()
for i := 0; i < 5; i++ {
sub := filepath.Join(dir, "dir", string('a'+rune(i)))
mustWriteFile(t, filepath.Join(sub, "AGENTS.md"), "x")
}
r := &agentcontext.Resolver{MaxResources: 3}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 5)
var excluded int
for _, res := range snap.Resources {
if res.Status == agentcontext.StatusExcluded {
excluded++
}
}
require.Equal(t, 2, excluded)
}
func TestResolver_SkipsVendorAndNodeModules(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), "root")
mustWriteFile(t, filepath.Join(dir, "node_modules", "deep", "AGENTS.md"), "should not appear")
mustWriteFile(t, filepath.Join(dir, "vendor", "AGENTS.md"), "should not appear either")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
require.Equal(t, filepath.Join(dir, "AGENTS.md"), snap.Resources[0].Source)
}
func TestResolver_UserSourceAttribution(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), "user-added")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir, UserSource: dir}})
require.Len(t, snap.Resources, 1)
require.Equal(t, dir, snap.Resources[0].SourcePath)
}
func TestResolver_MissingRootSilentlyIgnored(t *testing.T) {
t.Parallel()
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: "/nonexistent/path"}})
require.Empty(t, snap.Resources)
require.Empty(t, snap.SnapshotError)
}
func TestResolver_SingleFileRootClassified(t *testing.T) {
t.Parallel()
dir := t.TempDir()
path := filepath.Join(dir, "AGENTS.md")
mustWriteFile(t, path, "x")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: path}})
require.Len(t, snap.Resources, 1)
require.Equal(t, agentcontext.KindInstructionFile, snap.Resources[0].Kind)
}
func TestResolver_DuplicateRootsDeduplicated(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mustWriteFile(t, filepath.Join(dir, "AGENTS.md"), "x")
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{
{Path: dir},
{Path: dir},
{Path: dir},
})
require.Len(t, snap.Resources, 1)
}
func TestResolver_MCPProviderResources(t *testing.T) {
t.Parallel()
dir := t.TempDir()
mcpRes := agentcontext.Resource{
ID: "mcp_server:github",
Kind: agentcontext.KindMCPServer,
Source: "github",
Status: agentcontext.StatusOK,
Payload: []byte("tool-list-json"),
ContentHash: sha256.Sum256([]byte("tool-list-json")),
Description: "GitHub MCP server",
}
r := &agentcontext.Resolver{
MCP: &fakeMCPProvider{resources: []agentcontext.Resource{mcpRes}},
}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
got := findResource(t, snap.Resources, agentcontext.KindMCPServer, "github")
require.Equal(t, agentcontext.StatusOK, got.Status)
require.Equal(t, "GitHub MCP server", got.Description)
}
// TestResolver_MCPProviderRespectsAggregateByteCap guards the
// contract that a single oversized MCP payload cannot blow past
// MaxSnapshotBytes with StatusOK.
func TestResolver_MCPProviderRespectsAggregateByteCap(t *testing.T) {
t.Parallel()
dir := t.TempDir()
big := make([]byte, 1024)
for i := range big {
big[i] = 'x'
}
mcpRes := agentcontext.Resource{
ID: "mcp_server:big",
Kind: agentcontext.KindMCPServer,
Source: "big",
Status: agentcontext.StatusOK,
Payload: big,
ContentHash: sha256.Sum256(big),
}
r := &agentcontext.Resolver{
MaxSnapshotBytes: 512,
MCP: &fakeMCPProvider{resources: []agentcontext.Resource{mcpRes}},
}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
got := findResource(t, snap.Resources, agentcontext.KindMCPServer, "big")
require.Equal(t, agentcontext.StatusExcluded, got.Status,
"MCP payload exceeding MaxSnapshotBytes must be excluded")
require.Empty(t, got.Payload)
require.NotEmpty(t, snap.SnapshotError, "snapshot must surface the cap breach")
}
type fakeMCPProvider struct {
resources []agentcontext.Resource
}
func (f *fakeMCPProvider) MCPResources() []agentcontext.Resource {
return f.resources
}
// TestResolver_UnreadableInstructionFile verifies the
// permission-denied walk path produces a StatusUnreadable
// resource classified with the correct kind, matching the
// classification the resolver would emit on a successful read.
func TestResolver_UnreadableInstructionFile(t *testing.T) {
t.Parallel()
if runtime.GOOS == "windows" {
t.Skip("file mode 0o000 does not deny reads on Windows")
}
if os.Geteuid() == 0 {
t.Skip("root bypasses file mode permissions")
}
dir := t.TempDir()
path := filepath.Join(dir, "AGENTS.md")
mustWriteFile(t, path, "hello")
require.NoError(t, os.Chmod(path, 0o000))
t.Cleanup(func() { _ = os.Chmod(path, 0o600) })
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.KindInstructionFile, got.Kind)
require.Equal(t, agentcontext.StatusUnreadable, got.Status)
require.NotEmpty(t, got.Error)
}
// TestResolver_UnreadableMCPConfig confirms the walk-error path
// uses the file's real kind, not a hardcoded fallback. Without
// this, a permission flip on .mcp.json would produce a phantom
// resource ID swap when the permission is later restored.
func TestResolver_UnreadableMCPConfig(t *testing.T) {
t.Parallel()
if runtime.GOOS == "windows" {
t.Skip("file mode 0o000 does not deny reads on Windows")
}
if os.Geteuid() == 0 {
t.Skip("root bypasses file mode permissions")
}
dir := t.TempDir()
path := filepath.Join(dir, ".mcp.json")
mustWriteFile(t, path, `{"mcpServers": {}}`)
require.NoError(t, os.Chmod(path, 0o000))
t.Cleanup(func() { _ = os.Chmod(path, 0o600) })
r := &agentcontext.Resolver{}
snap := r.Resolve([]agentcontext.ScanRoot{{Path: dir}})
require.Len(t, snap.Resources, 1)
got := snap.Resources[0]
require.Equal(t, agentcontext.KindMCPConfig, got.Kind)
require.Equal(t, agentcontext.StatusUnreadable, got.Status)
require.NotEmpty(t, got.Error)
}
func TestResourceKindString(t *testing.T) {
t.Parallel()
tests := []struct {
kind agentcontext.ResourceKind
want string
}{
{agentcontext.KindUnspecified, "unknown"},
{agentcontext.KindInstructionFile, "instruction_file"},
{agentcontext.KindSkill, "skill"},
{agentcontext.KindMCPConfig, "mcp_config"},
{agentcontext.KindMCPServer, "mcp_server"},
{agentcontext.KindPlugin, "plugin"},
{agentcontext.KindHook, "hook"},
{agentcontext.KindSubagent, "subagent"},
{agentcontext.KindCommand, "command"},
{agentcontext.ResourceKind(999), "unknown"},
}
for _, tt := range tests {
require.Equal(t, tt.want, tt.kind.String())
}
}
func TestResourceStatusString(t *testing.T) {
t.Parallel()
tests := []struct {
status agentcontext.ResourceStatus
want string
}{
{agentcontext.StatusOK, "ok"},
{agentcontext.StatusOversize, "oversize"},
{agentcontext.StatusUnreadable, "unreadable"},
{agentcontext.StatusInvalid, "invalid"},
{agentcontext.StatusExcluded, "excluded"},
{agentcontext.ResourceStatus(999), "unknown"},
}
for _, tt := range tests {
require.Equal(t, tt.want, tt.status.String())
}
}
func TestComputeAggregateHash_DeterministicAcrossOrder(t *testing.T) {
t.Parallel()
a := agentcontext.Resource{
ID: "instruction_file:/a/AGENTS.md",
Kind: agentcontext.KindInstructionFile,
Source: "/a/AGENTS.md",
Status: agentcontext.StatusOK,
}
b := agentcontext.Resource{
ID: "instruction_file:/b/AGENTS.md",
Kind: agentcontext.KindInstructionFile,
Source: "/b/AGENTS.md",
Status: agentcontext.StatusOK,
}
got1 := agentcontext.ComputeAggregateHash([]agentcontext.Resource{a, b})
got2 := agentcontext.ComputeAggregateHash([]agentcontext.Resource{b, a})
require.Equal(t, got1, got2)
}
func TestComputeAggregateHash_ChangesOnContent(t *testing.T) {
t.Parallel()
base := agentcontext.Resource{
ID: "instruction_file:/a/AGENTS.md",
Kind: agentcontext.KindInstructionFile,
Source: "/a/AGENTS.md",
Status: agentcontext.StatusOK,
}
hash1 := agentcontext.ComputeAggregateHash([]agentcontext.Resource{base})
withContent := base
withContent.ContentHash = [32]byte{0x01}
hash2 := agentcontext.ComputeAggregateHash([]agentcontext.Resource{withContent})
require.NotEqual(t, hash1, hash2)
withStatus := base
withStatus.Status = agentcontext.StatusOversize
hash3 := agentcontext.ComputeAggregateHash([]agentcontext.Resource{withStatus})
require.NotEqual(t, hash1, hash3)
}
+391
View File
@@ -0,0 +1,391 @@
package agentcontext
import (
"context"
"errors"
"io/fs"
"os"
"path/filepath"
"strings"
"sync"
"syscall"
"time"
"github.com/fsnotify/fsnotify"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/quartz"
)
// DefaultWatchDebounce coalesces editor-style multi-event writes
// (truncate plus rename plus chmod) into a single re-resolve.
// Mirrors the debounce window the existing MCP config watcher
// uses so behavior is consistent across the agent.
const DefaultWatchDebounce = 250 * time.Millisecond
// WatcherOptions parameterizes the recursive watcher.
type WatcherOptions struct {
Logger slog.Logger
Clock quartz.Clock
Debounce time.Duration
// MaxDepth caps the recursion depth when discovering
// subdirectories to watch. Zero defaults to
// DefaultMaxScanDepth. Callers wiring the watcher to a
// Resolver should pass the resolver's MaxDepth so the
// watcher never misses edits below the scan horizon.
MaxDepth int
// OnChange runs at most once per debounce window. The
// caller must not block; the recommended pattern is a
// non-blocking send on a re-resolve trigger channel.
OnChange func()
}
// Watcher is a recursive fsnotify wrapper. fsnotify does not
// support recursive watches natively on Linux, so we walk every
// scan root at sync time and register each subdirectory
// individually. Inotify ENOSPC degrades the watcher into a
// poll-only mode that still re-resolves on Sync calls.
type Watcher struct {
logger slog.Logger
clock quartz.Clock
debounce time.Duration
maxDepth int
onChange func()
mu sync.Mutex
watcher *fsnotify.Watcher
watched map[string]struct{}
timer *quartz.Timer
degraded string // non-empty when the watcher dropped events
closed bool
closedCh chan struct{}
runDoneCh chan struct{}
}
// NewWatcher constructs a recursive watcher. The watcher does
// nothing until Sync is called.
func NewWatcher(opts WatcherOptions) (*Watcher, error) {
if opts.OnChange == nil {
return nil, xerrors.New("OnChange callback is required")
}
debounce := opts.Debounce
if debounce <= 0 {
debounce = DefaultWatchDebounce
}
clock := opts.Clock
if clock == nil {
clock = quartz.NewReal()
}
maxDepth := opts.MaxDepth
if maxDepth <= 0 {
maxDepth = DefaultMaxScanDepth
}
w, err := fsnotify.NewWatcher()
if err != nil {
// On Linux, fsnotify.NewWatcher only fails when the
// inotify subsystem is at the system-wide watch
// limit. Surface a Watcher in "degraded" mode so the
// caller can still rely on explicit Sync triggers.
degraded := &Watcher{
logger: opts.Logger,
clock: clock,
debounce: debounce,
maxDepth: maxDepth,
onChange: opts.OnChange,
watched: make(map[string]struct{}),
degraded: "fsnotify init failed: " + err.Error(),
closedCh: make(chan struct{}),
runDoneCh: closedChan(),
}
return degraded, nil
}
cw := &Watcher{
logger: opts.Logger,
clock: clock,
debounce: debounce,
maxDepth: maxDepth,
onChange: opts.OnChange,
watcher: w,
watched: make(map[string]struct{}),
closedCh: make(chan struct{}),
runDoneCh: make(chan struct{}),
}
go cw.run()
return cw, nil
}
// closedChan returns an already-closed channel for the
// degraded-watcher case where there is no run goroutine.
func closedChan() chan struct{} {
c := make(chan struct{})
close(c)
return c
}
// Degraded returns a non-empty string when the watcher is
// running with reduced functionality (typically inotify
// ENOSPC). The string is suitable for use as a snapshot-level
// error message.
func (w *Watcher) Degraded() string {
w.mu.Lock()
defer w.mu.Unlock()
return w.degraded
}
// Sync replaces the set of watched directories with a fresh
// recursive walk of every scan root. Files are not watched
// directly; watching the parent directory catches creates,
// renames, removes, and writes that touch any recognized
// basename. Files that are themselves scan roots are handled by
// watching their parent.
//
// Sync is idempotent and safe to call repeatedly. The lock is
// released around the recursive directory walk so concurrent
// Close, schedule, and the run goroutine are not blocked by a
// slow filesystem.
func (w *Watcher) Sync(ctx context.Context, roots []ScanRoot) {
w.mu.Lock()
if w.closed {
w.mu.Unlock()
return
}
if w.watcher == nil {
// Degraded mode: no fsnotify, so there is nothing
// to wire up. Do NOT fire the OnChange callback
// from here; the Manager's signal handler is the
// usual OnChange, and the Run loop calls back into
// Sync when it observes that signal. Firing here
// would re-arm an endless 250ms scan-and-push loop
// on hosts where inotify cannot initialize. Manual
// Resync, AddSource, and RemoveSource still drive
// re-resolves; auto-updates on file edits simply
// do not happen until fsnotify recovers.
w.mu.Unlock()
return
}
w.mu.Unlock()
// collectDirs touches the filesystem (filepath.WalkDir on
// every scan root). Compute the desired set outside the
// mutex so a slow walk does not block the run goroutine,
// Close, or schedule.
desired := w.collectDirs(roots)
w.mu.Lock()
defer w.mu.Unlock()
if w.closed {
return
}
// Remove directories no longer wanted.
for path := range w.watched {
if _, ok := desired[path]; ok {
continue
}
_ = w.watcher.Remove(path)
delete(w.watched, path)
}
// Track whether every Add in this pass succeeded so a
// recovered ENOSPC clears the degraded marker.
addedAll := true
// Add directories that are new.
for path := range desired {
if _, ok := w.watched[path]; ok {
continue
}
if err := w.watcher.Add(path); err != nil {
// ENOSPC means the kernel's per-user inotify
// watch budget is exhausted. Mark the watcher
// degraded; subsequent Sync calls still fire
// the change callback so resync still works.
if errors.Is(err, syscall.ENOSPC) {
w.degraded = "inotify watch limit exceeded (ENOSPC)"
addedAll = false
w.logger.Warn(ctx, "context watcher degraded: inotify watch limit exceeded",
slog.F("dir", path))
break
}
w.logger.Debug(ctx, "context watcher could not add dir",
slog.F("dir", path), slog.Error(err))
continue
}
w.watched[path] = struct{}{}
}
// Clear a previously-set ENOSPC mark when every Add in this
// pass succeeded. A user who bumps the kernel's inotify
// limit and re-syncs now sees a clean snapshot instead of a
// permanent SnapshotError.
if addedAll && w.degraded != "" {
w.degraded = ""
}
}
// Close stops the watcher and releases all kernel watch slots.
// Close is idempotent.
func (w *Watcher) Close() error {
w.mu.Lock()
if w.closed {
w.mu.Unlock()
return nil
}
w.closed = true
close(w.closedCh)
timer := w.timer
watcher := w.watcher
w.timer = nil
w.watcher = nil
w.mu.Unlock()
if timer != nil {
timer.Stop()
}
if watcher != nil {
_ = watcher.Close()
}
<-w.runDoneCh
return nil
}
// run forwards fsnotify events into the debounce timer. It exits
// when Close is called or the underlying watcher is closed.
func (w *Watcher) run() {
defer close(w.runDoneCh)
// Capture the watcher reference once. Close may set the
// field to nil concurrently; reading the captured local
// keeps the event loop safe through the race window.
w.mu.Lock()
fsw := w.watcher
w.mu.Unlock()
if fsw == nil {
return
}
for {
select {
case <-w.closedCh:
return
case ev, ok := <-fsw.Events:
if !ok {
return
}
if !w.eventRelevant(ev) {
continue
}
w.schedule()
case err, ok := <-fsw.Errors:
if !ok {
return
}
if err != nil {
w.logger.Debug(context.Background(), "context watcher error", slog.Error(err))
}
}
}
}
// eventRelevant filters out events that cannot affect any
// recognized resource. The check is conservative: any event on
// a directory triggers a re-resolve so newly created subtrees
// are picked up.
func (*Watcher) eventRelevant(ev fsnotify.Event) bool {
name := filepath.Base(ev.Name)
if recognizedInstructionFile(name) || name == mcpConfigFileName || name == skillMetaFileName {
return true
}
// Directory create/remove flips re-resolve so new subtrees
// arm watches and removed subtrees stop arming them.
if ev.Has(fsnotify.Create) || ev.Has(fsnotify.Remove) || ev.Has(fsnotify.Rename) {
return true
}
return false
}
// schedule arms or resets the debounce timer.
func (w *Watcher) schedule() {
w.mu.Lock()
if w.closed {
w.mu.Unlock()
return
}
cb := w.onChange
if w.timer != nil {
w.timer.Reset(w.debounce)
w.mu.Unlock()
return
}
w.timer = w.clock.AfterFunc(w.debounce, func() {
w.mu.Lock()
w.timer = nil
w.mu.Unlock()
cb()
})
w.mu.Unlock()
}
// collectDirs walks every scan root and returns the set of
// directories to watch. The maximum depth uses the watcher's
// configured maxDepth so it mirrors the resolver's horizon.
func (w *Watcher) collectDirs(roots []ScanRoot) map[string]struct{} {
out := make(map[string]struct{})
for _, root := range roots {
if root.Path == "" {
continue
}
info, err := os.Stat(root.Path)
if err != nil {
// Watch the deepest existing ancestor so the
// root being created later still fires.
if ancestor := existingAncestor(root.Path); ancestor != "" {
out[ancestor] = struct{}{}
}
continue
}
if !info.IsDir() {
out[filepath.Dir(root.Path)] = struct{}{}
continue
}
// Walk the directory and collect every descendant
// directory up to the depth cap.
rootDepth := strings.Count(filepath.Clean(root.Path), string(os.PathSeparator))
_ = filepath.WalkDir(root.Path, func(path string, d fs.DirEntry, err error) error {
if err != nil {
return nil
}
if !d.IsDir() {
return nil
}
if _, skip := skipDirNames[d.Name()]; skip && path != root.Path {
return fs.SkipDir
}
if strings.Count(path, string(os.PathSeparator))-rootDepth > w.maxDepth {
return fs.SkipDir
}
out[path] = struct{}{}
return nil
})
}
return out
}
// existingAncestor returns the deepest existing ancestor of
// path, or "" if no ancestor exists (e.g. an entirely missing
// drive on Windows).
func existingAncestor(path string) string {
cur := filepath.Dir(path)
for {
if cur == "" || cur == "." {
return ""
}
info, err := os.Stat(cur)
if err == nil && info.IsDir() {
return cur
}
parent := filepath.Dir(cur)
if parent == cur {
return ""
}
cur = parent
}
}
+97
View File
@@ -0,0 +1,97 @@
package agentcontext_test
import (
"context"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/agent/agentcontext"
"github.com/coder/coder/v2/testutil"
"github.com/coder/quartz"
)
func TestWatcher_FiresOnAgentsMdEdit(t *testing.T) {
t.Parallel()
dir := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v1"), 0o600))
var fires atomic.Int32
w, err := agentcontext.NewWatcher(agentcontext.WatcherOptions{
Logger: testutil.Logger(t).Named("watcher"),
Clock: quartz.NewReal(),
Debounce: 10 * time.Millisecond,
OnChange: func() { fires.Add(1) },
})
require.NoError(t, err)
t.Cleanup(func() { _ = w.Close() })
ctx := testutil.Context(t, testutil.WaitShort)
w.Sync(ctx, []agentcontext.ScanRoot{{Path: dir}})
// Rewrite the file inside Eventually so the test does not race
// fsnotify's watch-setup window. As soon as the watch is live,
// the next write fires the debounce timer.
require.Eventually(t, func() bool {
_ = os.WriteFile(filepath.Join(dir, "AGENTS.md"), []byte("v2"), 0o600)
return fires.Load() >= 1
}, testutil.WaitShort, testutil.IntervalFast, "expected at least one fire after AGENTS.md edit")
}
func TestWatcher_FiresOnNewSkillFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
skillsRoot := filepath.Join(dir, ".agents", "skills")
require.NoError(t, os.MkdirAll(skillsRoot, 0o755))
var fires atomic.Int32
w, err := agentcontext.NewWatcher(agentcontext.WatcherOptions{
Logger: testutil.Logger(t).Named("watcher"),
Debounce: 10 * time.Millisecond,
OnChange: func() { fires.Add(1) },
})
require.NoError(t, err)
t.Cleanup(func() { _ = w.Close() })
ctx := testutil.Context(t, testutil.WaitShort)
w.Sync(ctx, []agentcontext.ScanRoot{{Path: dir}})
// Create SKILL.md inside Eventually so the test does not race
// fsnotify's watch-setup window. The Manager pre-creates the
// skill dir, then rewrites SKILL.md each tick until the watcher
// fires at least once.
skillDir := filepath.Join(skillsRoot, "foo")
require.NoError(t, os.MkdirAll(skillDir, 0o755))
require.Eventually(t, func() bool {
_ = os.WriteFile(filepath.Join(skillDir, "SKILL.md"), []byte("---\nname: foo\ndescription: bar\n---\nbody"), 0o600)
return fires.Load() >= 1
}, testutil.WaitShort, testutil.IntervalFast, "expected fire after SKILL.md create")
}
func TestWatcher_CloseIsIdempotent(t *testing.T) {
t.Parallel()
w, err := agentcontext.NewWatcher(agentcontext.WatcherOptions{
Logger: testutil.Logger(t).Named("watcher"),
OnChange: func() {},
})
require.NoError(t, err)
require.NoError(t, w.Close())
require.NoError(t, w.Close())
}
func TestWatcher_SyncAfterCloseNoop(t *testing.T) {
t.Parallel()
w, err := agentcontext.NewWatcher(agentcontext.WatcherOptions{
Logger: testutil.Logger(t).Named("watcher"),
OnChange: func() {},
})
require.NoError(t, err)
require.NoError(t, w.Close())
// Must not panic.
w.Sync(context.Background(), []agentcontext.ScanRoot{{Path: t.TempDir()}})
}