Skip to content
Merged
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
3 changes: 0 additions & 3 deletions .github/workflows/e2e.yml
Original file line number Diff line number Diff line change
Expand Up @@ -43,9 +43,6 @@ jobs:
with:
go-version: ${{ env.GO_VERSION }}
cache: false
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
with:
persist-credentials: false
- run: make build
- name: Setup Bats and bats libs
id: setup-bats
Expand Down
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ dist/
.sam-one/
# Stray binaries from `go build ./cmd/<name>/` in the repo root
/sam-node
/sam-one
/sam-box
/sam-router
/sam-console
Expand Down
1 change: 0 additions & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -158,7 +158,6 @@ testnet:

test:
CGO_ENABLED=1 go test -v -race -count 1 $(if $(WHAT),-run $(WHAT)) ./...
cd tests/extproc && CGO_ENABLED=1 go test -v -race -count 1 $(if $(WHAT),-run $(WHAT)) ./...

e2e-test: build docker-build
bats -j 10 --verbose-run $(if $(WHAT),--filter "$(WHAT)") tests/e2e/
Expand Down
4 changes: 2 additions & 2 deletions api/datalog.go
Original file line number Diff line number Diff line change
Expand Up @@ -519,10 +519,10 @@ func init() {
// method() or path() fact derives nothing and a narrowed grant fails closed.
BaselineSources.HTTPRules = []string{
fmt.Sprintf(`%s($t, $k) <- %s($m), %s($t, $k, $set), $set.contains($m)`, FactHTTPMethodOK, FactMethod, FactGrantedMethod),
fmt.Sprintf(`%s($t, $k) <- %s($t, $k)`, FactHTTPMethodOK, FactGrantedMethodAny),
fmt.Sprintf(`%s($t, $k) <- %s($m), !($m == "CONNECT"), %s($t, $k)`, FactHTTPMethodOK, FactMethod, FactGrantedMethodAny),
fmt.Sprintf(`%s($t, $k) <- %s($p), %s($t, $k, $set), $set.contains($p)`, FactHTTPPathOK, FactPath, FactGrantedPathExact),
fmt.Sprintf(`%s($t, $k) <- %s($p), %s($t, $k, $prefix), $p.starts_with($prefix)`, FactHTTPPathOK, FactPath, FactGrantedPathPrefix),
fmt.Sprintf(`%s($t, $k) <- %s($t, $k)`, FactHTTPPathOK, FactGrantedPathAny),
fmt.Sprintf(`%s($t, $k) <- %s($p), $p.starts_with("/"), %s($t, $k)`, FactHTTPPathOK, FactPath, FactGrantedPathAny),
fmt.Sprintf(`%s($t, $n) <- %s($t, $n), %s($t, $n), %s($t, $n), %s($t, $n)`, FactGrantedServiceExact, FactService, FactHTTPGrantedServiceExact, FactHTTPMethodOK, FactHTTPPathOK),
fmt.Sprintf(`%s($t, $s) <- %s($t, $n), %s($t, $s), $n.ends_with($s), %s($t, $s), %s($t, $s)`, FactGrantedServiceSuffix, FactService, FactHTTPGrantedServiceSuffix, FactHTTPMethodOK, FactHTTPPathOK),
fmt.Sprintf(`%s($t, $p) <- %s($t, $n), %s($t, $p), $n.starts_with($p), %s($t, $p), %s($t, $p)`, FactGrantedServicePrefix, FactService, FactHTTPGrantedServicePrefix, FactHTTPMethodOK, FactHTTPPathOK),
Expand Down
41 changes: 41 additions & 0 deletions api/http_grants.go
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,9 @@ func ValidateHTTPGrant(g *HTTPGrant, allowedServices []string) error {
if !httpMethodSyntax.MatchString(m) {
return fmt.Errorf("http entry %q: method %q must be an uppercase HTTP method such as \"GET\"", g.GetService(), m)
}
if m == "CONNECT" {
return fmt.Errorf("http entry %q: method %q is not an HTTP request method; use EGRESS_MODE_TCP for tunnels", g.GetService(), m)
}
}
for _, p := range g.GetPaths() {
if err := validateHTTPGrantPath(p); err != nil {
Expand All @@ -82,6 +85,41 @@ func ValidateHTTPGrant(g *HTTPGrant, allowedServices []string) error {
return nil
}

// ValidateHTTPGrants validates all PolicyRole.http entries for a single role,
// ensuring each entry is valid and no two entries narrow the same service
// (which would otherwise combine their method and path facts across grants).
func ValidateHTTPGrants(grants []*HTTPGrant, allowedServices []string) error {
seen := make(map[string]bool, len(grants))
for _, g := range grants {
if err := ValidateHTTPGrant(g, allowedServices); err != nil {
return err
}
if seen[g.GetService()] {
return fmt.Errorf("duplicate http entry for service %q; combine methods and paths into a single entry", g.GetService())
}
seen[g.GetService()] = true
}
return nil
}

// hasEncodedPathTraversal reports whether p contains URL-encoded dot or slash
// sequences (%2e, %2f, %5c, case-insensitive) or ASCII control characters.
func hasEncodedPathTraversal(p string) bool {
for i := 0; i < len(p); i++ {
c := p[i]
if c < 0x20 || c == 0x7f || c == '\\' {
return true
}
if c == '%' && i+2 < len(p) {
hex := strings.ToLower(p[i+1 : i+3])
if hex == "2e" || hex == "2f" || hex == "5c" || hex == "00" {
return true
}
}
}
return false
}

// validateHTTPGrantPath accepts "/exact" or "/prefix/*". The path is matched
// against path($p) as the backend sees it, so it carries no query and no
// dot segment, and a wildcard is only meaningful at the end.
Expand All @@ -92,6 +130,9 @@ func validateHTTPGrantPath(p string) error {
if strings.ContainsAny(p, "?#") {
return fmt.Errorf("path %q must not carry a query or a fragment", p)
}
if hasEncodedPathTraversal(p) {
return fmt.Errorf("path %q must not contain encoded traversal sequences or control characters", p)
}
trimmed := strings.TrimSuffix(p, "*")
if strings.Contains(trimmed, "*") {
return fmt.Errorf("path %q: \"*\" is only allowed at the end, as in \"/v2/*\"", p)
Expand Down
13 changes: 13 additions & 0 deletions api/http_grants_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -271,6 +271,8 @@ func TestValidateHTTPGrant(t *testing.T) {
{"query in path", &HTTPGrant{Service: "mcp://tools", Paths: []string{"/user?x=1"}}, "query"},
{"wildcard in the middle", &HTTPGrant{Service: "mcp://tools", Paths: []string{"/a/*/b"}}, "only allowed at the end"},
{"dot segment", &HTTPGrant{Service: "mcp://tools", Paths: []string{"/a/../b"}}, "dot segment"},
{"encoded dot segment", &HTTPGrant{Service: "mcp://tools", Paths: []string{"/a/%2e%2e/b"}}, "encoded traversal"},
{"CONNECT method rejected", &HTTPGrant{Service: "egress://api.github.com", Methods: []string{"CONNECT"}}, "not an HTTP request method"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
Expand All @@ -287,3 +289,14 @@ func TestValidateHTTPGrant(t *testing.T) {
})
}
}

func TestValidateHTTPGrantsRejectsDuplicateService(t *testing.T) {
allowed := []string{"egress://api.github.com", "mcp://tools"}
dup := []*HTTPGrant{
{Service: "egress://api.github.com", Methods: []string{"GET"}, Paths: []string{"/repos/*"}},
{Service: "egress://api.github.com", Methods: []string{"POST"}, Paths: []string{"/issues/*"}},
}
if err := ValidateHTTPGrants(dup, allowed); err == nil || !strings.Contains(err.Error(), "duplicate http entry") {
t.Fatalf("ValidateHTTPGrants(dup) = %v, want duplicate error", err)
}
}
7 changes: 7 additions & 0 deletions api/labels.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,10 @@ var labelKeySyntax = regexp.MustCompile(`^[a-zA-Z0-9_.-]{1,63}$`)
// realistic cloud region/zone or on-prem naming convention.
const maxLabelValueLen = 255

// MaxNodeLabels bounds the number of labels a single node may declare at
// enrollment so label facts cannot exhaust the Biscuit authorizer fact budget.
const MaxNodeLabels = 64

// ValidateLabelKey checks that a label key is well-formed: 1-63 characters
// from [a-zA-Z0-9_.-].
func ValidateLabelKey(key string) error {
Expand Down Expand Up @@ -71,6 +75,9 @@ func ValidateLabels(labels map[string]string) error {
if len(labels) == 0 {
return nil
}
if len(labels) > MaxNodeLabels {
return fmt.Errorf("too many labels (%d): maximum is %d", len(labels), MaxNodeLabels)
}
keys := make([]string, 0, len(labels))
for k := range labels {
keys = append(keys, k)
Expand Down
28 changes: 18 additions & 10 deletions api/policy_rules.go
Original file line number Diff line number Diff line change
Expand Up @@ -246,23 +246,31 @@ func BuildPolicyRules(roles []*PolicyRole, bindings []*PolicyBinding) (rules []P
}

// Custom entries keep their source text: it may carry expressions,
// which biscuit-go cannot print back.
// which biscuit-go cannot print back. Every entry on a role is gated
// on role("<roleName>") so it applies only to holders of that role.
for _, dl := range role.CustomDatalog {
trimmed := strings.TrimRight(strings.TrimSpace(dl), ";")
if trimmed == "" {
continue
}
r, err := parser.FromStringRule(trimmed)
if err == nil {
rules = append(rules, PolicyRule{Rule: r, Text: trimmed})
continue
}
f, err2 := parser.FromStringFact(trimmed)
if err2 == nil {
add(f.Predicate)
if _, err := parser.FromStringRule(trimmed); err == nil {
headStr, bodyStr, _ := strings.Cut(trimmed, "<-")
scopedText := strings.TrimSpace(headStr) + " <- " + fromRole.String() + ", " + strings.TrimSpace(bodyStr)
scopedRule, scopedErr := parser.FromStringRule(scopedText)
if scopedErr != nil {
warnings = append(warnings, fmt.Sprintf("Failed to scope custom Datalog rule %q for role %s: %v", dl, roleName, scopedErr))
continue
}
rules = append(rules, PolicyRule{Rule: scopedRule, Text: scopedText})
continue
} else {
f, err2 := parser.FromStringFact(trimmed)
if err2 == nil {
add(f.Predicate, fromRole)
continue
}
warnings = append(warnings, fmt.Sprintf("Failed to parse custom Datalog rule/fact %q for role %s: rule_err=%v, fact_err=%v", dl, roleName, err, err2))
}
warnings = append(warnings, fmt.Sprintf("Failed to parse custom Datalog rule/fact %q for role %s: rule_err=%v, fact_err=%v", dl, roleName, err, err2))
}
}

Expand Down
4 changes: 2 additions & 2 deletions api/policy_rules_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -70,8 +70,8 @@ func TestBuildPolicyRules(t *testing.T) {
"target_restricted(true) <- role(\"test-role\")": false,
"granted_target_set(\"node\", [\"legacy-peer\", \"peer-abc\"]) <- role(\"test-role\")": false,
"granted_target_set(\"custom-fact\", [\"custom-val\"]) <- role(\"test-role\")": false,
"custom_rule($x) <- fact($x), $x > 3": false,
"custom_fact(\"hello\") <- true": false,
"custom_rule($x) <- role(\"test-role\"), fact($x), $x > 3": false,
"custom_fact(\"hello\") <- role(\"test-role\")": false,
}

for _, rule := range rules {
Expand Down
68 changes: 60 additions & 8 deletions api/tar.go
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,24 @@ func HTTPMethodSyntaxPattern() string {
return httpMethodSyntax.String()
}

func isPrintableASCIIString(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] < 0x20 || s[i] > 0x7e {
return false
}
}
return true
}

func hasASCIIControl(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] < 0x20 || s[i] == 0x7f {
return true
}
}
return false
}

// ValidateTARServicePattern validates an entry in TaskRule.allowed_services
// using the same dot-anchored service grammar as PolicyRole.allowed_services
// ("*", "<type>://*", "<type>://*.<suffix>", "<type>://<prefix>.*",
Expand All @@ -92,6 +110,9 @@ func ValidateTARServicePattern(svc string) error {
if err := ValidateServiceFormat(svc); err != nil {
return err
}
if err := ValidateEgressServicePattern(svc); err != nil {
return err
}
if svc == "*" {
return nil
}
Expand All @@ -115,9 +136,15 @@ func ValidateTaskAuthorizationRule(rule *TaskAuthorizationRule) error {
if len(rule.GetName()) > MaxTARNameLength {
return fmt.Errorf("name exceeds max length %d", MaxTARNameLength)
}
if !isPrintableASCIIString(rule.GetName()) {
return fmt.Errorf("name must contain only printable ASCII characters")
}
if len(rule.GetDisplayName()) > MaxTARDescriptionLength {
return fmt.Errorf("display_name exceeds max length %d", MaxTARDescriptionLength)
}
if hasASCIIControl(rule.GetDisplayName()) {
return fmt.Errorf("display_name must not contain control characters")
}
if exp := rule.GetExpireTime(); exp != nil {
if !exp.IsValid() {
return fmt.Errorf("expire_time is invalid")
Expand All @@ -144,6 +171,9 @@ func validateTaskRule(idx int, r *TaskRule) error {
if len(r.GetDescription()) > MaxTARDescriptionLength {
return fmt.Errorf("rule[%d]: description exceeds max length %d", idx, MaxTARDescriptionLength)
}
if hasASCIIControl(r.GetDescription()) {
return fmt.Errorf("rule[%d]: description must not contain control characters", idx)
}
services := r.GetAllowedServices()
if len(services) == 0 {
return fmt.Errorf("rule[%d]: allowed_services must not be empty", idx)
Expand All @@ -160,7 +190,7 @@ func validateTaskRule(idx int, r *TaskRule) error {
return fmt.Errorf("rule[%d]: allowed_resources count %d exceeds max %d", idx, len(r.GetAllowedResources()), MaxEntriesPerTARList)
}
for _, res := range r.GetAllowedResources() {
if res == "" || len(res) > MaxTARResourceLength {
if res == "" || len(res) > MaxTARResourceLength || hasASCIIControl(res) {
return fmt.Errorf("rule[%d]: invalid allowed_resources entry %q", idx, res)
}
}
Expand All @@ -172,7 +202,7 @@ func validateTaskRule(idx int, r *TaskRule) error {
return fmt.Errorf("rule[%d]: allowed_tools count %d exceeds max %d", idx, len(op.GetAllowedTools()), MaxEntriesPerTARList)
}
for _, tool := range op.GetAllowedTools() {
if tool == "" || len(tool) > MaxTARNameLength || strings.ContainsAny(tool, "/?# \t\r\n") {
if tool == "" || len(tool) > MaxTARNameLength || strings.ContainsAny(tool, "/?# \t\r\n") || !isPrintableASCIIString(tool) {
return fmt.Errorf("rule[%d]: invalid tool name %q in allowed_tools", idx, tool)
}
}
Expand All @@ -199,7 +229,7 @@ func validateTaskRule(idx int, r *TaskRule) error {
return fmt.Errorf("rule[%d]: allowed_permissions count %d exceeds max %d", idx, len(op.GetAllowedPermissions()), MaxEntriesPerTARList)
}
for _, perm := range op.GetAllowedPermissions() {
if perm == "" || len(perm) > MaxTARNameLength || strings.ContainsAny(perm, " \t\r\n") {
if perm == "" || len(perm) > MaxTARNameLength || strings.ContainsAny(perm, " \t\r\n") || !isPrintableASCIIString(perm) {
return fmt.Errorf("rule[%d]: invalid permission %q in allowed_permissions", idx, perm)
}
}
Expand Down Expand Up @@ -246,7 +276,7 @@ func DecodeTARBlockPayload(b64 string) (*TaskAuthorizationRule, error) {
if b64 == "" {
return nil, fmt.Errorf("empty tar_block payload")
}
raw, err := base64.RawURLEncoding.DecodeString(b64)
raw, err := base64.RawURLEncoding.Strict().DecodeString(b64)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical

The encoding/base64 package in the Go standard library does not have a Strict() method on base64.Encoding. This will cause a compilation error. You should remove .Strict() and use base64.RawURLEncoding.DecodeString(b64) directly.

Suggested change
raw, err := base64.RawURLEncoding.Strict().DecodeString(b64)
raw, err := base64.RawURLEncoding.DecodeString(b64)

if err != nil {
return nil, fmt.Errorf("invalid base64url in tar_block: %w", err)
}
Expand Down Expand Up @@ -309,11 +339,26 @@ func MatchServicePattern(pattern, reqType, reqName string) bool {
return reqName == pName
}

// IsSafeRequestHTTPPath reports whether reqPath is a normalized, leading-slash
// HTTP path free of dot segments ("." or ".."), URL-encoded dot/slash/backslash
// traversal sequences, query/fragment delimiters, and ASCII control characters.
func IsSafeRequestHTTPPath(reqPath string) bool {
if !strings.HasPrefix(reqPath, "/") || strings.ContainsAny(reqPath, "?#") || hasEncodedPathTraversal(reqPath) {
return false
}
for _, seg := range strings.Split(reqPath, "/") {
if seg == "." || seg == ".." {
return false
}
}
return true
}

// MatchHTTPPath reports whether a path pattern from TaskOperation.allowed_paths
// ("/exact" or "/prefix/*") matches reqPath using the same semantics as
// BuildHTTPGrantFacts and BaselineSources.HTTPRules.
func MatchHTTPPath(pattern, reqPath string) bool {
if pattern == "" || reqPath == "" {
if pattern == "" || !IsSafeRequestHTTPPath(reqPath) {
return false
}
if strings.HasSuffix(pattern, "*") {
Expand Down Expand Up @@ -393,7 +438,7 @@ func MatchTaskRule(rule *TaskRule, req TaskRequestContext) bool {
methods := op.GetAllowedMethods()
paths := op.GetAllowedPaths()
if len(methods) > 0 || len(paths) > 0 {
if !req.HasHTTP || req.Method == "" || req.Method == "CONNECT" || req.Path == "" {
if !req.HasHTTP || req.Method == "" || req.Method == "CONNECT" || !IsSafeRequestHTTPPath(req.Path) {
return false
}
if len(methods) > 0 && !slices.Contains(methods, req.Method) {
Expand Down Expand Up @@ -560,6 +605,8 @@ func BuildTARFromOAuthParams(defaultName, optionsParam string, resources []strin
perms = append(perms, strings.TrimPrefix(tok, "permission:"))
case strings.Contains(tok, "://") || tok == "*":
services = append(services, tok)
default:
return nil, fmt.Errorf("unrecognized scope token %q", tok)
}
}

Expand Down Expand Up @@ -591,11 +638,16 @@ func BuildTARFromOAuthParams(defaultName, optionsParam string, resources []strin
Rules: []*TaskRule{rule},
}
} else {
if len(services) > 0 || len(tools) > 0 || len(methods) > 0 || len(paths) > 0 || len(perms) > 0 {
return nil, fmt.Errorf("cannot combine options parameter with resource or scope parameters")
}
if tar.Name == "" && defaultName != "" {
tar.Name = defaultName
}
if tar.ExpireTime == nil && expireTime != nil {
tar.ExpireTime = expireTime
if expireTime != nil {
if tar.ExpireTime == nil || (expireTime.IsValid() && tar.ExpireTime.IsValid() && expireTime.AsTime().Before(tar.ExpireTime.AsTime())) {
tar.ExpireTime = expireTime
}
}
}

Expand Down
Loading
Loading