Skip to content

Commit 9a34d96

Browse files
fix(features): harden request resolution
Make feature checks reentrant and single-flight, validate rule declarations across short-circuit paths, cache inventory feature metadata, and codify the HTTP context boundary used by the remote server. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 1e4a1ca6-53f7-4158-af22-35d2448d0b13
1 parent 7bab7db commit 9a34d96

10 files changed

Lines changed: 377 additions & 107 deletions

File tree

docs/feature-flags.md

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,19 @@ The service deduplicates the declared flags, resolves each one at most once for
4040
the request, and shares those values with tool dependencies. Feature checks
4141
inside handlers continue to use `deps.IsFeatureEnabled`.
4242

43+
Feature predicates are pure and may depend only on their resolver. Construction
44+
validates every combination of up to 16 declared flags, so an undeclared lookup
45+
fails immediately even when ordinary evaluation would short-circuit that
46+
branch.
47+
48+
The inventory's checker owns request feature state. Once installed, that state
49+
is authoritative; a checker stored on tool dependencies is used only as a
50+
fallback when handlers are invoked directly without request state. HTTP
51+
availability is resolved after outer HTTP middleware and the inventory factory
52+
run, but before MCP receiving middleware, because the tool set must be known
53+
before constructing the MCP server. Handler-only lazy checks use the live
54+
tool-call context.
55+
4356
---
4457

4558
## Tools affected by each flag

