Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 23 additions & 26 deletions pkg/recipe/provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -174,15 +174,15 @@ func (p *EmbeddedDataProvider) ReadFile(ctx context.Context, path string) ([]byt
if err := ctx.Err(); err != nil {
return nil, aicrerrors.Wrap(aicrerrors.ErrCodeTimeout, fmt.Sprintf("context canceled before reading %q", path), err)
}
fullPath := filepath.Join(p.prefix, path)
fullPath := filepath.ToSlash(filepath.Clean(filepath.Join(p.prefix, path)))
slog.Debug("reading file from embedded provider", "path", path, "fullPath", fullPath)
return p.fs.ReadFile(fullPath)
}

// WalkDir walks the embedded filesystem.
func (p *EmbeddedDataProvider) WalkDir(ctx context.Context, root string, fn fs.WalkDirFunc) error {
fullRoot := filepath.Join(p.prefix, root)
if fullRoot == "" {
fullRoot := filepath.ToSlash(filepath.Clean(filepath.Join(p.prefix, root)))
if fullRoot == "" || fullRoot == "." {
fullRoot = "." // embed.FS expects "." for root
}
slog.Debug("walking embedded filesystem", "root", root, "fullRoot", fullRoot)
Expand All @@ -195,11 +195,12 @@ func (p *EmbeddedDataProvider) WalkDir(ctx context.Context, root string, fn fs.W
}
// Strip the prefix before passing to callback
var relPath string
if p.prefix == "" {
normalizedPrefix := filepath.ToSlash(p.prefix)
if normalizedPrefix == "" || normalizedPrefix == "." {
relPath = path
} else {
relPath = strings.TrimPrefix(path, p.prefix+"/")
if relPath == p.prefix {
relPath = strings.TrimPrefix(path, normalizedPrefix+"/")
if relPath == normalizedPrefix {
relPath = ""
}
}
Expand Down Expand Up @@ -374,7 +375,7 @@ func NewLayeredDataProvider(embedded *EmbeddedDataProvider, config LayeredProvid
fmt.Sprintf("file too large (%d bytes, max %d): %s", info.Size(), config.MaxFileSize, relPath))
}

externalFiles[relPath] = true
externalFiles[filepath.ToSlash(relPath)] = true
slog.Debug("discovered external file",
"path", relPath,
"size", info.Size())
Expand Down Expand Up @@ -427,21 +428,23 @@ func (p *LayeredDataProvider) ReadFile(ctx context.Context, path string) ([]byte
}
slog.Debug("reading file from layered provider", "path", path)

lookupPath := filepath.ToSlash(path)

// Special handling for registry file - merge instead of replace
if path == registryFileName {
if lookupPath == registryFileName {
slog.Debug("reading merged registry file")
return p.getMergedRegistry(ctx)
}

// Special handling for catalog file - merge instead of replace (when external exists)
if path == catalogFileName && p.externalFiles[catalogFileName] {
if lookupPath == catalogFileName && p.externalFiles[catalogFileName] {
slog.Debug("reading merged catalog file")
return p.getMergedCatalog(ctx)
}

// Check external directory first
if p.externalFiles[path] {
data, err := readExternalFile(p.externalDir, path, p.maxFileSize, p.allowSymlinks)
if p.externalFiles[lookupPath] {
data, err := readExternalFile(p.externalDir, lookupPath, p.maxFileSize, p.allowSymlinks)
if err != nil {
return nil, aicrerrors.PropagateOrWrap(err, aicrerrors.ErrCodeInternal, fmt.Sprintf("failed to read external file %s", path))
}
Expand All @@ -451,7 +454,7 @@ func (p *LayeredDataProvider) ReadFile(ctx context.Context, path string) ([]byte

// Fall back to embedded
slog.Debug("falling back to embedded data", "path", path)
return p.embedded.ReadFile(ctx, path)
return p.embedded.ReadFile(ctx, lookupPath)
}

// WalkDir walks both embedded and external directories.
Expand Down Expand Up @@ -484,18 +487,11 @@ func (p *LayeredDataProvider) WalkDir(ctx context.Context, root string, fn fs.Wa
if relErr != nil {
return aicrerrors.Wrap(aicrerrors.ErrCodeInternal, "failed to compute relative path", relErr)
}
slashRelPath := filepath.ToSlash(relPath)
Comment thread
SatyamPandey-07 marked this conversation as resolved.

// Strip root prefix if present
if root != "" {
relPath = strings.TrimPrefix(relPath, root+"/")
if relPath == root {
relPath = ""
}
}

visited[relPath] = true
slog.Debug("visiting external file", "path", relPath, "isDir", d.IsDir())
return fn(relPath, d, nil)
visited[slashRelPath] = true
slog.Debug("visiting external file", "path", slashRelPath, "isDir", d.IsDir())
return fn(slashRelPath, d, nil)
})
if err != nil {
return aicrerrors.PropagateOrWrap(err, aicrerrors.ErrCodeInternal,
Expand All @@ -521,15 +517,16 @@ func (p *LayeredDataProvider) WalkDir(ctx context.Context, root string, fn fs.Wa

// Source returns "external" or "embedded" depending on where the file comes from.
func (p *LayeredDataProvider) Source(path string) string {
normalizedPath := filepath.ToSlash(path)
var source string
switch {
case path == registryFileName:
case normalizedPath == registryFileName:
// Always merged: registry.yaml is required in external dir (enforced by constructor).
source = sourceMerged
case path == catalogFileName && p.externalFiles[catalogFileName]:
case normalizedPath == catalogFileName && p.externalFiles[catalogFileName]:
// Merged only when external catalog exists (catalog is optional).
source = sourceMerged
case p.externalFiles[path]:
case p.externalFiles[normalizedPath]:
source = sourceExternal
default:
source = sourceEmbedded
Expand Down
102 changes: 102 additions & 0 deletions pkg/recipe/provider_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,35 @@ func TestEmbeddedDataProvider(t *testing.T) {
}
})

t.Run("read host-native nested path", func(t *testing.T) {
data, err := provider.ReadFile(context.Background(), filepath.Join("overlays", "base.yaml"))
if err != nil {
t.Fatalf("failed to read nested overlays/base.yaml: %v", err)
}
if len(data) == 0 {
t.Error("overlays/base.yaml is empty")
}
})

t.Run("walk embedded directory with host-native root", func(t *testing.T) {
foundBase := false
err := provider.WalkDir(context.Background(), "overlays", func(path string, d fs.DirEntry, err error) error {
if err != nil {
return err
}
if strings.HasSuffix(path, "base.yaml") {
foundBase = true
}
return nil
})
if err != nil {
t.Fatalf("WalkDir failed: %v", err)
}
if !foundBase {
t.Error("expected to find base.yaml in overlays")
}
})

t.Run("source returns embedded", func(t *testing.T) {
source := provider.Source("registry.yaml")
if source != sourceEmbedded {
Expand Down Expand Up @@ -788,6 +817,79 @@ func TestLayeredDataProvider_SourceForRegistry(t *testing.T) {
}
}

// TestLayeredDataProvider_SourceHostNativeNestedPath tests that Source resolves nested catalog paths
// when formatted using the host-native path separator.
func TestLayeredDataProvider_SourceHostNativeNestedPath(t *testing.T) {
tmpDir := t.TempDir()
if err := os.WriteFile(filepath.Join(tmpDir, "registry.yaml"), []byte(testEmptyRegistryContent), 0600); err != nil {
t.Fatalf("failed to write registry.yaml: %v", err)
}
validatorsDir := filepath.Join(tmpDir, "validators")
if err := os.MkdirAll(validatorsDir, 0755); err != nil {
t.Fatalf("failed to create validators dir: %v", err)
}
if err := os.WriteFile(filepath.Join(validatorsDir, "catalog.yaml"), []byte("validators: []\n"), 0600); err != nil {
t.Fatalf("failed to write catalog.yaml: %v", err)
}

embedded := NewEmbeddedDataProvider(GetEmbeddedFS(), ".")
provider, err := NewLayeredDataProvider(embedded, LayeredProviderConfig{
ExternalDir: tmpDir,
})
if err != nil {
t.Fatalf("failed to create layered provider: %v", err)
}

// Test host-native nested path
source := provider.Source(filepath.Join("validators", "catalog.yaml"))
if source != sourceMerged {
t.Errorf("expected source %q, got %q", sourceMerged, source)
}
}

// TestLayeredDataProvider_WalkDirDeduplication verifies that external overrides retain root prefix
// and properly suppress embedded counterparts without emitting duplicate entries.
func TestLayeredDataProvider_WalkDirDeduplication(t *testing.T) {
tmpDir := t.TempDir()
if err := os.WriteFile(filepath.Join(tmpDir, "registry.yaml"), []byte(testEmptyRegistryContent), 0600); err != nil {
t.Fatalf("failed to write registry.yaml: %v", err)
}
overlaysDir := filepath.Join(tmpDir, "overlays")
if err := os.MkdirAll(overlaysDir, 0755); err != nil {
t.Fatalf("failed to create overlays dir: %v", err)
}
// Override an embedded overlay file (overlays/base.yaml)
if err := os.WriteFile(filepath.Join(overlaysDir, "base.yaml"), []byte("metadata:\n name: base-override\n"), 0600); err != nil {
t.Fatalf("failed to write base.yaml: %v", err)
}

embedded := NewEmbeddedDataProvider(GetEmbeddedFS(), ".")
provider, err := NewLayeredDataProvider(embedded, LayeredProviderConfig{
ExternalDir: tmpDir,
})
if err != nil {
t.Fatalf("failed to create layered provider: %v", err)
}

var seen []string
err = provider.WalkDir(context.Background(), "overlays", func(path string, d os.DirEntry, err error) error {
if err != nil {
return err
}
if !d.IsDir() && filepath.Base(path) == "base.yaml" {
seen = append(seen, path)
}
return nil
})
if err != nil {
t.Fatalf("WalkDir failed: %v", err)
}

if len(seen) != 1 || seen[0] != "overlays/base.yaml" {
t.Errorf("expected exactly [overlays/base.yaml], got %v", seen)
}
}

// TestLayeredDataProvider_CachedRegistry tests that merged registry is cached.
func TestLayeredDataProvider_CachedRegistry(t *testing.T) {
tmpDir := t.TempDir()
Expand Down
Loading