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: 49 additions & 0 deletions apt/extensions.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -466,16 +466,33 @@ def _distroless_extension(mctx):
arch_set = dependency_set["sets"].setdefault(arch, {})
arch_set[pkg_short_key] = package["Version"]

package_templates = []
for mod in mctx.modules:
for pt in mod.tags.package_template:
if pt.template and pt.template_file:
fail("apt.package_template: exactly one of 'template' or 'template_file' must be specified, not both.")
if not pt.template and not pt.template_file:
fail("apt.package_template: either 'template' or 'template_file' must be specified.")

tmpl = pt.template if pt.template else mctx.read(pt.template_file)
package_templates.append({
"packages": pt.packages,
"template": tmpl,
"additional_variables": dict(pt.additional_variables),
})

# Generate a hub repo for every dependency set
lock_content = glock.as_json()
package_repo_modes = compute_package_repo_modes(glock.packages(), package_repo_roots)
package_templates_json = json.encode(package_templates)
for depset_name in dependency_sets.keys():
depset_mergedusr = dependency_set_mergedusr.get(depset_name, False)
translate_dependency_set(
name = depset_name,
depset_name = depset_name,
lock_content = lock_content,
mergedusr = depset_mergedusr,
package_templates = package_templates_json,
)

# Generate a repo per package which will be aliased by hub repo.
Expand Down Expand Up @@ -659,12 +676,44 @@ lock = tag_class(
},
)

package_template = tag_class(
doc = """Configures a custom BUILD file template for packages matching specific name patterns.

The template is rendered into each architecture subpackage (`//<package>/<arch>/BUILD.bazel`).
Custom templates must define the following targets so the package root's multi-platform aliases and repo targets function correctly:
* `:data` (alias or target pointing to `{data_targets}`)
* `:control` (alias or target pointing to `{control_targets}`)
* `:{target_name}` (filegroup containing `{deps} + [":data"]`)

For reference on the standard structure and available template variables, see the default template at
`//apt/private:package.BUILD.tmpl` (https://github.com/bazel-contrib/rules_distroless/blob/main/apt/private/package.BUILD.tmpl).
""",
attrs = {
"packages": attr.string_list(
doc = "List of package names or glob patterns (e.g. ['nvidia-*', 'libc6', '*']) this template applies to.",
default = ["*"],
),
"template": attr.string(
doc = "Inline template string for the package BUILD file. Must define ':data', ':control', and ':{target_name}' targets (see `//apt/private:package.BUILD.tmpl`).",
),
"template_file": attr.label(
doc = "Template file for the package BUILD file. Must define ':data', ':control', and ':{target_name}' targets (see `//apt/private:package.BUILD.tmpl`).",
allow_single_file = True,
),
"additional_variables": attr.string_dict(
doc = "Additional variables to pass into template formatting.",
default = {},
),
},
)

apt = module_extension(
doc = _doc,
implementation = _distroless_extension,
tag_classes = {
"install": install,
"sources_list": sources_list,
"lock": lock,
"package_template": package_template,
},
)
45 changes: 32 additions & 13 deletions apt/private/translate_dependency_set.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -148,8 +148,17 @@ def package_deps_for_architecture(packages, package, architecture, mergedusr = F
if packages[dep_key]["architecture"] in [architecture, "all"]
]

def resolve_package_template(package_name, package_templates, default_template):
"""Resolves the package template and additional variables for a package name."""
for entry in package_templates:
for pattern in entry.get("packages", []):
if util.glob_match(pattern, package_name):
return (entry["template"], entry.get("additional_variables", {}))
return (default_template, {})

def _translate_dependency_set_impl(rctx):
package_template = rctx.read(rctx.attr.package_template)
default_package_template = rctx.read(rctx.attr.package_template)
package_templates = json.decode(rctx.attr.package_templates) if rctx.attr.package_templates else []
lockf = lockfile.from_json(rctx, rctx.attr.lock_content)

packages = lockf.packages()
Expand Down Expand Up @@ -200,20 +209,29 @@ Please unify the versions manually, or use separate `apt.install` calls (with di
),
).architectures[architecture] = package_key

(tmpl, additional_variables) = resolve_package_template(
package_name,
package_templates,
default_package_template,
)