pkg/github/dependencies.go

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -203,9 +203,9 @@ func (d BaseDeps) Metrics(ctx context.Context) metrics.Metrics {
203203
// GetRequestStateSealer implements RequestStateSealerProvider.
204204
func (d BaseDeps) GetRequestStateSealer() RequestStateSealer { return d.StateSealer }
205205

206-
// IsFeatureEnabled checks if a feature flag is enabled.
207-
// Returns false if the feature checker is nil, flag name is empty, or an error occurs.
208-
// This allows tools to conditionally change behavior based on feature flags.
206+
// IsFeatureEnabled checks if a feature flag is enabled. Request feature state
207+
// is authoritative when present; the dependency checker is a fallback for
208+
// direct handler invocation. Empty names and checker errors resolve false.
209209
func (d BaseDeps) IsFeatureEnabled(ctx context.Context, flag inventory.FeatureFlag) bool {
210210
return inventory.ResolveFeature(ctx, d.featureChecker, flag)
211211
}
@@ -483,7 +483,9 @@ func (d *RequestDeps) Metrics(ctx context.Context) metrics.Metrics {
483483
return d.obsv.Metrics(ctx)
484484
}
485485

486-
// IsFeatureEnabled checks if a feature flag is enabled.
486+
// IsFeatureEnabled checks if a feature flag is enabled. Request feature state
487+
// is authoritative when present; the dependency checker is a fallback for
488+
// direct handler invocation.
487489
func (d *RequestDeps) IsFeatureEnabled(ctx context.Context, flag inventory.FeatureFlag) bool {
488490
return inventory.ResolveFeature(ctx, d.featureChecker, flag)
489491
}

pkg/http/handler.go

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,11 @@ import (
2323

2424
const subscriptionsListenMethod = "subscriptions/listen"
2525

26+
// InventoryFactoryFunc builds the inventory for one HTTP request. All context
27+
// values required by its feature checker must be installed before this runs:
28+
// feature availability is resolved immediately afterward, before MCP receiving
29+
// middleware can run. Handler-only lazy checks still see receiving-middleware
30+
// context.
2631
type InventoryFactoryFunc func(r *http.Request) (*inventory.Inventory, error)
2732

2833
// GitHubMCPServerFactoryFunc is a function type for creating a new MCP Server instance.
@@ -214,6 +219,9 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
214219
if methodInfo, ok := ghcontext.MCPMethod(r.Context()); ok && methodInfo != nil {
215220
invToUse = inv.ForMCPRequest(methodInfo.Method, methodInfo.ItemName)
216221
}
222+
// Tool registration must know availability before the MCP server exists.
223+
// Remote consumers install user identity in HTTP middleware before the
224+
// inventory factory, so their per-user checker has its full context here.
217225
r = r.WithContext(invToUse.WithResolvedFeatures(r.Context()))
218226

219227
ghServer, err := h.githubMcpServerFactory(r, h.deps, invToUse, &github.MCPServerConfig{

pkg/http/handler_test.go

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1020,6 +1020,74 @@ func TestCrossOriginProtection(t *testing.T) {
10201020
}
10211021
}
10221022

1023+
func TestFeatureResolutionUsesOuterHTTPContext(t *testing.T) {
1024+
type userContextKey struct{}
1025+
const (
1026+
userValue = "remote-user"
1027+
featureFlag = inventory.FeatureFlag("remote-feature")
1028+
)
1029+
1030+
var checkerCalls int
1031+
tool := mockTool("feature_tool", "test", true)
1032+
tool.FeatureRule = inventory.NewFeatureRule(
1033+
[]inventory.FeatureFlag{featureFlag},
1034+
func(featureAsBool inventory.FeatureResolver) bool {
1035+
return featureAsBool(featureFlag)
1036+
},
1037+
)
1038+
inventoryFactory := func(_ *http.Request) (*inventory.Inventory, error) {
1039+
checker := func(ctx context.Context, flag inventory.FeatureFlag) (bool, error) {
1040+
checkerCalls++
1041+
return flag == featureFlag && ctx.Value(userContextKey{}) == userValue, nil
1042+
}
1043+
return inventory.NewBuilder().
1044+
SetTools([]inventory.ServerTool{tool}).
1045+
WithToolsets([]string{"all"}).
1046+
WithFeatureChecker(checker).
1047+
Build()
1048+
}
1049+
1050+
apiHost, err := utils.NewAPIHost("https://api.github.com")
1051+
require.NoError(t, err)
1052+
handler := NewHTTPMcpHandler(
1053+
context.Background(),
1054+
&ServerConfig{Version: "test"},
1055+
nil,
1056+
translations.NullTranslationHelper,
1057+
slog.Default(),
1058+
apiHost,
1059+
WithInventoryFactory(inventoryFactory),
1060+
WithGitHubMCPServerFactory(func(r *http.Request, _ github.ToolDependencies, inv *inventory.Inventory, _ *github.MCPServerConfig) (*mcp.Server, error) {
1061+
assert.True(t, inventory.ResolveFeature(r.Context(), nil, featureFlag))
1062+
require.Len(t, inv.AvailableTools(r.Context()), 1)
1063+
return mcp.NewServer(&mcp.Implementation{Name: "test", Version: "0.0.1"}, nil), nil
1064+
}),
1065+
WithScopeFetcher(allScopesFetcher{}),
1066+
)
1067+
1068+
router := chi.NewRouter()
1069+
router.Use(func(next http.Handler) http.Handler {
1070+
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
1071+
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), userContextKey{}, userValue)))
1072+
})
1073+
})
1074+
handler.RegisterMiddleware(router)
1075+
handler.RegisterRoutes(router)
1076+
1077+
body := `{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28","io.modelcontextprotocol/clientInfo":{"name":"test","version":"1.0.0"},"io.modelcontextprotocol/clientCapabilities":{}}}}`
1078+
req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(body))
1079+
req.Header.Set(headers.ContentTypeHeader, headers.ContentTypeJSON)
1080+
req.Header.Set(headers.AcceptHeader, strings.Join([]string{headers.ContentTypeJSON, headers.ContentTypeEventStream}, ", "))
1081+
req.Header.Set("Mcp-Protocol-Version", "2026-07-28")
1082+
req.Header.Set("Mcp-Method", "tools/list")
1083+
req.Header.Set(headers.AuthorizationHeader, "ghs_test-token")
1084+
1085+
recorder := httptest.NewRecorder()
1086+
router.ServeHTTP(recorder, req)
1087+
require.Equal(t, http.StatusOK, recorder.Code, "response body: %s", recorder.Body.String())
1088+
assert.Equal(t, 1, checkerCalls)
1089+
}
1090+
10231091
func TestHTTPToolMinimumProtocolVersion(t *testing.T) {
10241092
apiHost, err := utils.NewAPIHost("https://api.github.com")
10251093
require.NoError(t, err)

pkg/inventory/builder.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -255,6 +255,8 @@ func (b *Builder) Build() (*Inventory, error) {
255255
}
256256
}
257257

258+
r.cacheFeatureMetadata()
259+
258260
if b.generateInstructions {
259261
r.instructions = generateInstructions(r)
260262
}

pkg/inventory/features.go

Lines changed: 97 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -8,17 +8,22 @@ import (
88
"sync"
99
)
1010

