From 0af7dd62612207e324b4799ace6d7363aa025895 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Tue, 6 Oct 2026 10:57:57 +0200 Subject: [PATCH] feat(skills): load skills from public GitHub repositories Adds a GitHub source type to the skills loader so agents can pull skills straight from a public github.com repo (branch, tag, or commit SHA), with a short-lived ref cache and immutable, safely-extracted snapshots for warm, network-free reloads. Assisted-By: docker-agent --- docs/features/skills/index.md | 74 ++++ examples/skills_github.yaml | 19 + pkg/config/latest/types.go | 3 +- pkg/skills/frontmatter.go | 8 +- pkg/skills/frontmatter_test.go | 13 + pkg/skills/github.go | 286 +++++++++++++ pkg/skills/github_archive.go | 280 ++++++++++++ pkg/skills/github_test.go | 754 +++++++++++++++++++++++++++++++++ pkg/skills/skills.go | 35 +- pkg/teamloader/teamloader.go | 5 +- 10 files changed, 1467 insertions(+), 10 deletions(-) create mode 100644 examples/skills_github.yaml create mode 100644 pkg/skills/github.go create mode 100644 pkg/skills/github_archive.go create mode 100644 pkg/skills/github_test.go diff --git a/docs/features/skills/index.md b/docs/features/skills/index.md index 5263d2b29..530fcbfad 100644 --- a/docs/features/skills/index.md +++ b/docs/features/skills/index.md @@ -62,6 +62,80 @@ agents: A name that doesn't match any discovered skill is logged as a warning at startup but is otherwise ignored. +## GitHub Skill Sources + +Load skills directly from a public GitHub repository, without installing Git or +an external skills CLI: + +```yaml +agents: + root: + model: openai/gpt-4o-mini + skills: + - local + - https://github.com/docker/skills + toolsets: + - type: filesystem + - type: shell +``` + +Repository URLs use the default branch. To select a branch, tag, commit SHA, or +subdirectory, use a GitHub tree URL or query parameters: + +```yaml +skills: + - https://github.com/docker/skills/tree/main/skills + # Alternative examples (use one source, not all of them): + # - https://github.com/docker/skills?ref=v0.3.0 + # - https://github.com/docker/skills?ref= + # - https://github.com/owner/repo?ref=feature%2Fskills&path=skills +``` + +A tree URL treats the first segment after `/tree/` as the revision and the rest +as a directory. Use `?ref=...&path=...` for branch or tag names containing `/`, or +URL-escape the slash as `%2F`. Short commit hashes are resolved through the API; +use a full 40-character SHA for an immutable pin. + +Without a selected directory, discovery checks `skills/**/SKILL.md`, then +`.agents/skills/**/SKILL.md`, then a root `SKILL.md`, using the first layout +containing skills. An explicit directory is searched recursively without +fallback. Complete skill directories, including supporting files, are cached. + +### Authentication and caching + +Only **public repositories on github.com** are supported. The optional +`GITHUB_TOKEN` comes from the configured environment provider, including OS +environment variables, env files, and secrets. It raises GitHub API rate limits; +it is not required or prompted for, and does not enable private repositories. +The token is sent only to `api.github.com`, never to the archive host. + +Default branches, named branches, and tags are resolved to a commit SHA and +cached for **five minutes**. Downloaded snapshots are immutable and cached by +repository, SHA, and directory. A warm load makes no network requests; after +five minutes mutable revisions are checked again. A full-SHA source needs no +further network requests once cached. All files come from the same commit. + +Archives are fetched from `codeload.github.com`. Downloads are bounded to 32 MiB +compressed, 128 MiB expanded, and 4,096 selected files, with a 1 MiB limit per +skill file and 32 MiB total selected content. Failed downloads never publish a +partial snapshot; source failures appear as load-time warnings. Expired mutable +revisions are not silently reused on network errors. + +### Trust and sandbox behavior + +Cached GitHub skills remain **remote**: embedded `` !`command` `` expressions are +not expanded. Remote frontmatter `model` and `toolsets` overrides are ignored. +Fork skills inherit the parent's model and tools, with `allowed-tools` able to +restrict inherited tools. Symlinks and other non-regular entries inside selected +skill directories are rejected; Git submodules are not fetched. + +In sandbox mode, GitHub sources are fetched inside the VM rather than staged as +local skills. Permit network access to `api.github.com` and `codeload.github.com` +and make any optional token available through the sandbox's environment provider. + +See [`examples/skills_github.yaml`](https://github.com/docker/docker-agent/blob/main/examples/skills_github.yaml) +for a complete configuration. + ## Inline Skills Instead of (or alongside) loading skills from files and URLs, you can define skills directly in the agent config. An inline skill is a mapping item in the `skills` list, freely mixed with the string items above: diff --git a/examples/skills_github.yaml b/examples/skills_github.yaml new file mode 100644 index 000000000..33caf5926 --- /dev/null +++ b/examples/skills_github.yaml @@ -0,0 +1,19 @@ +#!yaml +# Public GitHub skill sources. GITHUB_TOKEN is optional and uses the configured +# environment provider. Branch/tag resolutions are cached for five minutes; +# immutable commit snapshots are reused without downloading again. +agents: + root: + model: openai/gpt-4o-mini + description: Assistant with Docker skills from GitHub. + instruction: Use the available skills when a task matches their description. + skills: + - local + - https://github.com/docker/skills + # To pin a revision or select a directory instead: + # - https://github.com/docker/skills/tree/main/skills + # - https://github.com/docker/skills?ref=v0.3.0 + # - https://github.com/owner/repo?ref=feature%2Fskills&path=skills + toolsets: + - type: filesystem + - type: shell diff --git a/pkg/config/latest/types.go b/pkg/config/latest/types.go index b1bc0417a..fd20c9636 100644 --- a/pkg/config/latest/types.go +++ b/pkg/config/latest/types.go @@ -914,7 +914,8 @@ type InlineSkill struct { // inline definitions enables skills without loading local or remote ones. // // The special source "local" loads skills from the filesystem (standard locations). -// HTTP/HTTPS URLs load skills from remote servers per the well-known skills discovery spec. +// HTTPS GitHub repository URLs discover skills from commit snapshots. Other +// HTTP/HTTPS URLs use the well-known skills discovery spec. type SkillsConfig struct { // Sources lists where to load skills from: "local" and/or HTTP/HTTPS URLs. Sources []string diff --git a/pkg/skills/frontmatter.go b/pkg/skills/frontmatter.go index 081ff1a9b..accec6143 100644 --- a/pkg/skills/frontmatter.go +++ b/pkg/skills/frontmatter.go @@ -11,18 +11,16 @@ func parseFrontmatter(content string) (Skill, bool) { content = strings.ReplaceAll(content, "\r\n", "\n") content = strings.ReplaceAll(content, "\r", "\n") - rest, found := strings.CutPrefix(content, "---") + rest, found := strings.CutPrefix(content, "---\n") if !found { return Skill{}, false } - endIndex := strings.Index(rest, "\n---") - if endIndex == -1 { + block, _, found := strings.Cut("\n"+rest+"\n", "\n---\n") + if !found { return Skill{}, false } - block := content[4 : endIndex+3] - var skill Skill var currentKey string // tracks multi-line keys like "metadata" or "allowed-tools" diff --git a/pkg/skills/frontmatter_test.go b/pkg/skills/frontmatter_test.go index f16d060fe..b66bafffc 100644 --- a/pkg/skills/frontmatter_test.go +++ b/pkg/skills/frontmatter_test.go @@ -33,3 +33,16 @@ func TestSplitKeyValue(t *testing.T) { }) } } + +func TestParseFrontmatterMalformed(t *testing.T) { + t.Parallel() + for _, content := range []string{"", "---", "---\n---garbage", "---\nname: test\n---garbage"} { + t.Run(content, func(t *testing.T) { + t.Parallel() + _, ok := parseFrontmatter(content) + assert.False(t, ok) + }) + } + _, ok := parseFrontmatter("---\n---") + assert.True(t, ok) +} diff --git a/pkg/skills/github.go b/pkg/skills/github.go new file mode 100644 index 000000000..d38ae64ae --- /dev/null +++ b/pkg/skills/github.go @@ -0,0 +1,286 @@ +package skills + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "os" + "path" + "path/filepath" + "strings" + "time" + + "golang.org/x/sync/singleflight" + + "github.com/docker/docker-agent/pkg/atomicfile" + "github.com/docker/docker-agent/pkg/environment" +) + +const githubRefTTL = 5 * time.Minute + +var githubLoads singleflight.Group + +type githubSource struct { + repository string + ref string + directory string +} + +type githubResolution struct { + Commit string `json:"commit"` + CheckedAt time.Time `json:"checked_at"` +} + +// GitHub tree URLs use one escaped segment for the ref; query parameters also +// support branches containing slashes without confusing them with directories. +func parseGitHubSource(source string) (githubSource, bool, error) { + u, err := url.Parse(source) + if err != nil || !strings.EqualFold(u.Hostname(), "github.com") { + return githubSource{}, false, nil + } + invalid := func() (githubSource, bool, error) { + return githubSource{}, true, errors.New("expected https://github.com/owner/repo[/tree/ref/directory], optionally with ref and path query parameters") + } + if u.Scheme != "https" || u.User != nil || u.Port() != "" || u.Fragment != "" { + return invalid() + } + segments := strings.Split(strings.Trim(u.EscapedPath(), "/"), "/") + if len(segments) < 2 { + return invalid() + } + for i, segment := range segments { + segments[i], err = url.PathUnescape(segment) + if err != nil { + return invalid() + } + } + segments[1] = strings.TrimSuffix(segments[1], ".git") + for _, segment := range segments[:2] { + if !isValidSkillName(segment) { + return invalid() + } + } + result := githubSource{repository: strings.ToLower(strings.Join(segments[:2], "/"))} + if len(segments) > 2 { + if len(segments) < 4 || segments[2] != "tree" { + return invalid() + } + result.ref = segments[3] + if len(segments) > 4 { + result.directory = strings.Join(segments[4:], "/") + } + } + query, err := url.ParseQuery(u.RawQuery) + if err != nil { + return invalid() + } + for key, values := range query { + if len(values) != 1 || (key != "ref" && key != "path") { + return invalid() + } + } + if values, ok := query["ref"]; ok { + if result.ref != "" || values[0] == "" { + return invalid() + } + result.ref = values[0] + } + if values, ok := query["path"]; ok { + if result.directory != "" || values[0] == "" { + return invalid() + } + result.directory = values[0] + } + if strings.ContainsAny(result.ref, "\x00\r\n\t ") || (result.directory != "" && !validGitHubPath(result.directory)) { + return invalid() + } + return result, true, nil +} + +func isCommitSHA(ref string) bool { + if len(ref) != 40 { + return false + } + _, err := hex.DecodeString(ref) + return err == nil +} + +func (s githubSource) cacheKey() string { + return "github:" + s.repository + "?ref=" + url.QueryEscape(s.ref) + "&path=" + url.QueryEscape(s.directory) +} + +func loadGitHubSkills(ctx context.Context, source githubSource, cache *diskCache, env environment.Provider) ([]Skill, error) { + key := cache.cacheDir(source.cacheKey(), "github") + result := githubLoads.DoChan(key, func() (any, error) { + loadCtx, cancel := context.WithTimeout(ctx, 2*time.Minute) + defer cancel() + commit, err := source.resolve(loadCtx, cache, env) + if err != nil { + return nil, err + } + snapshot := cache.cacheDir("github:"+source.repository+"@"+commit, "snapshots") + snapshot = filepath.Join(snapshot, githubDirectoryKey(source.directory)) + if _, err := os.Stat(filepath.Join(snapshot, "complete")); errors.Is(err, os.ErrNotExist) { + if err := source.download(loadCtx, commit, snapshot); err != nil { + return nil, err + } + } else if err != nil { + return nil, fmt.Errorf("reading GitHub snapshot: %w", err) + } + return source.readSnapshot(snapshot) + }) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case result := <-result: + if result.Err != nil { + return nil, result.Err + } + return result.Val.([]Skill), nil + } +} + +func githubDirectoryKey(directory string) string { + // Reuse the cache's URL hashing so repository paths never become cache paths. + digest := sha256.Sum256([]byte(directory)) + return hex.EncodeToString(digest[:]) +} + +func (s githubSource) resolve(ctx context.Context, cache *diskCache, env environment.Provider) (string, error) { + resolutionPath := filepath.Join(cache.cacheDir(s.cacheKey(), "github"), "resolution.json") + var cached githubResolution + if data, err := os.ReadFile(resolutionPath); err == nil { + if json.Unmarshal(data, &cached) == nil && isCommitSHA(cached.Commit) && + (isCommitSHA(s.ref) && strings.EqualFold(s.ref, cached.Commit) || time.Since(cached.CheckedAt) < githubRefTTL) { + return cached.Commit, nil + } + } + var repository struct { + Private bool `json:"private"` + Visibility string `json:"visibility"` + DefaultBranch string `json:"default_branch"` + } + if err := githubJSON(ctx, "https://api.github.com/repos/"+s.repository, env, &repository); err != nil { + return "", err + } + if repository.Private || repository.Visibility != "public" { + return "", errors.New("GitHub skills require a public repository") + } + ref := s.ref + if ref == "" { + ref = repository.DefaultBranch + } + if ref == "" { + return "", errors.New("GitHub repository has no default branch") + } + commit := strings.ToLower(ref) + if !isCommitSHA(ref) { + var revision struct { + SHA string `json:"sha"` + } + if err := githubJSON(ctx, "https://api.github.com/repos/"+s.repository+"/commits/"+url.PathEscape(ref), env, &revision); err != nil { + return "", err + } + commit = revision.SHA + if !isCommitSHA(commit) { + return "", errors.New("GitHub returned an invalid commit SHA") + } + } + data, err := json.Marshal(githubResolution{Commit: commit, CheckedAt: time.Now()}) + if err != nil { + return "", err + } + if err := os.MkdirAll(filepath.Dir(resolutionPath), 0o700); err != nil { + return "", err + } + if err := atomicfile.Write(resolutionPath, strings.NewReader(string(data)), 0o600); err != nil { + return "", err + } + return commit, nil +} + +func githubJSON(ctx context.Context, endpoint string, env environment.Provider, result any) error { + resp, err := githubGet(ctx, endpoint, env) + if err != nil { + return err + } + defer resp.Body.Close() + data, err := readGitHubBody(resp.Body, 1<<20) + if err != nil { + return err + } + if err := json.Unmarshal(data, result); err != nil { + return fmt.Errorf("decoding GitHub response: %w", err) + } + return nil +} + +func githubGet(ctx context.Context, endpoint string, env environment.Provider) (*http.Response, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, http.NoBody) + if err != nil { + return nil, err + } + if req.URL.Host == "api.github.com" { + req.Header.Set("Accept", "application/vnd.github+json") + req.Header.Set("X-GitHub-Api-Version", "2022-11-28") + if env != nil { + if token, _ := env.Get(ctx, "GITHUB_TOKEN"); token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + } + } + // Canonical GitHub endpoints need no redirects; never forward credentials. + client := *skillsHTTPClient + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + resp, err := client.Do(req) + if err != nil { + return nil, err + } + if resp.StatusCode != http.StatusOK { + resp.Body.Close() + if resp.StatusCode == http.StatusForbidden || resp.StatusCode == http.StatusTooManyRequests { + return nil, fmt.Errorf("GitHub HTTP %d (check GITHUB_TOKEN and API rate limits)", resp.StatusCode) + } + return nil, fmt.Errorf("fetching %s: HTTP %d", endpoint, resp.StatusCode) + } + return resp, nil +} + +func readGitHubBody(reader io.Reader, limit int64) ([]byte, error) { + body, err := io.ReadAll(io.LimitReader(reader, limit+1)) + if err != nil { + return nil, err + } + if int64(len(body)) > limit { + return nil, errors.New("GitHub response exceeds size limit") + } + return body, nil +} + +func validGitHubPath(value string) bool { + if value == "" || value != path.Clean(value) || strings.HasPrefix(value, "/") { + return false + } + for segment := range strings.SplitSeq(value, "/") { + if segment == "." || segment == ".." || strings.HasSuffix(segment, ".") || strings.HasSuffix(segment, " ") { + return false + } + for _, c := range segment { + if c < 0x20 || strings.ContainsRune(`\:*?"<>|`, c) { + return false + } + } + base := strings.ToUpper(strings.SplitN(segment, ".", 2)[0]) + if base == "CON" || base == "PRN" || base == "AUX" || base == "NUL" || + len(base) == 4 && (strings.HasPrefix(base, "COM") || strings.HasPrefix(base, "LPT")) && base[3] >= '0' && base[3] <= '9' { + return false + } + } + return true +} diff --git a/pkg/skills/github_archive.go b/pkg/skills/github_archive.go new file mode 100644 index 000000000..7184e94a4 --- /dev/null +++ b/pkg/skills/github_archive.go @@ -0,0 +1,280 @@ +package skills + +import ( + "archive/tar" + "bytes" + "compress/gzip" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "os" + "path" + "path/filepath" + "slices" + "strings" + + "github.com/docker/docker-agent/pkg/atomicfile" +) + +const ( + githubArchiveLimit = 32 << 20 + githubExpandedLimit = 128 << 20 + githubFileLimit = 1 << 20 + githubFilesLimit = 4096 +) + +func (s githubSource) download(ctx context.Context, commit, snapshot string) error { + resp, err := githubGet(ctx, "https://codeload.github.com/"+s.repository+"/tar.gz/"+commit, nil) + if err != nil { + return err + } + defer resp.Body.Close() + archive, err := readGitHubBody(resp.Body, githubArchiveLimit) + if err != nil { + return err + } + roots, err := s.archiveRoots(ctx, archive) + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(snapshot), 0o700); err != nil { + return err + } + staging, err := os.MkdirTemp(filepath.Dir(snapshot), ".download-") + if err != nil { + return err + } + defer os.RemoveAll(staging) + if err := extractGitHubSkills(ctx, archive, roots, filepath.Join(staging, "content")); err != nil { + return err + } + metadata, err := json.Marshal(roots) + if err != nil { + return err + } + if err := atomicfile.Write(filepath.Join(staging, "roots.json"), bytes.NewReader(metadata), 0o600); err != nil { + return err + } + if _, err := s.readSnapshot(staging); err != nil { + return err + } + if err := os.WriteFile(filepath.Join(staging, "complete"), nil, 0o600); err != nil { + return err + } + if err := os.Rename(staging, snapshot); err != nil { + // Another process may have published this immutable snapshot first. + if _, statErr := os.Stat(filepath.Join(snapshot, "complete")); statErr != nil { + return fmt.Errorf("publishing GitHub snapshot: %w", err) + } + } + return nil +} + +func (s githubSource) archiveRoots(ctx context.Context, archive []byte) ([]string, error) { + var candidates []string + err := walkGitHubArchive(ctx, archive, func(header *tar.Header, name string, _ io.Reader) error { + if header.Typeflag != tar.TypeReg || path.Base(name) != skillFile { + return nil + } + if s.directory != "" && !withinGitHubPath(name, s.directory) { + return nil + } + candidates = append(candidates, path.Dir(name)) + return nil + }) + if err != nil { + return nil, err + } + if s.directory == "" { + for _, prefix := range []string{"skills", ".agents/skills", "."} { + var roots []string + for _, candidate := range candidates { + if prefix == "." && candidate == "." || prefix != "." && withinGitHubPath(candidate, prefix) { + roots = append(roots, candidate) + } + } + if len(roots) != 0 { + candidates = roots + break + } + if prefix == "." { + candidates = nil + } + } + } + if len(candidates) == 0 { + return nil, errors.New("no SKILL.md files found in the GitHub skill source") + } + slices.Sort(candidates) + return candidates, nil +} + +func withinGitHubPath(name, directory string) bool { + return directory == "." || name == directory || strings.HasPrefix(name, directory+"/") +} + +func walkGitHubArchive(ctx context.Context, archive []byte, visit func(*tar.Header, string, io.Reader) error) error { + compressed, err := gzip.NewReader(bytes.NewReader(archive)) + if err != nil { + return fmt.Errorf("opening GitHub archive: %w", err) + } + defer compressed.Close() + limited := &io.LimitedReader{R: compressed, N: githubExpandedLimit + 1} + reader := tar.NewReader(limited) + var prefix string + for entries := 0; ; entries++ { + if err := ctx.Err(); err != nil { + return err + } + header, err := reader.Next() + if limited.N <= 1 || entries > 100000 { + return errors.New("GitHub archive exceeds expansion limit") + } + if errors.Is(err, io.EOF) { + // tar EOF precedes the gzip trailer; drain to verify its checksum. + if _, err := io.Copy(io.Discard, limited); err != nil { + return fmt.Errorf("validating GitHub archive: %w", err) + } + if limited.N <= 1 { + return errors.New("GitHub archive exceeds expansion limit") + } + return nil + } + if err != nil { + return fmt.Errorf("reading GitHub archive: %w", err) + } + if header.Typeflag == tar.TypeXGlobalHeader { + continue + } + root, name, ok := strings.Cut(strings.TrimSuffix(header.Name, "/"), "/") + if !ok { + if header.Typeflag == tar.TypeDir && prefix == "" { + prefix = root + continue + } + return errors.New("invalid GitHub archive root") + } + if prefix == "" { + prefix = root + } + if root != prefix || name != path.Clean(name) || strings.HasPrefix(name, "/") || strings.Contains(name, "\\") || name == ".." || strings.HasPrefix(name, "../") { + return fmt.Errorf("unsafe GitHub archive path %q", header.Name) + } + if err := visit(header, name, reader); err != nil { + return err + } + } +} + +func extractGitHubSkills(ctx context.Context, archive []byte, roots []string, destination string) error { + seen := make(map[string]string) + var files int + var total int64 + return walkGitHubArchive(ctx, archive, func(header *tar.Header, name string, reader io.Reader) error { + selected := slices.ContainsFunc(roots, func(root string) bool { return withinGitHubPath(name, root) }) + if !selected { + return nil + } + if !validGitHubPath(name) { + return fmt.Errorf("unsafe GitHub skill path %q", name) + } + if header.Typeflag != tar.TypeDir && header.Typeflag != tar.TypeReg { + return fmt.Errorf("unsupported GitHub archive entry %q (links are not allowed)", name) + } + // Check every component, not just files, for case-insensitive collisions. + for component := name; component != "."; component = path.Dir(component) { + folded := strings.ToLower(component) + if original, ok := seen[folded]; ok && original != component { + return fmt.Errorf("colliding GitHub archive paths %q and %q", original, component) + } + seen[folded] = component + } + target := filepath.Join(destination, filepath.FromSlash(name)) + if header.Typeflag == tar.TypeDir { + return os.MkdirAll(target, 0o700) + } + files++ + total += header.Size + if header.Size < 0 || header.Size > githubFileLimit || total > githubArchiveLimit || files > githubFilesLimit { + return errors.New("GitHub skill files exceed download limits") + } + body, err := readGitHubBody(reader, githubFileLimit) + if err != nil { + return err + } + if err := os.MkdirAll(filepath.Dir(target), 0o700); err != nil { + return err + } + file, err := os.OpenFile(target, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600) + if err != nil { + return err + } + _, writeErr := file.Write(body) + closeErr := file.Close() + return errors.Join(writeErr, closeErr) + }) +} + +func (s githubSource) readSnapshot(snapshot string) ([]Skill, error) { + content := filepath.Join(snapshot, "content") + data, err := os.ReadFile(filepath.Join(snapshot, "roots.json")) + if err != nil { + return nil, err + } + var roots []string + if err := json.Unmarshal(data, &roots); err != nil { + return nil, err + } + var loaded []Skill + for _, root := range roots { + if root != "." && !validGitHubPath(root) { + return nil, errors.New("invalid cached GitHub skill root") + } + filename := filepath.Join(content, filepath.FromSlash(root), skillFile) + name := path.Base(root) + if root == "." { + name = path.Base(s.repository) + } + skill, ok := loadSkillFile(filename, name) + if !ok { + return nil, fmt.Errorf("invalid or missing GitHub skill %q", root) + } + loaded = append(loaded, skill) + } + if len(loaded) == 0 { + return nil, errors.New("GitHub skill source contains no valid skills") + } + names := make(map[string]bool) + for i := range loaded { + skill := &loaded[i] + if !isValidSkillName(skill.Name) || names[skill.Name] { + return nil, fmt.Errorf("invalid or duplicate GitHub skill name %q", skill.Name) + } + names[skill.Name] = true + skill.Local = false + // Remote metadata must not activate extra toolsets or switch providers. + skill.Model = "" + skill.Toolsets = nil + err := filepath.WalkDir(skill.BaseDir, func(filename string, entry fs.DirEntry, err error) error { + if err != nil { + return err + } + if entry.IsDir() { + return nil + } + relative, err := filepath.Rel(skill.BaseDir, filename) + if err != nil { + return err + } + skill.Files = append(skill.Files, filepath.ToSlash(relative)) + return nil + }) + if err != nil { + return nil, err + } + } + return loaded, nil +} diff --git a/pkg/skills/github_test.go b/pkg/skills/github_test.go new file mode 100644 index 000000000..ba0ed95f2 --- /dev/null +++ b/pkg/skills/github_test.go @@ -0,0 +1,754 @@ +package skills + +import ( + "archive/tar" + "bytes" + "compress/gzip" + "context" + "encoding/json" + "fmt" + "io" + "io/fs" + "net/http" + "net/http/httptest" + "net/url" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/paths" +) + +const githubTestCommit = "0123456789abcdef0123456789abcdef01234567" + +// Preserve canonical URLs and headers while routing requests to loopback. +type githubTestTransport struct { + base http.RoundTripper + serverURL *url.URL +} + +func (transport githubTestTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if req.URL.Scheme != "https" || (req.URL.Host != "api.github.com" && req.URL.Host != "codeload.github.com") { + return nil, fmt.Errorf("unexpected GitHub test endpoint %s", req.URL) + } + local := req.Clone(req.Context()) + local.URL.Scheme = transport.serverURL.Scheme + local.URL.Host = transport.serverURL.Host + local.Host = req.URL.Host + return transport.base.RoundTrip(local) +} + +func githubTestClient(t *testing.T, handler http.HandlerFunc) { + t.Helper() + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + serverURL, err := url.Parse(server.URL) + require.NoError(t, err) + previous := skillsHTTPClient + client := *server.Client() + client.Transport = githubTestTransport{base: client.Transport, serverURL: serverURL} + skillsHTTPClient = &client + t.Cleanup(func() { skillsHTTPClient = previous }) +} + +type githubTestEntry struct { + name string + body string + typeflag byte + linkname string + size int64 +} + +type githubZeroReader struct{} + +func (githubZeroReader) Read(data []byte) (int, error) { + clear(data) + return len(data), nil +} + +func githubTestArchive(t *testing.T, entries ...githubTestEntry) []byte { + t.Helper() + var buffer bytes.Buffer + compressed := gzip.NewWriter(&buffer) + writer := tar.NewWriter(compressed) + for _, entry := range entries { + kind := entry.typeflag + if kind == 0 { + kind = tar.TypeReg + } + size := max(int64(len(entry.body)), entry.size) + if kind != tar.TypeReg { + size = 0 + } + header := &tar.Header{Name: entry.name, Mode: 0o777, Typeflag: kind, Linkname: entry.linkname, Size: size} + if kind == tar.TypeXGlobalHeader { + header = &tar.Header{Name: entry.name, Typeflag: kind, PAXRecords: map[string]string{"comment": "GitHub snapshot"}} + } + require.NoError(t, writer.WriteHeader(header)) + if kind == tar.TypeReg { + _, err := io.WriteString(writer, entry.body) + require.NoError(t, err) + _, err = io.CopyN(writer, githubZeroReader{}, size-int64(len(entry.body))) + require.NoError(t, err) + } + } + require.NoError(t, writer.Close()) + require.NoError(t, compressed.Close()) + return buffer.Bytes() +} + +func githubTestSkill(name, description string) string { + return "---\nname: " + name + "\ndescription: " + description + "\n---\n\n!`echo untrusted`\n" +} + +func githubTestSnapshot(cache *diskCache, source githubSource, commit string) string { + return filepath.Join(cache.cacheDir("github:"+source.repository+"@"+commit, "snapshots"), githubDirectoryKey(source.directory)) +} + +func githubTestExpireResolution(t *testing.T, cache *diskCache, source githubSource) { + t.Helper() + filename := filepath.Join(cache.cacheDir(source.cacheKey(), "github"), "resolution.json") + data, err := os.ReadFile(filename) + require.NoError(t, err) + var resolution githubResolution + require.NoError(t, json.Unmarshal(data, &resolution)) + resolution.CheckedAt = time.Now().Add(-2 * githubRefTTL) + data, err = json.Marshal(resolution) + require.NoError(t, err) + require.NoError(t, os.WriteFile(filename, data, 0o600)) +} + +func TestGitHubSourceParsing(t *testing.T) { + valid := []struct { + url string + ref string + directory string + }{ + {url: "https://github.com/Owner/Repo"}, + {url: "https://github.com/Owner/Repo.git/"}, + {url: "https://GITHUB.COM/Owner/Repo/tree/main/skills", ref: "main", directory: "skills"}, + {url: "https://github.com/Owner/Repo/tree/feature%2Ffoo/skills/nested", ref: "feature/foo", directory: "skills/nested"}, + {url: "https://github.com/Owner/Repo?ref=feature%2Ffoo&path=skills", ref: "feature/foo", directory: "skills"}, + {url: "https://github.com/Owner/Repo/tree/v1.0.0", ref: "v1.0.0"}, + {url: "https://github.com/Owner/Repo?ref=" + githubTestCommit, ref: githubTestCommit}, + } + for _, test := range valid { + t.Run(test.url, func(t *testing.T) { + source, recognized, err := parseGitHubSource(test.url) + require.NoError(t, err) + require.True(t, recognized) + assert.Equal(t, githubSource{repository: "owner/repo", ref: test.ref, directory: test.directory}, source) + }) + } + invalid := []string{ + "http://github.com/owner/repo", "https://user:secret@github.com/owner/repo", + "https://github.com:443/owner/repo", "https://github.com/owner/repo#fragment", + "https://github.com/owner", "https://github.com/owner/repo/blob/main/SKILL.md", + "https://github.com/owner/repo/tree", "https://github.com/owner/repo/tree/main?ref=other", + "https://github.com/owner/repo?ref=", "https://github.com/owner/repo?path=", + "https://github.com/owner/repo?ref=a&ref=b", "https://github.com/owner/repo?token=secret", + "https://github.com/owner/repo?path=..%2Foutside", "https://github.com/owner/repo?path=%2Fabsolute", + "https://github.com/owner/repo?path=skills%2F..%2Foutside", "https://github.com/owner/repo?path=skills%5Coutside", + "https://github.com/owner/repo?ref=main%0Ainjected", "https://github.com/owner/repo?ref=with+space", + "https://github.com/owner/repo?path=skills%2FCON.txt", "https://github.com/owner/repo?path=skills%2Ftrailing.", + } + for _, source := range invalid { + t.Run(source, func(t *testing.T) { + _, recognized, err := parseGitHubSource(source) + assert.True(t, recognized) + require.Error(t, err) + }) + } + for _, source := range []string{"https://example.com/skills", "https://github.com.evil.test/owner/repo", "https://api.github.com/repos/owner/repo"} { + _, recognized, err := parseGitHubSource(source) + require.NoError(t, err) + assert.False(t, recognized) + } +} + +func TestGitHubLoadRefsAndWarmCache(t *testing.T) { + for _, ref := range []string{"", "main", "v1.2.3", "feature/foo", strings.ToUpper(githubTestCommit)} { + t.Run("ref="+ref, func(t *testing.T) { + archive := githubTestArchive(t, githubTestEntry{name: "repo-commit/skills/build/SKILL.md", body: githubTestSkill("build", "Build images")}) + var requests []string + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + requests = append(requests, r.Host+r.URL.EscapedPath()) + switch r.Host { + case "api.github.com": + assert.Equal(t, "Bearer provider-token", r.Header.Get("Authorization")) + assert.Equal(t, "application/vnd.github+json", r.Header.Get("Accept")) + assert.Equal(t, "2022-11-28", r.Header.Get("X-GitHub-Api-Version")) + if r.URL.Path == "/repos/owner/repo" { + fmt.Fprint(w, `{"private":false,"visibility":"public","default_branch":"trunk"}`) + } else { + fmt.Fprintf(w, `{"sha":%q}`, githubTestCommit) + } + case "codeload.github.com": + assert.Empty(t, r.Header.Get("Authorization")) + assert.Equal(t, "/owner/repo/tar.gz/"+githubTestCommit, r.URL.Path) + _, _ = w.Write(archive) + default: + http.NotFound(w, r) + } + }) + t.Setenv("GITHUB_TOKEN", "must-not-use-process-token") + env := environment.NewMapEnvProvider(map[string]string{"GITHUB_TOKEN": "provider-token"}) + cache := newDiskCache(t.TempDir()) + source := githubSource{repository: "owner/repo", ref: ref} + loaded, err := loadGitHubSkills(t.Context(), source, cache, env) + require.NoError(t, err) + require.Len(t, loaded, 1) + assert.Equal(t, "build", loaded[0].Name) + wantRequests := []string{"api.github.com/repos/owner/repo"} + if !isCommitSHA(ref) { + resolvedRef := ref + if resolvedRef == "" { + resolvedRef = "trunk" + } + wantRequests = append(wantRequests, "api.github.com/repos/owner/repo/commits/"+url.PathEscape(resolvedRef)) + } + wantRequests = append(wantRequests, "codeload.github.com/owner/repo/tar.gz/"+githubTestCommit) + assert.Equal(t, wantRequests, requests) + requests = nil + if isCommitSHA(ref) { + githubTestExpireResolution(t, cache, source) + } + warm, err := loadGitHubSkills(t.Context(), source, newDiskCache(cache.baseDir), environment.NewNoEnvProvider()) + require.NoError(t, err) + assert.Equal(t, loaded, warm) + assert.Empty(t, requests, "warm disk cache must perform no HTTP requests") + }) + } +} + +func TestGitHubRejectsNonPublicRepositories(t *testing.T) { + for _, repository := range []string{ + `{"private":true,"visibility":"private","default_branch":"main"}`, + `{"private":false,"visibility":"internal","default_branch":"main"}`, + `{"private":true,"visibility":"public","default_branch":"main"}`, + `{"private":false,"default_branch":"main"}`, + } { + t.Run(repository, func(t *testing.T) { + requests := 0 + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + requests++ + assert.Equal(t, "api.github.com", r.Host) + assert.Equal(t, "/repos/owner/repo", r.URL.Path) + assert.Equal(t, "Bearer private-access-token", r.Header.Get("Authorization")) + fmt.Fprint(w, repository) + }) + cache := newDiskCache(t.TempDir()) + env := environment.NewMapEnvProvider(map[string]string{"GITHUB_TOKEN": "private-access-token"}) + loaded, err := loadGitHubSkills(t.Context(), githubSource{repository: "owner/repo", ref: githubTestCommit}, cache, env) + require.ErrorContains(t, err, "public repository") + assert.Empty(t, loaded) + assert.Equal(t, 1, requests) + entries, err := os.ReadDir(cache.baseDir) + require.NoError(t, err) + assert.Empty(t, entries, "rejected repositories must not create a cache") + }) + } +} + +func TestGitHubMutableRefRefresh(t *testing.T) { + firstCommit := githubTestCommit + secondCommit := strings.Repeat("b", 40) + commit := firstCommit + firstArchive := githubTestArchive(t, githubTestEntry{name: "repo-first/skills/build/SKILL.md", body: githubTestSkill("build", "First revision")}) + secondArchive := githubTestArchive(t, githubTestEntry{name: "repo-second/skills/build/SKILL.md", body: githubTestSkill("build", "Second revision")}) + var requests []string + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + requests = append(requests, r.Host+r.URL.Path) + switch r.URL.Path { + case "/repos/owner/repo": + fmt.Fprint(w, `{"private":false,"visibility":"public","default_branch":"main"}`) + case "/repos/owner/repo/commits/main": + fmt.Fprintf(w, `{"sha":%q}`, commit) + case "/owner/repo/tar.gz/" + firstCommit: + _, _ = w.Write(firstArchive) + case "/owner/repo/tar.gz/" + secondCommit: + _, _ = w.Write(secondArchive) + default: + http.NotFound(w, r) + } + }) + cache := newDiskCache(t.TempDir()) + source := githubSource{repository: "owner/repo", ref: "main"} + first, err := loadGitHubSkills(t.Context(), source, cache, nil) + require.NoError(t, err) + require.Len(t, first, 1) + commit = secondCommit + githubTestExpireResolution(t, cache, source) + requests = nil + second, err := loadGitHubSkills(t.Context(), source, cache, nil) + require.NoError(t, err) + require.Len(t, second, 1) + assert.Equal(t, "Second revision", second[0].Description) + assert.NotEqual(t, first[0].FilePath, second[0].FilePath) + oldContent, err := os.ReadFile(first[0].FilePath) + require.NoError(t, err) + assert.Contains(t, string(oldContent), "First revision") + require.Len(t, requests, 3) + githubTestExpireResolution(t, cache, source) + requests = nil + _, err = loadGitHubSkills(t.Context(), source, cache, nil) + require.NoError(t, err) + assert.Equal(t, []string{"api.github.com/repos/owner/repo", "api.github.com/repos/owner/repo/commits/main"}, requests, "unchanged commit reuses its immutable snapshot") +} + +func TestGitHubArchiveRootSelection(t *testing.T) { + entries := []githubTestEntry{ + {name: "repo-commit/SKILL.md", body: githubTestSkill("root", "Root skill")}, + {name: "repo-commit/.agents/skills/agent/SKILL.md", body: githubTestSkill("agent", "Agent skill")}, + {name: "repo-commit/skills/zeta/SKILL.md", body: githubTestSkill("zeta", "Zeta skill")}, + {name: "repo-commit/skills/alpha/SKILL.md", body: githubTestSkill("alpha", "Alpha skill")}, + {name: "repo-commit/custom/selected/SKILL.md", body: githubTestSkill("selected", "Selected skill")}, + {name: "repo-commit/custom/selected-neighbor/SKILL.md", body: githubTestSkill("neighbor", "Must not select")}, + } + for _, test := range []struct { + name string + entries []githubTestEntry + directory string + roots []string + }{ + {name: "skills wins", entries: entries, roots: []string{"skills/alpha", "skills/zeta"}}, + {name: "agents fallback", entries: entries[:2], roots: []string{".agents/skills/agent"}}, + {name: "root fallback", entries: entries[:1], roots: []string{"."}}, + {name: "selected skill", entries: entries, directory: "custom/selected", roots: []string{"custom/selected"}}, + {name: "selected collection", entries: entries, directory: "custom", roots: []string{"custom/selected", "custom/selected-neighbor"}}, + {name: "missing selection", entries: entries, directory: "missing"}, + {name: "no arbitrary fallback", entries: entries[4:]}, + } { + t.Run(test.name, func(t *testing.T) { + roots, err := (githubSource{directory: test.directory}).archiveRoots(t.Context(), githubTestArchive(t, test.entries...)) + if test.roots == nil { + require.ErrorContains(t, err, "no SKILL.md") + return + } + require.NoError(t, err) + assert.Equal(t, test.roots, roots) + }) + } +} + +func TestGitHubSnapshotFieldsAndFiles(t *testing.T) { + content := "---\ndescription: Build safely\nlicense: Apache-2.0\ncompatibility: Requires Docker\nmetadata:\n author: upstream\nallowed-tools: Read, Grep\ncontext: fork\nmodel: openai/untrusted\ntoolsets: shell, secrets\n---\n\n!`echo untrusted`\n" + archive := githubTestArchive(t, + githubTestEntry{name: "repo-commit/skills/build/SKILL.md", body: content}, + githubTestEntry{name: "repo-commit/skills/build/references/guide.md", body: "Reference contents"}, + githubTestEntry{name: "repo-commit/skills/build/scripts/build.sh", body: "#!/bin/sh\necho build\n"}, + githubTestEntry{name: "repo-commit/README.md", body: "not part of the skill"}, + githubTestEntry{name: "repo-commit/unrelated/link", typeflag: tar.TypeSymlink, linkname: "/etc/passwd"}, + ) + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "codeload.github.com", r.Host) + _, _ = w.Write(archive) + }) + source := githubSource{repository: "owner/repo"} + snapshot := filepath.Join(t.TempDir(), "snapshot") + require.NoError(t, source.download(t.Context(), githubTestCommit, snapshot)) + loaded, err := source.readSnapshot(snapshot) + require.NoError(t, err) + require.Len(t, loaded, 1) + skill := loaded[0] + assert.Equal(t, "build", skill.Name) + assert.Equal(t, "Build safely", skill.Description) + assert.False(t, skill.Local) + assert.False(t, skill.ExpandsCommands()) + assert.False(t, skill.IsInline()) + assert.True(t, skill.IsFork()) + assert.Empty(t, skill.Model) + assert.Nil(t, skill.Toolsets) + assert.Equal(t, "Apache-2.0", skill.License) + assert.Equal(t, "Requires Docker", skill.Compatibility) + assert.Equal(t, map[string]string{"author": "upstream"}, skill.Metadata) + assert.Equal(t, []string{"Read", "Grep"}, skill.AllowedTools) + assert.Equal(t, []string{"SKILL.md", "references/guide.md", "scripts/build.sh"}, skill.Files) + assert.Equal(t, filepath.Join(snapshot, "content", "skills", "build"), skill.BaseDir) + assert.Equal(t, filepath.Join(skill.BaseDir, skillFile), skill.FilePath) + body, err := os.ReadFile(skill.FilePath) + require.NoError(t, err) + assert.Equal(t, content, string(body)) + reference, err := os.ReadFile(filepath.Join(skill.BaseDir, "references", "guide.md")) + require.NoError(t, err) + assert.Equal(t, "Reference contents", string(reference)) + _, err = os.Stat(filepath.Join(snapshot, "complete")) + require.NoError(t, err) + _, err = os.Stat(filepath.Join(snapshot, "content", "README.md")) + require.ErrorIs(t, err, os.ErrNotExist) +} + +func TestGitHubRejectsUnsafeArchives(t *testing.T) { + base := githubTestEntry{name: "repo-commit/skills/build/SKILL.md", body: githubTestSkill("build", "Safe skill")} + for _, test := range []struct { + name string + entry githubTestEntry + }{ + {name: "parent traversal", entry: githubTestEntry{name: "repo-commit/../escaped", body: "bad"}}, + {name: "nested traversal", entry: githubTestEntry{name: "repo-commit/skills/build/../../escaped", body: "bad"}}, + {name: "absolute path", entry: githubTestEntry{name: "/absolute/escaped", body: "bad"}}, + {name: "backslash", entry: githubTestEntry{name: `repo-commit/skills/build/..\escaped`, body: "bad"}}, + {name: "different archive root", entry: githubTestEntry{name: "other-root/skills/build/extra", body: "bad"}}, + {name: "symlink", entry: githubTestEntry{name: "repo-commit/skills/build/link", typeflag: tar.TypeSymlink, linkname: "/etc/passwd"}}, + {name: "hardlink", entry: githubTestEntry{name: "repo-commit/skills/build/link", typeflag: tar.TypeLink, linkname: "repo-commit/skills/build/SKILL.md"}}, + {name: "fifo", entry: githubTestEntry{name: "repo-commit/skills/build/pipe", typeflag: tar.TypeFifo}}, + {name: "windows device", entry: githubTestEntry{name: "repo-commit/skills/build/NUL.txt", body: "bad"}}, + {name: "trailing dot", entry: githubTestEntry{name: "repo-commit/skills/build/file.", body: "bad"}}, + {name: "alternate data stream", entry: githubTestEntry{name: "repo-commit/skills/build/file:stream", body: "bad"}}, + {name: "case file collision", entry: githubTestEntry{name: "repo-commit/skills/build/skill.md", body: "bad"}}, + {name: "case directory collision", entry: githubTestEntry{name: "repo-commit/skills/Build/SKILL.md", body: githubTestSkill("other", "Case collision")}}, + {name: "duplicate file", entry: base}, + {name: "oversized file", entry: githubTestEntry{name: "repo-commit/skills/build/huge", size: githubFileLimit + 1}}, + } { + t.Run(test.name, func(t *testing.T) { + archive := githubTestArchive(t, base, test.entry) + githubTestClient(t, func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write(archive) }) + snapshot := filepath.Join(t.TempDir(), "snapshot") + err := (githubSource{repository: "owner/repo"}).download(t.Context(), githubTestCommit, snapshot) + require.Error(t, err) + _, err = os.Stat(snapshot) + require.ErrorIs(t, err, os.ErrNotExist, "invalid archives must not publish a snapshot") + staging, err := filepath.Glob(filepath.Join(filepath.Dir(snapshot), ".download-*")) + require.NoError(t, err) + assert.Empty(t, staging, "failed extraction must clean its staging directory") + }) + } +} + +func TestGitHubArchiveLimits(t *testing.T) { + base := githubTestEntry{name: "repo-commit/skills/build/SKILL.md", body: githubTestSkill("build", "Safe skill")} + t.Run("file count", func(t *testing.T) { + entries := []githubTestEntry{base} + for i := range githubFilesLimit { + entries = append(entries, githubTestEntry{name: fmt.Sprintf("repo-commit/skills/build/file-%04d", i)}) + } + err := extractGitHubSkills(t.Context(), githubTestArchive(t, entries...), []string{"skills/build"}, t.TempDir()) + require.ErrorContains(t, err, "download limits") + }) + t.Run("total selected size", func(t *testing.T) { + entries := []githubTestEntry{base} + for i := range githubArchiveLimit / githubFileLimit { + entries = append(entries, githubTestEntry{name: fmt.Sprintf("repo-commit/skills/build/file-%02d", i), size: githubFileLimit}) + } + err := extractGitHubSkills(t.Context(), githubTestArchive(t, entries...), []string{"skills/build"}, t.TempDir()) + require.ErrorContains(t, err, "download limits") + }) + t.Run("unselected expansion bomb", func(t *testing.T) { + archive := githubTestArchive(t, base, githubTestEntry{name: "repo-commit/unselected/bomb", size: githubExpandedLimit + 1}) + _, err := (githubSource{}).archiveRoots(t.Context(), archive) + require.ErrorContains(t, err, "expansion limit") + }) + t.Run("body size boundary", func(t *testing.T) { + data, err := readGitHubBody(strings.NewReader("1234"), 4) + require.NoError(t, err) + assert.Equal(t, []byte("1234"), data) + _, err = readGitHubBody(strings.NewReader("12345"), 4) + require.ErrorContains(t, err, "size limit") + }) +} + +func TestGitHubInvalidSkillsDoNotPublish(t *testing.T) { + for _, test := range []struct { + name string + entries []githubTestEntry + }{ + {name: "no frontmatter", entries: []githubTestEntry{{name: "repo/skills/build/SKILL.md", body: "# Not a skill"}}}, + {name: "unclosed frontmatter", entries: []githubTestEntry{{name: "repo/skills/build/SKILL.md", body: "---\ndescription: Invalid\n"}}}, + {name: "short malformed frontmatter", entries: []githubTestEntry{{name: "repo/skills/build/SKILL.md", body: "---\n---"}}}, + {name: "missing description", entries: []githubTestEntry{{name: "repo/skills/build/SKILL.md", body: "---\nname: build\n---\n"}}}, + {name: "invalid name", entries: []githubTestEntry{{name: "repo/skills/build/SKILL.md", body: githubTestSkill("../escape", "Invalid name")}}}, + {name: "duplicate names", entries: []githubTestEntry{ + {name: "repo/skills/first/SKILL.md", body: githubTestSkill("same", "First")}, + {name: "repo/skills/second/SKILL.md", body: githubTestSkill("same", "Second")}, + }}, + } { + t.Run(test.name, func(t *testing.T) { + archive := githubTestArchive(t, test.entries...) + githubTestClient(t, func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write(archive) }) + snapshot := filepath.Join(t.TempDir(), "snapshot") + require.Error(t, (githubSource{repository: "owner/repo"}).download(t.Context(), githubTestCommit, snapshot)) + _, err := os.Stat(snapshot) + require.ErrorIs(t, err, os.ErrNotExist) + }) + } +} + +func TestGitHubDownloadFailureCanRetry(t *testing.T) { + archive := githubTestArchive(t, githubTestEntry{name: "repo/skills/build/SKILL.md", body: githubTestSkill("build", "Build images")}) + fail := true + var requests []string + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + requests = append(requests, r.Host+r.URL.Path) + switch { + case r.Host == "api.github.com": + fmt.Fprint(w, `{"private":false,"visibility":"public","default_branch":"main"}`) + case fail: + _, _ = w.Write([]byte("not a gzip archive")) + default: + _, _ = w.Write(archive) + } + }) + cache := newDiskCache(t.TempDir()) + source := githubSource{repository: "owner/repo", ref: githubTestCommit} + _, err := loadGitHubSkills(t.Context(), source, cache, nil) + require.Error(t, err) + _, err = os.Stat(githubTestSnapshot(cache, source, githubTestCommit)) + require.ErrorIs(t, err, os.ErrNotExist) + fail = false + requests = nil + loaded, err := loadGitHubSkills(t.Context(), source, cache, nil) + require.NoError(t, err) + require.Len(t, loaded, 1) + assert.Equal(t, []string{"codeload.github.com/owner/repo/tar.gz/" + githubTestCommit}, requests) +} + +func TestGitHubAPIFailuresAndRedirects(t *testing.T) { + for _, test := range []struct { + name string + status int + body string + want string + }{ + {name: "rate limit", status: http.StatusForbidden, want: "GITHUB_TOKEN"}, + {name: "too many requests", status: http.StatusTooManyRequests, want: "rate limits"}, + {name: "missing repository", status: http.StatusNotFound, want: "HTTP 404"}, + {name: "redirect rejected", status: http.StatusFound, want: "HTTP 302"}, + {name: "malformed JSON", status: http.StatusOK, body: "{", want: "decoding GitHub response"}, + {name: "missing default branch", status: http.StatusOK, body: `{"private":false,"visibility":"public"}`, want: "no default branch"}, + {name: "oversized API response", status: http.StatusOK, body: strings.Repeat(" ", (1<<20)+1), want: "size limit"}, + } { + t.Run(test.name, func(t *testing.T) { + requests := 0 + githubTestClient(t, func(w http.ResponseWriter, _ *http.Request) { + requests++ + w.Header().Set("Location", "https://api.github.com/redirected") + w.WriteHeader(test.status) + fmt.Fprint(w, test.body) + }) + cache := newDiskCache(t.TempDir()) + _, err := loadGitHubSkills(t.Context(), githubSource{repository: "owner/repo"}, cache, nil) + require.ErrorContains(t, err, test.want) + assert.Equal(t, 1, requests, "redirects must never be followed") + entries, err := os.ReadDir(cache.baseDir) + require.NoError(t, err) + assert.Empty(t, entries) + }) + } +} + +func TestGitHubConcurrentLoads(t *testing.T) { + archive := githubTestArchive(t, githubTestEntry{name: "repo/skills/build/SKILL.md", body: githubTestSkill("build", "Build images")}) + var mutex sync.Mutex + requests := 0 + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + mutex.Lock() + requests++ + mutex.Unlock() + switch r.Host { + case "api.github.com": + fmt.Fprint(w, `{"private":false,"visibility":"public","default_branch":"main"}`) + case "codeload.github.com": + _, _ = w.Write(archive) + } + }) + cache := newDiskCache(t.TempDir()) + source := githubSource{repository: "owner/repo", ref: githubTestCommit} + type result struct { + skills []Skill + err error + } + results := make(chan result, 12) + start := make(chan struct{}) + for range cap(results) { + go func() { + <-start + loaded, err := loadGitHubSkills(t.Context(), source, cache, nil) + results <- result{skills: loaded, err: err} + }() + } + close(start) + for range cap(results) { + loaded := <-results + require.NoError(t, loaded.err) + require.Len(t, loaded.skills, 1) + } + mutex.Lock() + defer mutex.Unlock() + assert.Equal(t, 2, requests, "concurrent identical sources share public check and archive download") +} + +func TestGitHubCanceledArchiveWalk(t *testing.T) { + archive := githubTestArchive(t, githubTestEntry{name: "repo/skills/build/SKILL.md", body: githubTestSkill("build", "Build images")}) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + _, err := (githubSource{}).archiveRoots(ctx, archive) + require.ErrorIs(t, err, context.Canceled) + err = extractGitHubSkills(ctx, archive, []string{"skills/build"}, t.TempDir()) + require.ErrorIs(t, err, context.Canceled) +} + +func TestGitHubLoadWithWarningsIntegration(t *testing.T) { + previousCache := paths.GetCacheDir() + paths.SetCacheDir(t.TempDir()) + t.Cleanup(func() { paths.SetCacheDir(previousCache) }) + archive := githubTestArchive(t, + githubTestEntry{name: "repo/skills/zeta/SKILL.md", body: githubTestSkill("zeta", "Zeta")}, + githubTestEntry{name: "repo/skills/alpha/SKILL.md", body: githubTestSkill("alpha", "Alpha")}, + ) + requests := 0 + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + requests++ + switch r.Host { + case "api.github.com": + assert.Equal(t, "Bearer integration-token", r.Header.Get("Authorization")) + fmt.Fprint(w, `{"private":false,"visibility":"public"}`) + case "codeload.github.com": + assert.Empty(t, r.Header.Get("Authorization")) + _, _ = w.Write(archive) + } + }) + env := environment.NewMapEnvProvider(map[string]string{"GITHUB_TOKEN": "integration-token"}) + valid := "https://github.com/owner/repo?ref=" + githubTestCommit + invalid := "https://github.com/owner/repo?path=../outside" + loaded, warnings := LoadWithWarnings(t.Context(), []string{invalid, valid}, env) + require.Len(t, loaded, 2) + assert.Equal(t, "alpha", loaded[0].Name) + assert.Equal(t, "zeta", loaded[1].Name) + require.Len(t, warnings, 1) + assert.Contains(t, warnings[0], invalid) + assert.NotContains(t, warnings[0], "integration-token") + assert.Equal(t, 2, requests) +} + +func TestGitHubInvalidCachedRoots(t *testing.T) { + for _, roots := range []string{`["../outside"]`, `["/absolute"]`, `["skills\\escape"]`, `null`, `[]`, `{}`} { + t.Run(roots, func(t *testing.T) { + snapshot := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(snapshot, "roots.json"), []byte(roots), 0o600)) + _, err := (githubSource{repository: "owner/repo"}).readSnapshot(snapshot) + require.Error(t, err) + }) + } +} + +func TestGitHubBodyReadFailure(t *testing.T) { + reader := io.MultiReader(strings.NewReader("partial"), githubErrorReader{}) + _, err := readGitHubBody(reader, 100) + require.ErrorIs(t, err, fs.ErrInvalid) +} + +type githubErrorReader struct{} + +func (githubErrorReader) Read([]byte) (int, error) { return 0, fs.ErrInvalid } + +func TestGitHubInvalidResolvedCommit(t *testing.T) { + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/repos/owner/repo" { + fmt.Fprint(w, `{"private":false,"visibility":"public","default_branch":"main"}`) + } else { + fmt.Fprint(w, `{"sha":"not-a-full-sha"}`) + } + }) + cache := newDiskCache(t.TempDir()) + _, err := loadGitHubSkills(t.Context(), githubSource{repository: "owner/repo"}, cache, nil) + require.ErrorContains(t, err, "invalid commit SHA") + _, err = os.Stat(filepath.Join(cache.cacheDir((githubSource{repository: "owner/repo"}).cacheKey(), "github"), "resolution.json")) + require.ErrorIs(t, err, os.ErrNotExist) +} + +func TestGitHubSelectedDirectoriesHaveSeparateSnapshots(t *testing.T) { + archive := githubTestArchive(t, + githubTestEntry{name: "repo/skills/first/SKILL.md", body: githubTestSkill("first", "First selection")}, + githubTestEntry{name: "repo/skills/second/SKILL.md", body: githubTestSkill("second", "Second selection")}, + ) + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Host == "api.github.com" { + fmt.Fprint(w, `{"private":false,"visibility":"public"}`) + } else { + _, _ = w.Write(archive) + } + }) + cache := newDiskCache(t.TempDir()) + var selections []Skill + for _, directory := range []string{"skills/first", "skills/second"} { + source := githubSource{repository: "owner/repo", ref: githubTestCommit, directory: directory} + loaded, err := loadGitHubSkills(t.Context(), source, cache, nil) + require.NoError(t, err) + require.Len(t, loaded, 1) + selections = append(selections, loaded[0]) + } + assert.Equal(t, "first", selections[0].Name) + assert.Equal(t, "second", selections[1].Name) + assert.NotEqual(t, selections[0].BaseDir, selections[1].BaseDir) +} + +func TestGitHubCanceledWaiterDoesNotCancelActiveLoad(t *testing.T) { + archive := githubTestArchive(t, githubTestEntry{name: "repo/skills/build/SKILL.md", body: githubTestSkill("build", "Build images")}) + entered := make(chan struct{}) + release := make(chan struct{}) + var unblock sync.Once + defer unblock.Do(func() { close(release) }) + githubTestClient(t, func(w http.ResponseWriter, r *http.Request) { + if r.Host == "api.github.com" { + close(entered) + select { + case <-release: + case <-r.Context().Done(): + return + } + fmt.Fprint(w, `{"private":false,"visibility":"public"}`) + } else { + _, _ = w.Write(archive) + } + }) + cache := newDiskCache(t.TempDir()) + source := githubSource{repository: "owner/repo", ref: githubTestCommit} + result := make(chan error, 1) + go func() { + _, err := loadGitHubSkills(t.Context(), source, cache, nil) + result <- err + }() + select { + case <-entered: + case <-time.After(5 * time.Second): + t.Fatal("first load did not reach GitHub") + } + ctx, cancel := context.WithCancel(t.Context()) + cancel() + _, err := loadGitHubSkills(ctx, source, cache, nil) + require.ErrorIs(t, err, context.Canceled) + unblock.Do(func() { close(release) }) + select { + case err := <-result: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("active load did not finish after releasing GitHub") + } +} + +func TestGitHubArchiveChecksum(t *testing.T) { + t.Parallel() + archive := githubTestArchive(t, githubTestEntry{name: "repo/skills/example/SKILL.md", body: githubTestSkill("example", "Example")}) + archive[len(archive)-8] ^= 0xff + _, err := (githubSource{repository: "owner/repo"}).archiveRoots(t.Context(), archive) + require.ErrorContains(t, err, "invalid checksum") +} + +func TestGitHubArchiveGlobalHeader(t *testing.T) { + t.Parallel() + archive := githubTestArchive(t, + githubTestEntry{name: "pax_global_header", typeflag: tar.TypeXGlobalHeader}, + githubTestEntry{name: "repo/", typeflag: tar.TypeDir}, + githubTestEntry{name: "repo/skills/example/SKILL.md", body: githubTestSkill("example", "Example")}, + ) + roots, err := (githubSource{repository: "owner/repo"}).archiveRoots(t.Context(), archive) + require.NoError(t, err) + assert.Equal(t, []string{"skills/example"}, roots) +} diff --git a/pkg/skills/skills.go b/pkg/skills/skills.go index fff85df04..15c08cf2e 100644 --- a/pkg/skills/skills.go +++ b/pkg/skills/skills.go @@ -2,11 +2,14 @@ package skills import ( "context" + "fmt" + "log/slog" "maps" "path/filepath" "slices" "strings" + "github.com/docker/docker-agent/pkg/environment" "github.com/docker/docker-agent/pkg/paths" ) @@ -65,7 +68,7 @@ func (s Skill) ExpandsCommands() bool { // Load discovers and loads skills from the given sources. // Each source is either "local" (for filesystem-based skills) or an HTTP/HTTPS -// URL (for remote skills per the well-known skills discovery spec). +// URL (a public GitHub repository or a well-known skills discovery endpoint). // // Local skills are loaded from (in order, later overrides earlier): // @@ -83,7 +86,18 @@ func (s Skill) ExpandsCommands() bool { // // The returned slice is sorted by skill name for deterministic ordering. func Load(ctx context.Context, sources []string) []Skill { + loaded, warnings := LoadWithWarnings(ctx, sources, environment.NewOsEnvProvider()) + for _, warning := range warnings { + slog.WarnContext(ctx, warning) + } + return loaded +} + +// LoadWithWarnings loads skills and reports GitHub source failures to the caller. +// The environment provider supplies the optional GITHUB_TOKEN credential. +func LoadWithWarnings(ctx context.Context, sources []string, env environment.Provider) ([]Skill, []string) { skillMap := make(map[string]Skill) + var warnings []string var remoteCache *diskCache for _, source := range sources { @@ -94,7 +108,22 @@ func Load(ctx context.Context, sources []string) []Skill { if remoteCache == nil { remoteCache = newDiskCache(filepath.Join(paths.GetCacheDir(), "skills")) } - for _, skill := range loadRemoteSkills(ctx, source, remoteCache) { + github, recognized, err := parseGitHubSource(source) + if err != nil { + warnings = append(warnings, fmt.Sprintf("GitHub skill source %s: %v", source, err)) + continue + } + var loaded []Skill + if recognized { + loaded, err = loadGitHubSkills(ctx, github, remoteCache, env) + if err != nil { + warnings = append(warnings, fmt.Sprintf("GitHub skill source %s: %v", source, err)) + continue + } + } else { + loaded = loadRemoteSkills(ctx, source, remoteCache) + } + for _, skill := range loaded { skillMap[source+"/"+skill.Name] = skill } } @@ -108,7 +137,7 @@ func Load(ctx context.Context, sources []string) []Skill { // deterministic ordering even when a local and a remote source // expose a skill with the same name. return strings.Compare(a.FilePath, b.FilePath) - }) + }), warnings } // isHTTPSource reports whether s is an HTTP(S) URL source. diff --git a/pkg/teamloader/teamloader.go b/pkg/teamloader/teamloader.go index 3409ff16b..d848b0b93 100644 --- a/pkg/teamloader/teamloader.go +++ b/pkg/teamloader/teamloader.go @@ -533,7 +533,10 @@ func LoadWithConfig(ctx context.Context, agentSource config.Source, runConfig *c // Add skills toolset if skills are enabled if agentConfig.Skills.Enabled() { - loadedSkills := skills.Load(ctx, agentConfig.Skills.Sources) + loadedSkills, skillWarnings := skills.LoadWithWarnings(ctx, agentConfig.Skills.Sources, env) + if len(skillWarnings) > 0 { + opts = append(opts, agent.WithLoadTimeWarnings(skillWarnings)) + } loadedSkills = filterSkillsByName(loadedSkills, agentConfig.Skills.Include) // Inline skills are defined in the agent config itself; they are // always exposed and never subject to the include filter.