format_vars = {
"target_name": architecture,
"data_targets": '"@%s//:data"' % repo_name,
"control_targets": '"@%s//:control"' % repo_name,
"src": '"@%s//:data"' % repo_name,
"deps": package_deps_for_architecture(packages, package, architecture, mergedusr = rctx.attr.mergedusr),
"urls": package["urls"],
"name": package["name"],
"arch": package["architecture"],
"sha256": package["sha256"],
"repo_name": repo_name,
}
format_vars.update(additional_variables)

rctx.file(
"%s/%s/BUILD.bazel" % (package["name"], architecture),
package_template.format(
target_name = architecture,
data_targets = '"@%s//:data"' % repo_name,
control_targets = '"@%s//:control"' % repo_name,
src = '"@%s//:data"' % repo_name,
deps = package_deps_for_architecture(packages, package, architecture, mergedusr = rctx.attr.mergedusr),
urls = package["urls"],
name = package["name"],
arch = package["architecture"],
sha256 = package["sha256"],
repo_name = repo_name,
),
tmpl.format(**format_vars),
)

for (_, info) in packages_to_architectures.items():
Expand Down Expand Up @@ -270,5 +288,6 @@ translate_dependency_set = repository_rule(
"lock_content": attr.string(doc = "INTERNAL: DO NOT USE"),
"mergedusr": attr.bool(default = False, doc = "INTERNAL: Whether package layers were normalized with merged-/usr semantics."),
"package_template": attr.label(default = "//apt/private:package.BUILD.tmpl"),
"package_templates": attr.string(default = "[]", doc = "INTERNAL: JSON-encoded list of package template configurations"),
},
)
24 changes: 24 additions & 0 deletions apt/private/util.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,29 @@ def _warning(rctx, message):
"\033[0;33mWARNING:\033[0m {}".format(message),
], quiet = False)

def _glob_match(pattern, text):
"""Matches text against a glob pattern with '*' wildcards."""
if pattern == "*":
return True
if "*" not in pattern:
return pattern == text

parts = pattern.split("*")
if not text.startswith(parts[0]):
return False
text = text[len(parts[0]):]

for i in range(1, len(parts) - 1):
sub = parts[i]
if not sub:
continue
idx = text.find(sub)
if idx == -1:
return False
text = text[idx + len(sub):]

return text.endswith(parts[-1])