11+
const maxFeatureRuleFlags = 16
12+
1113
// FeatureFlag identifies a feature consistently across inventory consumers.
1214
type FeatureFlag string
1315

14-
// FeatureFlagChecker resolves one feature flag for the current request.
16+
// FeatureFlagChecker resolves one feature flag for the current request. Every
17+
// context value needed for availability checks must be installed before the
18+
// inventory is resolved. Handler-only checks receive the live tool-call context.
1519
type FeatureFlagChecker func(ctx context.Context, flag FeatureFlag) (bool, error)
1620

1721
// FeatureResolver returns the resolved value of a feature flag.
1822
// Implementations absorb resolution errors and fail closed.
1923
type FeatureResolver func(flag FeatureFlag) bool
2024

21-
// FeaturePredicate determines whether an inventory item is available.
25+
// FeaturePredicate determines whether an inventory item is available. Predicates
26+
// must be pure: their result may depend only on calls to the supplied resolver.
2227
type FeaturePredicate func(featureAsBool FeatureResolver) bool
2328

2429
// FeatureRule declares the feature flags used by an availability predicate.
@@ -44,11 +49,33 @@ func NewFeatureRule(features []FeatureFlag, predicate FeaturePredicate) FeatureR
4449
featureSet[feature] = struct{}{}
4550
declared = append(declared, feature)
4651
}
47-
return FeatureRule{
52+
rule := FeatureRule{
4853
features: declared,
4954
featureSet: featureSet,
5055
predicate: predicate,
5156
}
57+
rule.validate()
58+
return rule
59+
}
60+
61+
func (r FeatureRule) validate() {
62+
if r.predicate == nil {
63+
return
64+
}
65+
if len(r.features) > maxFeatureRuleFlags {
66+
panic(fmt.Sprintf("feature rule declares %d flags; maximum is %d", len(r.features), maxFeatureRuleFlags))
67+
}
68+
69+
for assignment := range 1 << len(r.features) {
70+
r.evaluate(func(feature FeatureFlag) bool {
71+
for i, declared := range r.features {
72+
if feature == declared {
73+
return assignment&(1<<i) != 0
74+
}
75+
}
76+
return false
77+
})
78+
}
5279
}
5380

5481
// Features returns the feature flags referenced by the rule.
@@ -69,7 +96,10 @@ func (r FeatureRule) Enabled(featureAsBool FeatureResolver) bool {
6996
if featureAsBool == nil {
7097
return false
7198
}
99+
return r.evaluate(featureAsBool)
100+
}
72101

