From 4bcea053df17229f51a5508b1b6301d681003d9d Mon Sep 17 00:00:00 2001 From: Andrew Nesbitt Date: Wed, 12 Aug 2026 21:46:05 +0100 Subject: [PATCH 1/3] Optimize POM parsing and resolution --- bench_test.go | 173 +++++++++++++++++++ interpolate.go | 131 ++++++++++---- interpolate_test.go | 3 + parse.go | 411 ++++++++++++++++++++++++++++++++++++++++++++ pom.go | 6 +- pom_test.go | 84 ++++++++- resolver.go | 160 +++++++++++++---- resolver_test.go | 91 ++++++++++ 8 files changed, 994 insertions(+), 65 deletions(-) create mode 100644 parse.go diff --git a/bench_test.go b/bench_test.go index 8143dc5..1fc35a5 100644 --- a/bench_test.go +++ b/bench_test.go @@ -4,9 +4,19 @@ import ( "context" "os" "path/filepath" + "strconv" "testing" ) +const ( + benchmarkParentDepth = 16 + benchmarkBOMCount = 6 + benchmarkManagedDepsPerBOM = 40 + benchmarkPropertyCount = 128 + benchmarkManagedDepCount = 256 + benchmarkArtifactCount = 64 +) + // memFetcher preloads every fixture POM into memory so benchmarks measure // resolution work, not disk I/O. type memFetcher map[GAV]*POM @@ -137,6 +147,169 @@ func BenchmarkResolveCorpus(b *testing.B) { } } +func BenchmarkResolveDeepParents(b *testing.B) { + f, root := benchmarkDeepParents() + ctx := context.Background() + b.ResetTimer() + for b.Loop() { + if _, err := NewResolver(f).Resolve(ctx, root, Options{}); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkResolveImportedBOMs(b *testing.B) { + f, root := benchmarkImportedBOMs() + ctx := context.Background() + b.ResetTimer() + for b.Loop() { + if _, err := NewResolver(f).Resolve(ctx, root, Options{}); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkResolveRepeatedProperties(b *testing.B) { + f, root := benchmarkRepeatedProperties() + ctx := context.Background() + b.ResetTimer() + for b.Loop() { + if _, err := NewResolver(f).Resolve(ctx, root, Options{}); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkResolveDependencyManagement(b *testing.B) { + f, root := benchmarkDependencyManagement() + ctx := context.Background() + b.ResetTimer() + for b.Loop() { + if _, err := NewResolver(f).Resolve(ctx, root, Options{}); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkResolveManyArtifacts(b *testing.B) { + f, roots := benchmarkManyArtifacts() + ctx := context.Background() + b.ResetTimer() + for b.Loop() { + r := NewResolver(f) + for _, root := range roots { + if _, err := r.Resolve(ctx, root, Options{}); err != nil { + b.Fatal(err) + } + } + } +} + +func benchmarkDeepParents() (memFetcher, GAV) { + f := memFetcher{} + var parent *Parent + for i := range benchmarkParentDepth { + id := "parent-" + strconv.Itoa(i) + gav := GAV{GroupID: "org.example", ArtifactID: id, Version: "1"} + p := &POM{ + GroupID: gav.GroupID, + ArtifactID: gav.ArtifactID, + Version: gav.Version, + Parent: parent, + Properties: Properties{"shared.version": strconv.Itoa(i + 1)}, + Dependencies: []Dep{{ + GroupID: "org.example.lib", ArtifactID: "lib-" + strconv.Itoa(i), + }}, + DependencyManagement: DepMgmt{Dependencies: []Dep{{ + GroupID: "org.example.lib", ArtifactID: "lib-" + strconv.Itoa(i), Version: "${shared.version}", + }}}, + } + f[gav] = p + parent = &Parent{GroupID: gav.GroupID, ArtifactID: gav.ArtifactID, Version: gav.Version} + } + root := GAV{GroupID: "org.example", ArtifactID: "deep-app", Version: "1"} + f[root] = &POM{GroupID: root.GroupID, ArtifactID: root.ArtifactID, Version: root.Version, Parent: parent} + return f, root +} + +func benchmarkImportedBOMs() (memFetcher, GAV) { + f := memFetcher{} + imports := make([]Dep, 0, benchmarkBOMCount) + deps := make([]Dep, 0, benchmarkBOMCount*benchmarkManagedDepsPerBOM) + for i := range benchmarkBOMCount { + bom := GAV{GroupID: "org.example.bom", ArtifactID: "bom-" + strconv.Itoa(i), Version: "1"} + managed := make([]Dep, 0, benchmarkManagedDepsPerBOM) + for j := range benchmarkManagedDepsPerBOM { + artifact := "lib-" + strconv.Itoa(i) + "-" + strconv.Itoa(j) + managed = append(managed, Dep{GroupID: "org.example.lib", ArtifactID: artifact, Version: "1." + strconv.Itoa(j)}) + deps = append(deps, Dep{GroupID: "org.example.lib", ArtifactID: artifact}) + } + f[bom] = &POM{GroupID: bom.GroupID, ArtifactID: bom.ArtifactID, Version: bom.Version, DependencyManagement: DepMgmt{Dependencies: managed}} + imports = append(imports, Dep{GroupID: bom.GroupID, ArtifactID: bom.ArtifactID, Version: bom.Version, Type: "pom", Scope: scopeImport}) + } + root := GAV{GroupID: "org.example", ArtifactID: "bom-app", Version: "1"} + f[root] = &POM{ + GroupID: root.GroupID, ArtifactID: root.ArtifactID, Version: root.Version, + DependencyManagement: DepMgmt{Dependencies: imports}, Dependencies: deps, + } + return f, root +} + +func benchmarkRepeatedProperties() (memFetcher, GAV) { + props := make(Properties, benchmarkPropertyCount) + props["property.0"] = "1.0.0" + for i := 1; i < benchmarkPropertyCount; i++ { + props["property."+strconv.Itoa(i)] = "${property." + strconv.Itoa(i-1) + "}" + } + deps := make([]Dep, benchmarkPropertyCount) + for i := range deps { + deps[i] = Dep{GroupID: "org.example.lib", ArtifactID: "lib-" + strconv.Itoa(i), Version: "${property.127}"} + } + root := GAV{GroupID: "org.example", ArtifactID: "property-app", Version: "1"} + f := memFetcher{root: { + GroupID: root.GroupID, ArtifactID: root.ArtifactID, Version: root.Version, + Properties: props, Dependencies: deps, + }} + return f, root +} + +func benchmarkDependencyManagement() (memFetcher, GAV) { + managed := make([]Dep, benchmarkManagedDepCount) + deps := make([]Dep, benchmarkManagedDepCount) + for i := range benchmarkManagedDepCount { + artifact := "lib-" + strconv.Itoa(i) + managed[i] = Dep{GroupID: "org.example.lib", ArtifactID: artifact, Version: "${shared.version}", Scope: "runtime"} + deps[i] = Dep{GroupID: "org.example.lib", ArtifactID: artifact} + } + root := GAV{GroupID: "org.example", ArtifactID: "managed-app", Version: "1"} + f := memFetcher{root: { + GroupID: root.GroupID, ArtifactID: root.ArtifactID, Version: root.Version, + Properties: Properties{"shared.version": "2.0.0"}, + DependencyManagement: DepMgmt{Dependencies: managed}, Dependencies: deps, + }} + return f, root +} + +func benchmarkManyArtifacts() (memFetcher, []GAV) { + f, bomRoot := benchmarkImportedBOMs() + parent := GAV{GroupID: "org.example", ArtifactID: "shared-parent", Version: "1"} + f[parent] = &POM{ + GroupID: parent.GroupID, ArtifactID: parent.ArtifactID, Version: parent.Version, + DependencyManagement: f[bomRoot].DependencyManagement, + } + roots := make([]GAV, benchmarkArtifactCount) + for i := range roots { + root := GAV{GroupID: "org.example", ArtifactID: "app-" + strconv.Itoa(i), Version: "1"} + roots[i] = root + f[root] = &POM{ + GroupID: root.GroupID, ArtifactID: root.ArtifactID, Version: root.Version, + Parent: &Parent{GroupID: parent.GroupID, ArtifactID: parent.ArtifactID, Version: parent.Version}, + Dependencies: []Dep{{GroupID: "org.example.lib", ArtifactID: "lib-0-0"}}, + } + } + return f, roots +} + func splitFixtureName(s string) [3]string { var out [3]string first := -1 diff --git a/interpolate.go b/interpolate.go index f093966..00aefff 100644 --- a/interpolate.go +++ b/interpolate.go @@ -1,50 +1,111 @@ package pom import ( - "regexp" "strings" ) const ( - maxInterpolationPasses = 10 - maxInterpolatedLength = 1 << 20 // 1 MiB + maxInterpolationPasses = 10 + maxInterpolatedLength = 1 << 20 // 1 MiB + expressionStart = "${" ) -var exprRE = regexp.MustCompile(`\$\{([^}]+)\}`) - // interpolate substitutes ${name} expressions in s using props. It iterates // until no further substitutions occur or maxInterpolationPasses is reached, // so chained references like ${a} -> ${b} -> value resolve correctly. func interpolate(s string, props map[string]string) string { - if !strings.Contains(s, "${") { + if !strings.Contains(s, expressionStart) { return s } for range maxInterpolationPasses { - changed := false - capped := false - growth := 0 - baseLen := len(s) - s = exprRE.ReplaceAllStringFunc(s, func(m string) string { - if capped { - return m - } - name := m[2 : len(m)-1] - if v, ok := lookup(props, name); ok { - growth += len(v) - len(m) - if baseLen+growth > maxInterpolatedLength { - capped = true - return m + var changed, capped bool + s, changed, capped = interpolatePass(s, props) + if capped || !changed || !strings.Contains(s, expressionStart) { + break + } + } + return s +} + +func interpolatePass(s string, props map[string]string) (string, bool, bool) { + first := strings.Index(s, expressionStart) + if first < 0 { + return s, false, false + } + + // A property reference is commonly the whole value. Returning the map's + // string directly avoids building an identical intermediate string. + if v, ok := wholeExpression(s, props); ok { + if len(v) > maxInterpolatedLength { + return s, false, true + } + return v, v != s, false + } + + baseLen := len(s) + growth := 0 + search := 0 + last := 0 + changed := false + capped := false + var out strings.Builder + for { + open, close, ok := nextExpression(s, search) + if !ok { + break + } + if close == open+len(expressionStart) { + search = close + 1 + continue + } + + replacement, ok := lookup(props, s[open+len(expressionStart):close]) + match := s[open : close+1] + if ok && !capped { + growth += len(replacement) - len(match) + if baseLen+growth > maxInterpolatedLength { + capped = true + } else if replacement != match { + if !changed { + out.Grow(baseLen) } + out.WriteString(s[last:open]) + out.WriteString(replacement) + last = close + 1 changed = true - return v } - return m - }) - if capped || !changed || !strings.Contains(s, "${") { - break } + search = close + 1 } - return s + if !changed { + return s, false, capped + } + out.WriteString(s[last:]) + return out.String(), true, capped +} + +func wholeExpression(s string, props map[string]string) (string, bool) { + if !strings.HasPrefix(s, expressionStart) { + return "", false + } + close := strings.IndexByte(s[len(expressionStart):], '}') + if close != len(s)-len(expressionStart)-1 || close == 0 { + return "", false + } + return lookup(props, s[len(expressionStart):len(s)-1]) +} + +func nextExpression(s string, search int) (int, int, bool) { + relOpen := strings.Index(s[search:], expressionStart) + if relOpen < 0 { + return 0, 0, false + } + open := search + relOpen + relClose := strings.IndexByte(s[open+len(expressionStart):], '}') + if relClose < 0 { + return 0, 0, false + } + return open, open + len(expressionStart) + relClose, true } // lookup resolves a single property name, applying the alias rules Maven @@ -59,7 +120,7 @@ func lookup(props map[string]string, name string) (string, bool) { } } switch name { - case "version", "groupId", "artifactId": + case elementVersion, elementGroupID, elementArtifactID: if v, ok := props["project."+name]; ok { return v, true } @@ -69,14 +130,20 @@ func lookup(props map[string]string, name string) (string, bool) { // containsExpr reports whether s still contains an unresolved ${...}. func containsExpr(s string) bool { - return strings.Contains(s, "${") + return strings.Contains(s, expressionStart) } // firstExpr returns the first ${name} property name in s, or "" if none. func firstExpr(s string) string { - m := exprRE.FindStringSubmatch(s) - if m == nil { - return "" + for search := 0; search < len(s); { + open, close, ok := nextExpression(s, search) + if !ok { + break + } + if close > open+len(expressionStart) { + return s[open+len(expressionStart) : close] + } + search = close + 1 } - return m[1] + return "" } diff --git a/interpolate_test.go b/interpolate_test.go index 1758523..8f4ea20 100644 --- a/interpolate_test.go +++ b/interpolate_test.go @@ -21,8 +21,11 @@ func TestInterpolate(t *testing.T) { {"${b}", "1"}, {"${c}", "1.1"}, {"v${a}-final", "v1-final"}, + {"${a}${b}", "11"}, {"${missing}", "${missing}"}, {"${a}.${missing}", "1.${missing}"}, + {"before ${a", "before ${a"}, + {"${} ${a}", "${} 1"}, {"${pom.version}", "2.0"}, {"${version}", "2.0"}, {"${groupId}", "org.example"}, diff --git a/parse.go b/parse.go new file mode 100644 index 0000000..e01e4ec --- /dev/null +++ b/parse.go @@ -0,0 +1,411 @@ +package pom + +import ( + "bytes" + "encoding/xml" + "fmt" + "strings" +) + +const ( + elementArtifactID = "artifactId" + elementDependencies = "dependencies" + elementGroupID = "groupId" + elementName = "name" + elementProject = "project" + elementURL = "url" + elementVersion = "version" +) + +func decodePOM(dec *xml.Decoder) (*POM, error) { + for { + tok, err := dec.Token() + if err != nil { + return nil, err + } + switch tok := tok.(type) { + case xml.StartElement: + if tok.Name.Local != elementProject { + return nil, fmt.Errorf("expected element type but have <%s>", tok.Name.Local) + } + p := &POM{XMLName: tok.Name} + if err := decodeProject(dec, tok, p); err != nil { + return nil, err + } + return p, nil + case xml.CharData: + if len(bytes.TrimSpace(tok)) != 0 { + return nil, fmt.Errorf("expected element type but have character data") + } + } + } +} + +func decodeProject(dec *xml.Decoder, start xml.StartElement, p *POM) error { + for { + tok, err := dec.Token() + if err != nil { + return err + } + switch tok := tok.(type) { + case xml.StartElement: + switch tok.Name.Local { + case elementGroupID: + p.GroupID, err = decodeText(dec, tok) + case elementArtifactID: + p.ArtifactID, err = decodeText(dec, tok) + case elementVersion: + p.Version, err = decodeText(dec, tok) + case "packaging": + p.Packaging, err = decodeText(dec, tok) + case "parent": + p.Parent, err = decodeParent(dec, tok) + case elementName: + p.Name, err = decodeText(dec, tok) + case "description": + p.Description, err = decodeText(dec, tok) + case elementURL: + p.URL, err = decodeText(dec, tok) + case "licenses": + err = decodeLicenses(dec, tok, &p.Licenses) + case "scm": + err = decodeSCM(dec, tok, &p.SCM) + case "distributionManagement": + err = decodeDistMgmt(dec, tok, &p.DistributionManagement) + case "properties": + p.Properties, err = decodeProperties(dec, tok) + case elementDependencies: + err = decodeDependencies(dec, tok, &p.Dependencies) + case "dependencyManagement": + err = decodeDepMgmt(dec, tok, &p.DependencyManagement) + case "profiles": + err = decodeProfiles(dec, tok, &p.Profiles) + default: + err = dec.Skip() + } + if err != nil { + return err + } + case xml.EndElement: + if tok.Name == start.Name { + return nil + } + } + } +} + +func decodeParent(dec *xml.Decoder, start xml.StartElement) (*Parent, error) { + p := &Parent{} + err := decodeFields(dec, start, func(child xml.StartElement) error { + var err error + switch child.Name.Local { + case elementGroupID: + p.GroupID, err = decodeText(dec, child) + case elementArtifactID: + p.ArtifactID, err = decodeText(dec, child) + case elementVersion: + p.Version, err = decodeText(dec, child) + case "relativePath": + var path string + path, err = decodeText(dec, child) + p.RelativePath = &path + default: + err = dec.Skip() + } + return err + }) + return p, err +} + +func decodeLicenses(dec *xml.Decoder, start xml.StartElement, licenses *[]License) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + if child.Name.Local != "license" { + return dec.Skip() + } + var license License + if err := decodeFields(dec, child, func(field xml.StartElement) error { + var err error + switch field.Name.Local { + case elementName: + license.Name, err = decodeText(dec, field) + case elementURL: + license.URL, err = decodeText(dec, field) + default: + err = dec.Skip() + } + return err + }); err != nil { + return err + } + *licenses = append(*licenses, license) + return nil + }) +} + +func decodeSCM(dec *xml.Decoder, start xml.StartElement, scm *SCM) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + var err error + switch child.Name.Local { + case elementURL: + scm.URL, err = decodeText(dec, child) + case "connection": + scm.Connection, err = decodeText(dec, child) + case "developerConnection": + scm.DeveloperConnection, err = decodeText(dec, child) + default: + err = dec.Skip() + } + return err + }) +} + +func decodeDistMgmt(dec *xml.Decoder, start xml.StartElement, dist *DistMgmt) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + if child.Name.Local != "relocation" { + return dec.Skip() + } + relocation := &Relocation{} + if err := decodeFields(dec, child, func(field xml.StartElement) error { + var err error + switch field.Name.Local { + case elementGroupID: + relocation.GroupID, err = decodeText(dec, field) + case elementArtifactID: + relocation.ArtifactID, err = decodeText(dec, field) + case elementVersion: + relocation.Version, err = decodeText(dec, field) + case "message": + relocation.Message, err = decodeText(dec, field) + default: + err = dec.Skip() + } + return err + }); err != nil { + return err + } + dist.Relocation = relocation + return nil + }) +} + +func decodeProperties(dec *xml.Decoder, start xml.StartElement) (Properties, error) { + properties := Properties{} + err := decodeFields(dec, start, func(child xml.StartElement) error { + value, err := decodeText(dec, child) + if err == nil { + properties[child.Name.Local] = strings.TrimSpace(value) + } + return err + }) + return properties, err +} + +func decodeDependencies(dec *xml.Decoder, start xml.StartElement, deps *[]Dep) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + if child.Name.Local != "dependency" { + return dec.Skip() + } + dep, err := decodeDep(dec, child) + if err == nil { + *deps = append(*deps, dep) + } + return err + }) +} + +func decodeDepMgmt(dec *xml.Decoder, start xml.StartElement, depMgmt *DepMgmt) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + if child.Name.Local != elementDependencies { + return dec.Skip() + } + return decodeDependencies(dec, child, &depMgmt.Dependencies) + }) +} + +func decodeDep(dec *xml.Decoder, start xml.StartElement) (Dep, error) { + var dep Dep + err := decodeFields(dec, start, func(child xml.StartElement) error { + var err error + switch child.Name.Local { + case elementGroupID: + dep.GroupID, err = decodeText(dec, child) + case elementArtifactID: + dep.ArtifactID, err = decodeText(dec, child) + case elementVersion: + dep.Version, err = decodeText(dec, child) + case "type": + dep.Type, err = decodeText(dec, child) + case "classifier": + dep.Classifier, err = decodeText(dec, child) + case "scope": + dep.Scope, err = decodeText(dec, child) + case "optional": + dep.Optional, err = decodeText(dec, child) + case "exclusions": + err = decodeExclusions(dec, child, &dep.Exclusions) + default: + err = dec.Skip() + } + return err + }) + return dep, err +} + +func decodeExclusions(dec *xml.Decoder, start xml.StartElement, exclusions *[]Exclusion) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + if child.Name.Local != "exclusion" { + return dec.Skip() + } + var exclusion Exclusion + if err := decodeFields(dec, child, func(field xml.StartElement) error { + var err error + switch field.Name.Local { + case elementGroupID: + exclusion.GroupID, err = decodeText(dec, field) + case elementArtifactID: + exclusion.ArtifactID, err = decodeText(dec, field) + default: + err = dec.Skip() + } + return err + }); err != nil { + return err + } + *exclusions = append(*exclusions, exclusion) + return nil + }) +} + +func decodeProfiles(dec *xml.Decoder, start xml.StartElement, profiles *[]Profile) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + if child.Name.Local != "profile" { + return dec.Skip() + } + profile, err := decodeProfile(dec, child) + if err == nil { + *profiles = append(*profiles, profile) + } + return err + }) +} + +func decodeProfile(dec *xml.Decoder, start xml.StartElement) (Profile, error) { + var profile Profile + err := decodeFields(dec, start, func(child xml.StartElement) error { + var err error + switch child.Name.Local { + case "id": + profile.ID, err = decodeText(dec, child) + case "activation": + err = decodeActivation(dec, child, &profile.Activation) + case "properties": + profile.Properties, err = decodeProperties(dec, child) + case elementDependencies: + err = decodeDependencies(dec, child, &profile.Dependencies) + case "dependencyManagement": + err = decodeDepMgmt(dec, child, &profile.DependencyManagement) + default: + err = dec.Skip() + } + return err + }) + return profile, err +} + +func decodeActivation(dec *xml.Decoder, start xml.StartElement, activation *Activation) error { + return decodeFields(dec, start, func(child xml.StartElement) error { + var err error + switch child.Name.Local { + case "activeByDefault": + activation.ActiveByDefault, err = decodeText(dec, child) + case "jdk": + activation.JDK, err = decodeText(dec, child) + case "os": + err = decodeFields(dec, child, func(field xml.StartElement) error { + var fieldErr error + switch field.Name.Local { + case elementName: + activation.OS.Name, fieldErr = decodeText(dec, field) + case "family": + activation.OS.Family, fieldErr = decodeText(dec, field) + case "arch": + activation.OS.Arch, fieldErr = decodeText(dec, field) + default: + fieldErr = dec.Skip() + } + return fieldErr + }) + case "property": + err = decodeFields(dec, child, func(field xml.StartElement) error { + var fieldErr error + switch field.Name.Local { + case elementName: + activation.Property.Name, fieldErr = decodeText(dec, field) + case "value": + activation.Property.Value, fieldErr = decodeText(dec, field) + default: + fieldErr = dec.Skip() + } + return fieldErr + }) + default: + err = dec.Skip() + } + return err + }) +} + +func decodeFields(dec *xml.Decoder, start xml.StartElement, decode func(xml.StartElement) error) error { + for { + tok, err := dec.Token() + if err != nil { + return err + } + switch tok := tok.(type) { + case xml.StartElement: + if err := decode(tok); err != nil { + return err + } + case xml.EndElement: + if tok.Name == start.Name { + return nil + } + } + } +} + +func decodeText(dec *xml.Decoder, start xml.StartElement) (string, error) { + var first string + var text strings.Builder + for { + tok, err := dec.Token() + if err != nil { + return "", err + } + switch tok := tok.(type) { + case xml.CharData: + part := string(tok) + if first == "" && text.Len() == 0 { + first = part + continue + } + if text.Len() == 0 { + text.Grow(len(first) + len(tok)) + text.WriteString(first) + first = "" + } + text.WriteString(part) + case xml.StartElement: + if err := dec.Skip(); err != nil { + return "", err + } + case xml.EndElement: + if tok.Name == start.Name { + if text.Len() != 0 { + return text.String(), nil + } + return first, nil + } + } + } +} diff --git a/pom.go b/pom.go index 06c6b77..db5b484 100644 --- a/pom.go +++ b/pom.go @@ -266,11 +266,11 @@ func ParsePOM(data []byte) (*POM, error) { dec.CharsetReader = func(charset string, input io.Reader) (io.Reader, error) { return input, nil } - var p POM - if err := dec.Decode(&p); err != nil { + p, err := decodePOM(dec) + if err != nil { return nil, fmt.Errorf("pom: parse pom: %w", err) } - return &p, nil + return p, nil } // EffectiveGAV returns the POM's own coordinates, falling back to the diff --git a/pom_test.go b/pom_test.go index 7989f91..e3d364c 100644 --- a/pom_test.go +++ b/pom_test.go @@ -1,6 +1,14 @@ package pom -import "testing" +import ( + "bytes" + "encoding/xml" + "io" + "os" + "path/filepath" + "reflect" + "testing" +) func TestParsePOM(t *testing.T) { src := []byte(` @@ -109,6 +117,80 @@ func TestParsePOMError(t *testing.T) { } } +func TestParsePOMMatchesXMLDecoder(t *testing.T) { + files, err := filepath.Glob("testdata/poms/*.pom") + if err != nil { + t.Fatal(err) + } + for _, path := range files { + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + want, err := decodePOMReflect(data) + if err != nil { + t.Fatalf("reference parse %s: %v", path, err) + } + got, err := ParsePOM(data) + if err != nil { + t.Fatalf("ParsePOM %s: %v", path, err) + } + if !reflect.DeepEqual(got, want) { + t.Errorf("ParsePOM %s differs from encoding/xml", path) + } + } +} + +func TestParsePOMMatchesXMLDecoderErrors(t *testing.T) { + inputs := []string{ + "", + "plain text", + "", + "", + "", + "&unknown;", + "", + } + for _, input := range inputs { + _, wantErr := decodePOMReflect([]byte(input)) + _, gotErr := ParsePOM([]byte(input)) + if (gotErr != nil) != (wantErr != nil) { + t.Errorf("ParsePOM(%q) error = %v, reference error = %v", input, gotErr, wantErr) + } + } +} + +func TestParsePOMMatchesXMLDecoderNestedText(t *testing.T) { + data := []byte(` + beforenestedafter + leftnestedright + `) + want, err := decodePOMReflect(data) + if err != nil { + t.Fatal(err) + } + got, err := ParsePOM(data) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(got, want) { + t.Errorf("ParsePOM nested text = %+v, want %+v", got, want) + } +} + +func decodePOMReflect(data []byte) (*POM, error) { + dec := xml.NewDecoder(bytes.NewReader(data)) + dec.Strict = false + dec.CharsetReader = func(_ string, input io.Reader) (io.Reader, error) { + return input, nil + } + var p POM + if err := dec.Decode(&p); err != nil { + return nil, err + } + return &p, nil +} + func TestManagementKey(t *testing.T) { d := Dep{GroupID: "g", ArtifactID: "a"} if d.managementKey() != "g:a:jar:" { diff --git a/resolver.go b/resolver.go index 1b354f2..6a095ef 100644 --- a/resolver.go +++ b/resolver.go @@ -69,14 +69,20 @@ type Fetcher interface { // through one place. type Resolver struct { fetcher Fetcher - cache map[GAV]*EffectivePOM + cache map[GAV]cachedModels +} + +type cachedModels struct { + defaults *EffectivePOM + pessimistic *EffectivePOM + bom *EffectivePOM } // NewResolver constructs a Resolver around f. Resolved POMs are memoised // for the lifetime of the Resolver since released coordinates are // immutable. func NewResolver(f Fetcher) *Resolver { - return &Resolver{fetcher: f, cache: map[GAV]*EffectivePOM{}} + return &Resolver{fetcher: f, cache: map[GAV]cachedModels{}} } // Options tunes a single Resolve call. @@ -149,7 +155,7 @@ func (d ResolvedDep) GAV() GAV { // Resolve fetches gav and computes its effective POM under opts. func (r *Resolver) Resolve(ctx context.Context, gav GAV, opts Options) (*EffectivePOM, error) { - if ep, ok := r.cache[gav]; ok { + if ep := r.cachedModel(gav, opts.Profiles); ep != nil { return ep, nil } root, err := r.fetcher.Fetch(ctx, gav) @@ -160,17 +166,52 @@ func (r *Resolver) Resolve(ctx context.Context, gav GAV, opts Options) (*Effecti if err != nil { return nil, err } - r.cache[gav] = ep + r.cacheModel(gav, opts.Profiles, ep) return ep, nil } +func (r *Resolver) cachedModel(gav GAV, activation ProfileActivation) *EffectivePOM { + models := r.cache[gav] + switch activation.Mode { + case Pessimistic: + return models.pessimistic + case Explicit: + if len(activation.IDs) != 0 { + return nil + } + return models.defaults + case OnlyDefault: + fallthrough + default: + return models.defaults + } +} + +func (r *Resolver) cacheModel(gav GAV, activation ProfileActivation, ep *EffectivePOM) { + models := r.cache[gav] + switch activation.Mode { + case Pessimistic: + models.pessimistic = ep + case Explicit: + if len(activation.IDs) != 0 { + return + } + models.defaults = ep + case OnlyDefault: + fallthrough + default: + models.defaults = ep + } + r.cache[gav] = models +} + // ResolvePOM computes the effective POM for an already-parsed root POM. // Useful when the caller holds a pom.xml from a source checkout that is // not itself fetchable by coordinate. func (r *Resolver) ResolvePOM(ctx context.Context, root *POM, opts Options) (*EffectivePOM, error) { chain, warnings := r.parentChain(ctx, root) - m := newMerger(opts.Profiles) + m := newMerger(opts.Profiles, chain) parentFailed := len(warnings) > 0 for _, p := range chain { m.apply(p) @@ -279,15 +320,16 @@ func (r *Resolver) expandBOMs(ctx context.Context, m *merger, depth int) []strin } func (r *Resolver) resolveBOM(ctx context.Context, gav GAV, depth int) (*EffectivePOM, error) { - if ep, ok := r.cache[gav]; ok { - return ep, nil + models := r.cache[gav] + if models.bom != nil { + return models.bom, nil } root, err := r.fetcher.Fetch(ctx, gav) if err != nil { return nil, err } chain, warnings := r.parentChain(ctx, root) - m := newMerger(ProfileActivation{Mode: OnlyDefault}) + m := newMerger(ProfileActivation{Mode: OnlyDefault}, chain) for _, p := range chain { m.apply(p) } @@ -301,7 +343,9 @@ func (r *Resolver) resolveBOM(ctx context.Context, gav GAV, depth int) (*Effecti DependencyManagement: m.depMgmt, Warnings: warnings, } - r.cache[gav] = ep + models = r.cache[gav] + models.bom = ep + r.cache[gav] = models return ep, nil } @@ -324,21 +368,63 @@ type merger struct { deps []Dep depKeys map[string]int - depProf map[string]string - profDefs map[string]bool + depProf []string + profDefs map[string]struct{} activeProfiles []string } -func newMerger(act ProfileActivation) *merger { - return &merger{ +func newMerger(act ProfileActivation, chain []*POM) *merger { + props, managed, imports, deps, profileDefs, activeProfiles := mergerCapacities(act, chain) + m := &merger{ activation: act, - props: map[string]string{}, - depMgmt: map[string]Dep{}, - depKeys: map[string]int{}, - depProf: map[string]string{}, - profDefs: map[string]bool{}, + props: make(map[string]string, props), + depMgmt: make(map[string]Dep, managed), + depKeys: make(map[string]int, deps), + profDefs: make(map[string]struct{}, profileDefs), + } + if imports != 0 { + m.imports = make([]Dep, 0, imports) + } + if deps != 0 { + m.deps = make([]Dep, 0, deps) + m.depProf = make([]string, 0, deps) } + if activeProfiles != 0 { + m.activeProfiles = make([]string, 0, activeProfiles) + } + return m +} + +func mergerCapacities(act ProfileActivation, chain []*POM) (props, managed, imports, deps, profileDefs, activeProfiles int) { + for _, p := range chain { + props += len(p.Properties) + managed += len(p.DependencyManagement.Dependencies) + for _, d := range p.DependencyManagement.Dependencies { + if d.Scope == scopeImport { + imports++ + } + } + deps += len(p.Dependencies) + for i := range p.Profiles { + profile := &p.Profiles[i] + if !act.active(profile) { + profileDefs += len(profile.Properties) + continue + } + activeProfiles++ + props += len(profile.Properties) + managed += len(profile.DependencyManagement.Dependencies) + for _, d := range profile.DependencyManagement.Dependencies { + if d.Scope == scopeImport { + imports++ + } + } + deps += len(profile.Dependencies) + } + } + managed -= imports + return props, managed, imports, deps, profileDefs, activeProfiles } // apply merges one POM into the accumulator. Called root-first, so later @@ -405,7 +491,7 @@ func (m *merger) interpolateSCM() SCM { func (m *merger) recordProfileGated(pr *Profile) { for k := range pr.Properties { - m.profDefs[k] = true + m.profDefs[k] = struct{}{} } } @@ -434,15 +520,13 @@ func (m *merger) mergeDeps(entries []Dep, profile string) { if i, ok := m.depKeys[k]; ok { m.deps[i] = overlayDep(m.deps[i], d) if profile != "" { - m.depProf[k] = profile + m.depProf[i] = profile } continue } m.depKeys[k] = len(m.deps) m.deps = append(m.deps, d) - if profile != "" { - m.depProf[k] = profile - } + m.depProf = append(m.depProf, profile) } } @@ -516,6 +600,20 @@ func (m *merger) interpolateProps() { } func (m *merger) interpolateDepMgmt() { + for _, d := range m.depMgmt { + if containsExpr(d.GroupID) || containsExpr(d.ArtifactID) || containsExpr(d.Type) || containsExpr(d.Classifier) { + m.interpolateDepMgmtKeys() + return + } + } + for key, d := range m.depMgmt { + d.Version = interpolate(d.Version, m.props) + d.Scope = interpolate(d.Scope, m.props) + m.depMgmt[key] = d + } +} + +func (m *merger) interpolateDepMgmtKeys() { out := make(map[string]Dep, len(m.depMgmt)) for _, d := range m.depMgmt { d.GroupID = interpolate(d.GroupID, m.props) @@ -531,15 +629,14 @@ func (m *merger) interpolateDepMgmt() { func (m *merger) resolveDeps(parentFailed bool) []ResolvedDep { out := make([]ResolvedDep, 0, len(m.deps)) - for _, d := range m.deps { - rd := m.resolveDep(d, parentFailed) + for i, d := range m.deps { + rd := m.resolveDep(d, parentFailed, m.depProf[i]) out = append(out, rd) } return out } -func (m *merger) resolveDep(d Dep, parentFailed bool) ResolvedDep { - rawKey := d.managementKey() +func (m *merger) resolveDep(d Dep, parentFailed bool, profile string) ResolvedDep { d.GroupID = interpolate(d.GroupID, m.props) d.ArtifactID = interpolate(d.ArtifactID, m.props) d.Type = interpolate(d.Type, m.props) @@ -569,7 +666,7 @@ func (m *merger) resolveDep(d Dep, parentFailed bool) ResolvedDep { Scope: defaultScope(d.Scope), Optional: strings.EqualFold(strings.TrimSpace(d.Optional), "true"), Exclusions: d.Exclusions, - Profile: m.depProf[rawKey], + Profile: profile, } rd.Resolution, rd.Expression = m.classify(d, rawVersion, parentFailed) @@ -610,7 +707,7 @@ func (m *merger) classifyExpr(s string, parentFailed bool) (Resolution, string) switch { case strings.HasPrefix(name, "env."): return UnresolvedEnv, s - case m.profDefs[name]: + case hasKey(m.profDefs, name): return UnresolvedProfileGated, s case parentFailed: return UnresolvedParent, s @@ -619,6 +716,11 @@ func (m *merger) classifyExpr(s string, parentFailed bool) (Resolution, string) } } +func hasKey[K comparable, V any](m map[K]V, key K) bool { + _, ok := m[key] + return ok +} + func defaultType(t string) string { if t == "" { return "jar" diff --git a/resolver_test.go b/resolver_test.go index d1bf8a6..353bb4c 100644 --- a/resolver_test.go +++ b/resolver_test.go @@ -347,6 +347,97 @@ func TestResolverCacheAndError(t *testing.T) { } } +func TestResolverCacheSeparatesProfileModes(t *testing.T) { + f := mapFetcher{ + "org.x:app:1": ` + org.xapp1 + extra + org.xextra1 + + `, + } + r := NewResolver(f) + ctx := context.Background() + gav := GAV{"org.x", "app", "1"} + defaults, err := r.Resolve(ctx, gav, Options{}) + if err != nil { + t.Fatal(err) + } + pessimistic, err := r.Resolve(ctx, gav, Options{Profiles: ProfileActivation{Mode: Pessimistic}}) + if err != nil { + t.Fatal(err) + } + if len(defaults.Dependencies) != 0 || len(pessimistic.Dependencies) != 1 { + t.Fatalf("profile modes shared a cache entry: default=%d pessimistic=%d", len(defaults.Dependencies), len(pessimistic.Dependencies)) + } + again, err := r.Resolve(ctx, gav, Options{}) + if err != nil { + t.Fatal(err) + } + if again != defaults { + t.Error("same profile mode should reuse its cached model") + } +} + +func TestResolverDoesNotCacheNamedProfiles(t *testing.T) { + f := mapFetcher{ + "org.x:app:1": ` + org.xapp1 + + oneorg.xone1 + twoorg.xtwo1 + + `, + } + r := NewResolver(f) + ctx := context.Background() + gav := GAV{"org.x", "app", "1"} + one, err := r.Resolve(ctx, gav, Options{Profiles: ProfileActivation{Mode: Explicit, IDs: []string{"one"}}}) + if err != nil { + t.Fatal(err) + } + two, err := r.Resolve(ctx, gav, Options{Profiles: ProfileActivation{Mode: Explicit, IDs: []string{"two"}}}) + if err != nil { + t.Fatal(err) + } + if depByGA(one, "org.x:one") == nil || depByGA(one, "org.x:two") != nil { + t.Errorf("explicit profile one: %+v", one.Dependencies) + } + if depByGA(two, "org.x:two") == nil || depByGA(two, "org.x:one") != nil { + t.Errorf("explicit profile two: %+v", two.Dependencies) + } +} + +func TestResolverCacheSeparatesBOMModel(t *testing.T) { + f := mapFetcher{ + "org.x:bom:1": ` + org.xbom1 + extra + org.xextra1 + + `, + "org.x:app:1": ` + org.xapp1 + + org.xbom1pomimport + + org.xextra + `, + } + r := NewResolver(f) + ctx := context.Background() + if _, err := r.Resolve(ctx, GAV{"org.x", "bom", "1"}, Options{Profiles: ProfileActivation{Mode: Pessimistic}}); err != nil { + t.Fatal(err) + } + app, err := r.Resolve(ctx, GAV{"org.x", "app", "1"}, Options{}) + if err != nil { + t.Fatal(err) + } + if dep := depByGA(app, "org.x:extra"); dep == nil || dep.Resolution != UnresolvedMissing { + t.Errorf("default BOM import reused pessimistic root model: %+v", dep) + } +} + func TestLookupManagedFallback(t *testing.T) { f := mapFetcher{ "org.x:app:1": ` From 5aa7369d97ef157307e5d765d3c4a2a74d555ab4 Mon Sep 17 00:00:00 2001 From: Andrew Nesbitt Date: Wed, 12 Aug 2026 23:00:16 +0100 Subject: [PATCH 2/3] Accept UTF-8 BOM-prefixed POMs --- pom.go | 1 + pom_test.go | 15 +++++++++++++++ 2 files changed, 16 insertions(+) diff --git a/pom.go b/pom.go index db5b484..7c36789 100644 --- a/pom.go +++ b/pom.go @@ -261,6 +261,7 @@ func ParsePOM(data []byte) (*POM, error) { if int64(len(data)) > MaxPOMBytes { return nil, ErrPOMTooLarge } + data = bytes.TrimPrefix(data, []byte{0xEF, 0xBB, 0xBF}) dec := xml.NewDecoder(bytes.NewReader(data)) dec.Strict = false dec.CharsetReader = func(charset string, input io.Reader) (io.Reader, error) { diff --git a/pom_test.go b/pom_test.go index e3d364c..117e40e 100644 --- a/pom_test.go +++ b/pom_test.go @@ -141,6 +141,21 @@ func TestParsePOMMatchesXMLDecoder(t *testing.T) { } } +func TestParsePOMMatchesXMLDecoderWithUTF8BOM(t *testing.T) { + data := []byte("\xEF\xBB\xBForg.example") + want, err := decodePOMReflect(data) + if err != nil { + t.Fatal(err) + } + got, err := ParsePOM(data) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(got, want) { + t.Errorf("ParsePOM with UTF-8 BOM = %+v, want %+v", got, want) + } +} + func TestParsePOMMatchesXMLDecoderErrors(t *testing.T) { inputs := []string{ "", From 4b75c06177f62042a1b3aaa37bcf6ec9a43ea8c1 Mon Sep 17 00:00:00 2001 From: Andrew Nesbitt Date: Fri, 14 Aug 2026 09:30:08 +0100 Subject: [PATCH 3/3] Guard incomplete interpolation expressions --- interpolate.go | 2 +- interpolate_test.go | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/interpolate.go b/interpolate.go index 00aefff..c7b2b60 100644 --- a/interpolate.go +++ b/interpolate.go @@ -89,7 +89,7 @@ func wholeExpression(s string, props map[string]string) (string, bool) { return "", false } close := strings.IndexByte(s[len(expressionStart):], '}') - if close != len(s)-len(expressionStart)-1 || close == 0 { + if close < 0 || close != len(s)-len(expressionStart)-1 || close == 0 { return "", false } return lookup(props, s[len(expressionStart):len(s)-1]) diff --git a/interpolate_test.go b/interpolate_test.go index 8f4ea20..c4e755c 100644 --- a/interpolate_test.go +++ b/interpolate_test.go @@ -25,6 +25,7 @@ func TestInterpolate(t *testing.T) { {"${missing}", "${missing}"}, {"${a}.${missing}", "1.${missing}"}, {"before ${a", "before ${a"}, + {"${", "${"}, {"${} ${a}", "${} 1"}, {"${pom.version}", "2.0"}, {"${version}", "2.0"},