util = struct(
sanitize = _sanitize,
package_repo_name = _package_repo_name,
Expand All @@ -101,4 +124,5 @@ util = struct(
is_snapshot_uri = _is_snapshot_uri,
index_fact_key = _index_fact_key,
prune_uncacheable_facts = _prune_uncacheable_facts,
glob_match = _glob_match,
)
3 changes: 3 additions & 0 deletions apt/tests/BUILD.bazel
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ load(":linker_script_test.bzl", "linker_script_tests")
load(":lockfile_test.bzl", "lockfile_tests")
load(":resolution_test.bzl", "resolution_tests")
load(":translate_dependency_set_test.bzl", "translate_dependency_set_tests")
load(":util_test.bzl", "util_tests")
load(":version_test.bzl", "version_tests")

version_tests()
Expand All @@ -28,3 +29,5 @@ lockfile_tests()
translate_dependency_set_tests()

extensions_tests()

util_tests()
64 changes: 63 additions & 1 deletion apt/tests/translate_dependency_set_test.bzl
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
"unit tests for dependency set translation"

load("@bazel_skylib//lib:unittest.bzl", "asserts", "unittest")
load("//apt/private:translate_dependency_set.bzl", "package_deps_for_architecture")
load("//apt/private:translate_dependency_set.bzl", "package_deps_for_architecture", "resolve_package_template")
load("//apt/private:util.bzl", "util")

_TEST_SUITE_PREFIX = "translate_dependency_set/"
Expand Down Expand Up @@ -61,6 +61,68 @@ def _package_repo_name_modes_test(ctx):

package_repo_name_modes_test = unittest.make(_package_repo_name_modes_test)

def _resolve_package_template_test(ctx):
env = unittest.begin(ctx)

default_template = "default: {name}"
custom_nvidia_template = "nvidia: {name}"
custom_dev_template = "dev: {name}"

templates = [
{
"packages": ["nvidia-*"],
"template": custom_nvidia_template,
"additional_variables": {"cuda_version": "12.0"},
},
{
"packages": ["*-dev", "libc6"],
"template": custom_dev_template,
"additional_variables": {"is_dev": "true"},
},
]

# Matching nvidia-* prefix
(tmpl, vars) = resolve_package_template("nvidia-driver", templates, default_template)
asserts.equals(env, custom_nvidia_template, tmpl)
asserts.equals(env, {"cuda_version": "12.0"}, vars)

# Matching *-dev suffix
(tmpl, vars) = resolve_package_template("libssl-dev", templates, default_template)
asserts.equals(env, custom_dev_template, tmpl)
asserts.equals(env, {"is_dev": "true"}, vars)

# Matching exact "libc6"
(tmpl, vars) = resolve_package_template("libc6", templates, default_template)
asserts.equals(env, custom_dev_template, tmpl)
asserts.equals(env, {"is_dev": "true"}, vars)

# Fallback to default template when unmatched
(tmpl, vars) = resolve_package_template("bash", templates, default_template)
asserts.equals(env, default_template, tmpl)
asserts.equals(env, {}, vars)

# First match takes precedence
overlapping_templates = [
{
"packages": ["lib*"],
"template": "lib_template",
"additional_variables": {"tier": "1"},
},
{
"packages": ["libc6"],
"template": "libc6_template",
"additional_variables": {"tier": "2"},
},
]
(tmpl, vars) = resolve_package_template("libc6", overlapping_templates, default_template)
asserts.equals(env, "lib_template", tmpl)
asserts.equals(env, {"tier": "1"}, vars)

return unittest.end(env)

resolve_package_template_test = unittest.make(_resolve_package_template_test)

def translate_dependency_set_tests():
no_mixed_architectures_test(name = _TEST_SUITE_PREFIX + "no_mixed_architectures")
package_repo_name_modes_test(name = _TEST_SUITE_PREFIX + "package_repo_name_modes")
resolve_package_template_test(name = _TEST_SUITE_PREFIX + "resolve_package_template")
55 changes: 55 additions & 0 deletions apt/tests/util_test.bzl
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
"unit tests for apt utility functions"

load("@bazel_skylib//lib:unittest.bzl", "asserts", "unittest")
load("//apt/private:util.bzl", "util")

_TEST_SUITE_PREFIX = "util/"

def _glob_match_test(ctx):
env = unittest.begin(ctx)

# Universal wildcard
asserts.true(env, util.glob_match("*", "anything"))
asserts.true(env, util.glob_match("*", ""))
asserts.true(env, util.glob_match("*", "libc6"))

# Exact matches
asserts.true(env, util.glob_match("libc6", "libc6"))
asserts.false(env, util.glob_match("libc6", "libc6-dev"))
asserts.false(env, util.glob_match("libc6", "libm"))
asserts.true(env, util.glob_match("", ""))
asserts.false(env, util.glob_match("", "foo"))

# Prefix match
asserts.true(env, util.glob_match("nvidia-*", "nvidia-driver"))
asserts.true(env, util.glob_match("nvidia-*", "nvidia-smi"))
asserts.true(env, util.glob_match("nvidia-*", "nvidia-"))
asserts.false(env, util.glob_match("nvidia-*", "libnvidia-driver"))
asserts.false(env, util.glob_match("nvidia-*", "nvid"))

# Suffix match
asserts.true(env, util.glob_match("*-dev", "libc6-dev"))
asserts.true(env, util.glob_match("*-dev", "libssl-dev"))
asserts.true(env, util.glob_match("*-dev", "-dev"))
asserts.false(env, util.glob_match("*-dev", "libc6-dev-doc"))
asserts.false(env, util.glob_match("*-dev", "dev"))

# Middle wildcard
asserts.true(env, util.glob_match("lib*-dev", "libc6-dev"))
asserts.true(env, util.glob_match("lib*-dev", "libssl-dev"))
asserts.true(env, util.glob_match("lib*-dev", "lib-dev"))
asserts.false(env, util.glob_match("lib*-dev", "libc6"))
asserts.false(env, util.glob_match("lib*-dev", "libssl-dbg"))

# Multiple wildcards
asserts.true(env, util.glob_match("*foo*bar*", "1foo2bar3"))
asserts.true(env, util.glob_match("*foo*bar*", "foobar"))
asserts.false(env, util.glob_match("*foo*bar*", "barfoo"))
asserts.false(env, util.glob_match("*foo*bar*", "foo"))

return unittest.end(env)

glob_match_test = unittest.make(_glob_match_test)

def util_tests():
glob_match_test(name = _TEST_SUITE_PREFIX + "glob_match")