102+
func (r FeatureRule) evaluate(featureAsBool FeatureResolver) bool {
73103
var undeclared FeatureFlag
74104
usedUndeclared := false
75105
enabled := r.predicate(func(feature FeatureFlag) bool {
@@ -81,25 +111,35 @@ func (r FeatureRule) Enabled(featureAsBool FeatureResolver) bool {
81111
return featureAsBool(feature)
82112
})
83113
if usedUndeclared {
84-
fmt.Fprintf(os.Stderr, "Feature rule used undeclared feature %q\n", undeclared)
85-
return false
114+
panic(fmt.Sprintf("feature rule used undeclared feature %q", undeclared))
86115
}
87116
return enabled
88117
}
89118

90119
type featureStateContextKey struct{}
120+
type resolvingFeatureContextKey struct{}
91121

92122
type featureState struct {
93123
checker FeatureFlagChecker
94124

95-
mu sync.Mutex
96-
values map[FeatureFlag]bool
125+
mu sync.Mutex
126+
results map[FeatureFlag]*featureResult
127+
}
128+
129+
type featureResult struct {
130+
ready chan struct{}
131+
enabled bool
132+
}
133+
134+
type resolvingFeature struct {
135+
flag FeatureFlag
136+
parent *resolvingFeature
97137
}
98138

99139
func newFeatureState(checker FeatureFlagChecker) *featureState {
100140
return &featureState{
101141
checker: checker,
102-
values: make(map[FeatureFlag]bool),
142+
results: make(map[FeatureFlag]*featureResult),
103143
}
104144
}
105145

@@ -108,24 +148,60 @@ func (s *featureState) enabled(ctx context.Context, feature FeatureFlag) bool {
108148
return false
109149
}
110150

111-
s.mu.Lock()
112-
defer s.mu.Unlock()
151+
for current := resolvingFeatureFromContext(ctx); current != nil; current = current.parent {
152+
if current.flag == feature {
153+
fmt.Fprintf(os.Stderr, "Feature flag resolution cycle detected for %q\n", feature)
154+
return false
155+
}
156+
}
113157

114-
if enabled, ok := s.values[feature]; ok {
115-
return enabled
158+
s.mu.Lock()
159+
result, found := s.results[feature]
160+
if !found {
161+
result = &featureResult{ready: make(chan struct{})}
162+
s.results[feature] = result
163+
}
164+
s.mu.Unlock()
165+
166+
if found {
167+
select {
168+
case <-result.ready:
169+
return result.enabled
170+
case <-ctx.Done():
171+
return false
172+
}
116173
}
117174

118-
enabled, err := s.checker(ctx, feature)
175+
completed := false
176+
defer func() {
177+
if !completed {
178+
close(result.ready)
179+
}
180+
}()
181+
182+
resolutionCtx := context.WithValue(ctx, resolvingFeatureContextKey{}, &resolvingFeature{
183+
flag: feature,
184+
parent: resolvingFeatureFromContext(ctx),
185+
})
186+
enabled, err := s.checker(resolutionCtx, feature)
119187
if err != nil {
120188
fmt.Fprintf(os.Stderr, "Feature flag check error for %q: %v\n", feature, err)
121189
enabled = false
122190
}
123-
s.values[feature] = enabled
191+
result.enabled = enabled
192+
completed = true
193+
close(result.ready)
124194
return enabled
125195
}
126196

197+
func resolvingFeatureFromContext(ctx context.Context) *resolvingFeature {
198+
feature, _ := ctx.Value(resolvingFeatureContextKey{}).(*resolvingFeature)
199+
return feature
200+
}
201+
127202
// WithResolvedFeatures resolves the deduplicated feature names into state owned
128-
// by the returned context. Repeated calls extend and reuse that state.
203+
// by the returned context. Repeated calls extend and reuse that state. When state
204+
// already exists, its checker is authoritative and checker is ignored.
129205
func WithResolvedFeatures(ctx context.Context, checker FeatureFlagChecker, features []FeatureFlag) context.Context {
130206
state, _ := ctx.Value(featureStateContextKey{}).(*featureState)
131207
if state == nil {
@@ -145,18 +221,20 @@ func WithResolvedFeatures(ctx context.Context, checker FeatureFlagChecker, featu
145221
}
146222

147223
// ResolveFeature returns a feature value from request-owned resolution state.
148-
// Features not resolved up front are resolved lazily and cached.
149-
func ResolveFeature(ctx context.Context, checker FeatureFlagChecker, feature FeatureFlag) bool {
224+
// Context state and its checker are authoritative. fallbackChecker is used only
225+
// when the context has no state; that uncached compatibility path lets handlers
226+
// invoked directly outside a server continue to resolve features.
227+
func ResolveFeature(ctx context.Context, fallbackChecker FeatureFlagChecker, feature FeatureFlag) bool {
150228
if feature == "" {
151229
return false
152230
}
153231
if state, _ := ctx.Value(featureStateContextKey{}).(*featureState); state != nil {
154232
return state.enabled(ctx, feature)
155233
}
156-
if checker == nil {
234+
if fallbackChecker == nil {
157235
return false
158236
}
159-
return newFeatureState(checker).enabled(ctx, feature)
237+
return newFeatureState(fallbackChecker).enabled(ctx, feature)
160238
}
161239

162240
func featureResolver(ctx context.Context, checker FeatureFlagChecker) FeatureResolver {
@@ -173,19 +251,3 @@ func featureResolver(ctx context.Context, checker FeatureFlagChecker) FeatureRes
173251
return state.enabled(ctx, feature)
174252
}
175253
}
176-
177-
func collectFeatures(rules ...FeatureRule) []FeatureFlag {
178-
seen := make(map[FeatureFlag]struct{})
179-
for _, rule := range rules {
180-
for _, feature := range rule.features {
181-
seen[feature] = struct{}{}
182-
}
183-
}
184-
185-
features := make([]FeatureFlag, 0, len(seen))
186-
for feature := range seen {
187-
features = append(features, feature)
188-
}
189-
slices.Sort(features)
190-
return features
191-
}

0 commit comments

Comments
 (0)