diff --git a/adapter/inbound.go b/adapter/inbound.go index 8717a05ee..d33fbea2c 100644 --- a/adapter/inbound.go +++ b/adapter/inbound.go @@ -86,6 +86,7 @@ type InboundContext struct { DestinationAddresses []netip.Addr DNSResponse *dns.Msg + NamedDNSResponses map[string]*dns.Msg DestinationAddressMatchFromResponse bool SourceGeoIPCode string GeoIPCode string diff --git a/adapter/rule.go b/adapter/rule.go index 2117ba45a..14b406749 100644 --- a/adapter/rule.go +++ b/adapter/rule.go @@ -23,6 +23,10 @@ type DNSRule interface { LegacyPreMatch(metadata *InboundContext) bool WithAddressLimit() bool MatchAddressLimit(metadata *InboundContext, response *dns.Msg) bool + MatchResponseTag() string + MatchResponseTags() []string + MatchResponseAnonymous() bool + Race() bool } type RuleAction interface { diff --git a/dns/router.go b/dns/router.go index 13d4abef0..bd86da93c 100644 --- a/dns/router.go +++ b/dns/router.go @@ -3,6 +3,7 @@ package dns import ( "context" "errors" + "maps" "net/netip" "strings" "sync" @@ -186,10 +187,14 @@ func (r *Router) buildRules(startRules bool) ([]adapter.DNSRule, bool, dnsRuleMo return nil, false, dnsRuleModeFlags{}, err } if !legacyDNSMode { - err = validateLegacyDNSModeDisabledRules(router, r.rawRules, nil) + var validationWarnings []string + validationWarnings, err = validateLegacyDNSModeDisabledRules(router, r.rawRules, nil) if err != nil { return nil, false, dnsRuleModeFlags{}, err } + for _, warning := range validationWarnings { + r.logger.Warn(warning) + } } err = validateEvaluateFakeIPRules(r.rawRules, r.transport) if err != nil { @@ -248,7 +253,8 @@ func (r *Router) ValidateRuleSetMetadataUpdate(tag string, metadata adapter.Rule return err } if !candidateLegacyDNSMode { - return validateLegacyDNSModeDisabledRules(router, r.rawRules, overrides) + _, err = validateLegacyDNSModeDisabledRules(router, r.rawRules, overrides) + return err } return nil } @@ -258,7 +264,7 @@ func (r *Router) ValidateRuleSetMetadataUpdate(tag string, metadata adapter.Rule } if legacyDNSMode { if !candidateLegacyDNSMode && flags.disabled { - err := validateLegacyDNSModeDisabledRules(router, r.rawRules, overrides) + _, err = validateLegacyDNSModeDisabledRules(router, r.rawRules, overrides) if err != nil { return err } @@ -269,7 +275,8 @@ func (r *Router) ValidateRuleSetMetadataUpdate(tag string, metadata adapter.Rule if candidateLegacyDNSMode { return E.New(deprecated.OptionLegacyDNSAddressFilter.MessageWithLink()) } - return validateLegacyDNSModeDisabledRules(router, r.rawRules, overrides) + _, err = validateLegacyDNSModeDisabledRules(router, r.rawRules, overrides) + return err } func (r *Router) matchDNS(ctx context.Context, rules []adapter.DNSRule, allowFakeIP bool, ruleIndex int, isAddressQuery bool, options *adapter.DNSQueryOptions) (adapter.DNSTransport, adapter.DNSRule, int) { @@ -411,16 +418,138 @@ type exchangeWithRulesResult struct { const dnsRespondMissingResponseMessage = "respond action requires an evaluated response from a preceding evaluate action" type dnsRuleWalkState struct { - ruleIndex int - effectiveOptions adapter.DNSQueryOptions - evaluatedResponse *mDNS.Msg - evaluatedTransport adapter.DNSTransport + ruleIndex int + lastLoggedIndex int + effectiveOptions adapter.DNSQueryOptions + anonymousFuture *dnsEvaluatedFuture + namedFutures map[string]*dnsEvaluatedFuture + namedResponses map[string]*mDNS.Msg + namedTransports map[string]adapter.DNSTransport + futures []*dnsEvaluatedFuture + armedRules []*dnsArmedRule + terminalFuture *dnsEvaluatedFuture + terminalIndex int + wake chan struct{} +} + +func (s *dnsRuleWalkState) anonymousResponse() *mDNS.Msg { + if s.anonymousFuture == nil { + return nil + } + return s.anonymousFuture.view() +} + +type dnsEvaluatedFuture struct { + tag string + terminal bool + transport adapter.DNSTransport + cancel context.CancelFunc + done chan struct{} + response *mDNS.Msg + err error + settled bool +} + +func (f *dnsEvaluatedFuture) resolved() bool { + select { + case <-f.done: + return true + default: + return false + } +} + +func (f *dnsEvaluatedFuture) view() *mDNS.Msg { + if !f.resolved() || f.err != nil { + return nil + } + return f.response +} + +type dnsArmedRule struct { + ruleIndex int + rule adapter.DNSRule + futures []*dnsEvaluatedFuture + anonymousFuture *dnsEvaluatedFuture + bindsAnonymous bool + options adapter.DNSQueryOptions } type dnsPendingExchange struct { transport adapter.DNSTransport options adapter.DNSQueryOptions - evaluate bool + future *dnsEvaluatedFuture +} + +type dnsWalkSuspension struct { + await *dnsEvaluatedFuture + drain bool + pending *dnsPendingExchange +} + +func (r *Router) launchDNSEvaluate(ctx context.Context, state *dnsRuleWalkState, tag string, transport adapter.DNSTransport, message *mDNS.Msg, options adapter.DNSQueryOptions) *dnsEvaluatedFuture { + if state.wake == nil { + state.wake = make(chan struct{}, 1) + } + wake := state.wake + exchangeCtx, cancel := context.WithCancel(adapter.OverrideContext(ctx)) + future := &dnsEvaluatedFuture{ + tag: tag, + transport: transport, + cancel: cancel, + done: make(chan struct{}), + } + state.futures = append(state.futures, future) + r.client.ExchangeAsync(exchangeCtx, transport, message, r.finalizeExchangeOptions(options), nil, func(response *mDNS.Msg, err error) { + future.response = response + future.err = err + close(future.done) + select { + case wake <- struct{}{}: + default: + } + }) + return future +} + +func (r *Router) settleDNSFutures(ctx context.Context, message *mDNS.Msg, state *dnsRuleWalkState) { + for _, future := range state.futures { + if future.settled || !future.resolved() { + continue + } + future.settled = true + if future.err != nil && !future.terminal { + r.logger.ErrorContext(ctx, E.Cause(future.err, "exchange failed for ", FormatQuestion(message.Question[0].String()))) + } + if future.tag == "" { + continue + } + newResponses := make(map[string]*mDNS.Msg, len(state.namedResponses)+1) + maps.Copy(newResponses, state.namedResponses) + newResponses[future.tag] = future.view() + state.namedResponses = newResponses + newTransports := make(map[string]adapter.DNSTransport, len(state.namedTransports)+1) + maps.Copy(newTransports, state.namedTransports) + newTransports[future.tag] = future.transport + state.namedTransports = newTransports + } +} + +func cancelDNSFutures(state *dnsRuleWalkState) { + for _, future := range state.futures { + future.cancel() + } +} + +func dnsRefusedResponse(message *mDNS.Msg) *mDNS.Msg { + return &mDNS.Msg{ + MsgHdr: mDNS.MsgHdr{ + Id: message.Id, + Rcode: mDNS.RcodeRefused, + Response: true, + }, + Question: []mDNS.Question{message.Question[0]}, + } } func (r *Router) finalizeExchangeOptions(options adapter.DNSQueryOptions) adapter.DNSQueryOptions { @@ -430,43 +559,132 @@ func (r *Router) finalizeExchangeOptions(options adapter.DNSQueryOptions) adapte return options } -func (r *Router) walkDNSRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool) (exchangeWithRulesResult, *dnsPendingExchange) { +func (r *Router) walkDNSRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool) (exchangeWithRulesResult, *dnsWalkSuspension) { metadata := adapter.ContextFrom(ctx) if metadata == nil { panic("no context") } for ; state.ruleIndex < len(rules); state.ruleIndex++ { currentRule := rules[state.ruleIndex] + hasBindings := len(currentRule.MatchResponseTags()) > 0 || currentRule.MatchResponseAnonymous() + if hasBindings { + r.settleDNSFutures(ctx, message, state) + if currentRule.Race() { + var ( + pendingFutures []*dnsEvaluatedFuture + anonymousFuture *dnsEvaluatedFuture + ) + for _, responseTag := range currentRule.MatchResponseTags() { + future := state.namedFutures[responseTag] + if future != nil && !future.resolved() { + pendingFutures = append(pendingFutures, future) + } + } + bindsAnonymous := currentRule.MatchResponseAnonymous() + if bindsAnonymous { + anonymousFuture = state.anonymousFuture + if anonymousFuture != nil && !anonymousFuture.resolved() { + pendingFutures = append(pendingFutures, anonymousFuture) + } + } + if len(pendingFutures) > 0 { + r.logger.DebugContext(ctx, "armed[", state.ruleIndex, "] ", currentRule, " => ", currentRule.Action()) + state.armedRules = append(state.armedRules, &dnsArmedRule{ + ruleIndex: state.ruleIndex, + rule: currentRule, + futures: pendingFutures, + anonymousFuture: anonymousFuture, + bindsAnonymous: bindsAnonymous, + options: state.effectiveOptions, + }) + continue + } + } else { + var awaitFuture *dnsEvaluatedFuture + for _, responseTag := range currentRule.MatchResponseTags() { + future := state.namedFutures[responseTag] + if future != nil && !future.resolved() { + awaitFuture = future + break + } + } + if awaitFuture == nil && currentRule.MatchResponseAnonymous() { + if future := state.anonymousFuture; future != nil && !future.resolved() { + awaitFuture = future + } + } + if awaitFuture != nil { + return exchangeWithRulesResult{}, &dnsWalkSuspension{await: awaitFuture} + } + } + } metadata.ResetRuleCache() - metadata.DNSResponse = state.evaluatedResponse + metadata.DNSResponse = state.anonymousResponse() + metadata.NamedDNSResponses = state.namedResponses metadata.DestinationAddressMatchFromResponse = false if !currentRule.Match(metadata) { continue } - r.logRuleMatch(ctx, state.ruleIndex, currentRule) + if state.lastLoggedIndex != state.ruleIndex { + state.lastLoggedIndex = state.ruleIndex + r.logRuleMatch(ctx, state.ruleIndex, currentRule) + } switch action := currentRule.Action().(type) { case *R.RuleActionDNSRouteOptions: r.applyDNSRouteOptions(&state.effectiveOptions, *action) case *R.RuleActionEvaluate: - queryOptions := state.effectiveOptions transport, loaded := r.transport.Transport(action.Server) if !loaded { r.logger.ErrorContext(ctx, "transport not found: ", action.Server) - state.evaluatedResponse = nil - state.evaluatedTransport = nil + if action.Tag == "" { + state.anonymousFuture = nil + } continue } + if !action.Speculative && len(state.armedRules) > 0 { + return exchangeWithRulesResult{}, &dnsWalkSuspension{drain: true} + } + queryOptions := state.effectiveOptions r.applyDNSRouteOptions(&queryOptions, action.RuleActionDNSRouteOptions) - return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions, evaluate: true} + future := r.launchDNSEvaluate(ctx, state, action.Tag, transport, message, queryOptions) + if action.Tag == "" { + state.anonymousFuture = future + } else { + if state.namedFutures == nil { + state.namedFutures = make(map[string]*dnsEvaluatedFuture) + } + state.namedFutures[action.Tag] = future + } case *R.RuleActionRespond: - if state.evaluatedResponse == nil { + if len(state.armedRules) > 0 { + return exchangeWithRulesResult{}, &dnsWalkSuspension{drain: true} + } + if responseTag := currentRule.MatchResponseTag(); responseTag != "" { + namedResponse := state.namedResponses[responseTag] + if namedResponse == nil { + return exchangeWithRulesResult{ + err: E.New(dnsRespondMissingResponseMessage), + }, nil + } + return exchangeWithRulesResult{ + response: namedResponse, + transport: state.namedTransports[responseTag], + }, nil + } + if !hasBindings { + if future := state.anonymousFuture; future != nil && !future.resolved() { + return exchangeWithRulesResult{}, &dnsWalkSuspension{await: future} + } + } + response := state.anonymousResponse() + if response == nil { return exchangeWithRulesResult{ err: E.New(dnsRespondMissingResponseMessage), }, nil } return exchangeWithRulesResult{ - response: state.evaluatedResponse, - transport: state.evaluatedTransport, + response: response, + transport: state.anonymousFuture.transport, }, nil case *R.RuleActionDNSRoute: queryOptions := state.effectiveOptions @@ -478,19 +696,27 @@ func (r *Router) walkDNSRules(ctx context.Context, rules []adapter.DNSRule, mess case dnsRouteStatusSkipped: continue } - return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions} + if len(state.armedRules) > 0 { + if action.Speculative && state.terminalFuture == nil { + future := r.launchDNSEvaluate(ctx, state, "", transport, message, queryOptions) + future.terminal = true + state.terminalFuture = future + state.terminalIndex = state.ruleIndex + } + return exchangeWithRulesResult{}, &dnsWalkSuspension{drain: true} + } + if state.terminalFuture != nil && state.terminalIndex == state.ruleIndex { + return exchangeWithRulesResult{}, &dnsWalkSuspension{pending: &dnsPendingExchange{transport: state.terminalFuture.transport, future: state.terminalFuture}} + } + return exchangeWithRulesResult{}, &dnsWalkSuspension{pending: &dnsPendingExchange{transport: transport, options: queryOptions}} case *R.RuleActionReject: + if len(state.armedRules) > 0 { + return exchangeWithRulesResult{}, &dnsWalkSuspension{drain: true} + } switch action.Method { case C.RuleActionRejectMethodDefault: return exchangeWithRulesResult{ - response: &mDNS.Msg{ - MsgHdr: mDNS.MsgHdr{ - Id: message.Id, - Rcode: mDNS.RcodeRefused, - Response: true, - }, - Question: []mDNS.Question{message.Question[0]}, - }, + response: dnsRefusedResponse(message), rejectAction: action, }, nil case C.RuleActionRejectMethodDrop: @@ -500,70 +726,201 @@ func (r *Router) walkDNSRules(ctx context.Context, rules []adapter.DNSRule, mess }, nil } case *R.RuleActionPredefined: + if len(state.armedRules) > 0 { + return exchangeWithRulesResult{}, &dnsWalkSuspension{drain: true} + } return exchangeWithRulesResult{ response: action.Response(message), }, nil } } - return exchangeWithRulesResult{}, &dnsPendingExchange{transport: r.transport.Default(), options: state.effectiveOptions} + if len(state.armedRules) > 0 { + return exchangeWithRulesResult{}, &dnsWalkSuspension{drain: true} + } + return exchangeWithRulesResult{}, &dnsWalkSuspension{pending: &dnsPendingExchange{transport: r.transport.Default(), options: state.effectiveOptions}} } func (r *Router) exchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool) exchangeWithRulesResult { - state := dnsRuleWalkState{effectiveOptions: options} - result, pending := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP) - if pending == nil { + state := dnsRuleWalkState{effectiveOptions: options, lastLoggedIndex: -1} + result, suspension := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP) + if suspension == nil { + cancelDNSFutures(&state) return result } - return r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending) + return r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, suspension) } -func (r *Router) resumeExchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool, pending *dnsPendingExchange) exchangeWithRulesResult { +func (r *Router) resumeExchangeWithRules(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool, suspension *dnsWalkSuspension) exchangeWithRulesResult { + defer cancelDNSFutures(state) for { - response, err := r.client.Exchange(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil) - if !pending.evaluate { - return exchangeWithRulesResult{ - response: response, - transport: pending.transport, - err: err, + r.settleDNSFutures(ctx, message, state) + sweepResult, sweepPending, committed := r.sweepArmedDNSRules(ctx, message, state, allowFakeIP) + if committed { + if sweepPending != nil { + return r.finishPendingExchange(ctx, message, state, sweepPending) + } + return sweepResult + } + if suspension != nil { + if suspension.pending != nil { + return r.finishPendingExchange(ctx, message, state, suspension.pending) + } + if (suspension.await != nil && !suspension.await.resolved()) || (suspension.drain && len(state.armedRules) > 0) { + select { + case <-state.wake: + case <-ctx.Done(): + return exchangeWithRulesResult{err: ctx.Err()} + } + continue } } - if err != nil { - r.logger.ErrorContext(ctx, E.Cause(err, "exchange failed for ", FormatQuestion(message.Question[0].String()))) - state.evaluatedResponse = nil - state.evaluatedTransport = nil - } else { - state.evaluatedResponse = response - state.evaluatedTransport = pending.transport - } - state.ruleIndex++ var result exchangeWithRulesResult - result, pending = r.walkDNSRules(ctx, rules, message, state, allowFakeIP) - if pending == nil { + result, suspension = r.walkDNSRules(ctx, rules, message, state, allowFakeIP) + if suspension == nil { return result } } } +func (r *Router) sweepArmedDNSRules(ctx context.Context, message *mDNS.Msg, state *dnsRuleWalkState, allowFakeIP bool) (exchangeWithRulesResult, *dnsPendingExchange, bool) { + metadata := adapter.ContextFrom(ctx) + for index := 0; index < len(state.armedRules); { + armed := state.armedRules[index] + ready := true + for _, future := range armed.futures { + if !future.resolved() { + ready = false + break + } + } + if !ready { + index++ + continue + } + state.armedRules = append(state.armedRules[:index], state.armedRules[index+1:]...) + metadata.ResetRuleCache() + if armed.bindsAnonymous { + if armed.anonymousFuture != nil { + metadata.DNSResponse = armed.anonymousFuture.view() + } else { + metadata.DNSResponse = nil + } + } else { + metadata.DNSResponse = state.anonymousResponse() + } + metadata.NamedDNSResponses = state.namedResponses + metadata.DestinationAddressMatchFromResponse = false + if !armed.rule.Match(metadata) { + continue + } + r.logRuleMatch(ctx, armed.ruleIndex, armed.rule) + switch action := armed.rule.Action().(type) { + case *R.RuleActionRespond: + var ( + response *mDNS.Msg + transport adapter.DNSTransport + ) + if responseTag := armed.rule.MatchResponseTag(); responseTag != "" { + response = state.namedResponses[responseTag] + transport = state.namedTransports[responseTag] + } else if armed.anonymousFuture != nil { + response = armed.anonymousFuture.view() + transport = armed.anonymousFuture.transport + } else if state.anonymousFuture != nil { + response = state.anonymousResponse() + transport = state.anonymousFuture.transport + } + if response == nil { + return exchangeWithRulesResult{ + err: E.New(dnsRespondMissingResponseMessage), + }, nil, true + } + return exchangeWithRulesResult{ + response: response, + transport: transport, + }, nil, true + case *R.RuleActionDNSRoute: + queryOptions := armed.options + transport, status := r.resolveDNSRoute(action.Server, action.RuleActionDNSRouteOptions, allowFakeIP, &queryOptions) + switch status { + case dnsRouteStatusMissing: + r.logger.ErrorContext(ctx, "transport not found: ", action.Server) + continue + case dnsRouteStatusSkipped: + continue + } + return exchangeWithRulesResult{}, &dnsPendingExchange{transport: transport, options: queryOptions}, true + case *R.RuleActionReject: + switch action.Method { + case C.RuleActionRejectMethodDefault: + return exchangeWithRulesResult{ + response: dnsRefusedResponse(message), + rejectAction: action, + }, nil, true + case C.RuleActionRejectMethodDrop: + return exchangeWithRulesResult{ + rejectAction: action, + err: R.ErrDrop, + }, nil, true + } + case *R.RuleActionPredefined: + return exchangeWithRulesResult{ + response: action.Response(message), + }, nil, true + } + } + return exchangeWithRulesResult{}, nil, false +} + +func (r *Router) finishPendingExchange(ctx context.Context, message *mDNS.Msg, state *dnsRuleWalkState, pending *dnsPendingExchange) exchangeWithRulesResult { + for _, future := range state.futures { + if future != pending.future { + future.cancel() + } + } + if pending.future != nil { + select { + case <-pending.future.done: + case <-ctx.Done(): + return exchangeWithRulesResult{err: ctx.Err()} + } + return exchangeWithRulesResult{ + response: pending.future.view(), + transport: pending.future.transport, + err: pending.future.err, + } + } + response, err := r.client.Exchange(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil) + return exchangeWithRulesResult{ + response: response, + transport: pending.transport, + err: err, + } +} + func (r *Router) exchangeWithRulesAsync(ctx context.Context, rules []adapter.DNSRule, message *mDNS.Msg, options adapter.DNSQueryOptions, allowFakeIP bool, callback func(result exchangeWithRulesResult)) { - state := dnsRuleWalkState{effectiveOptions: options} - result, pending := r.walkDNSRules(ctx, rules, message, &state, allowFakeIP) - if pending == nil { + state := &dnsRuleWalkState{effectiveOptions: options, lastLoggedIndex: -1} + result, suspension := r.walkDNSRules(ctx, rules, message, state, allowFakeIP) + if suspension == nil { + cancelDNSFutures(state) callback(result) return } - if pending.evaluate { - go func() { - callback(r.resumeExchangeWithRules(ctx, rules, message, &state, allowFakeIP, pending)) - }() + if suspension.pending != nil && suspension.pending.future == nil { + cancelDNSFutures(state) + pending := suspension.pending + r.client.ExchangeAsync(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil, func(response *mDNS.Msg, err error) { + callback(exchangeWithRulesResult{ + response: response, + transport: pending.transport, + err: err, + }) + }) return } - r.client.ExchangeAsync(adapter.OverrideContext(ctx), pending.transport, message, r.finalizeExchangeOptions(pending.options), nil, func(response *mDNS.Msg, err error) { - callback(exchangeWithRulesResult{ - response: response, - transport: pending.transport, - err: err, - }) - }) + go func() { + callback(r.resumeExchangeWithRules(ctx, rules, message, state, allowFakeIP, suspension)) + }() } func (r *Router) resolveLookupStrategy(options adapter.DNSQueryOptions) C.DomainStrategy { @@ -694,6 +1051,7 @@ func (r *Router) prepareExchange(ctx context.Context, message *mDNS.Msg) (*dnsEx metadata.Destination = M.Socksaddr{} metadata.QueryType = message.Question[0].Qtype metadata.DNSResponse = nil + metadata.NamedDNSResponses = nil metadata.DestinationAddressMatchFromResponse = false switch metadata.QueryType { case mDNS.TypeA: @@ -877,6 +1235,7 @@ func (r *Router) Lookup(ctx context.Context, domain string, options adapter.DNSQ metadata.Destination = M.Socksaddr{} metadata.Domain = FqdnToDomain(domain) metadata.DNSResponse = nil + metadata.NamedDNSResponses = nil metadata.DestinationAddressMatchFromResponse = false if options.Transport != nil { transport := options.Transport @@ -987,7 +1346,7 @@ func defaultRuleNeedsLegacyDNSModeFromAddressFilter(rule option.DefaultDNSRule) if rule.RuleSetIPCIDRAcceptEmpty { //nolint:staticcheck return true } - return !rule.MatchResponse && (rule.IPAcceptAny || len(rule.IPCIDR) > 0 || rule.IPIsPrivate) + return !rule.MatchResponse.IsEnabled() && (rule.IPAcceptAny || len(rule.IPCIDR) > 0 || rule.IPIsPrivate) } func hasResponseMatchFields(rule option.DefaultDNSRule) bool { @@ -998,7 +1357,7 @@ func hasResponseMatchFields(rule option.DefaultDNSRule) bool { } func defaultRuleDisablesLegacyDNSMode(rule option.DefaultDNSRule) bool { - return rule.MatchResponse || + return rule.MatchResponse.IsEnabled() || hasResponseMatchFields(rule) || rule.Action == C.RuleActionTypeEvaluate || rule.Action == C.RuleActionTypeRespond || @@ -1108,21 +1467,72 @@ func lookupDNSRuleSetMetadata(router adapter.Router, tag string, metadataOverrid return ruleSet.Metadata(), nil } -func validateLegacyDNSModeDisabledRules(router adapter.Router, rules []option.DNSRule, metadataOverrides map[string]adapter.RuleSetMetadata) error { - var seenEvaluate bool +type dnsRuleResponseUse struct { + needsAnonymous bool + referencedTags []string +} + +func validateLegacyDNSModeDisabledRules(router adapter.Router, rules []option.DNSRule, metadataOverrides map[string]adapter.RuleSetMetadata) ([]string, error) { + var ( + warnings []string + seenAnonymousEvaluate bool + seenRace bool + definedTags = make(map[string]bool) + definedTagOrder []string + referencedTags = make(map[string]bool) + lastAnonymousEvaluate = -1 + anonymousReadSinceLast bool + ) for i, rule := range rules { - requiresPriorEvaluate, err := validateLegacyDNSModeDisabledRuleTree(router, rule, metadataOverrides) + use, err := validateLegacyDNSModeDisabledRuleTree(router, rule, metadataOverrides) if err != nil { - return E.Cause(err, "validate dns rule[", i, "]") + return nil, E.Cause(err, "validate dns rule[", i, "]") } - if requiresPriorEvaluate && !seenEvaluate { - return E.New("dns rule[", i, "]: response-based matching requires a preceding evaluate action") + if dnsRuleActionSpeculative(rule) && !seenRace { + warnings = append(warnings, F.ToString("dns rule[", i, "]: `speculative` has no effect without a preceding `race` rule")) + } + if dnsRuleRace(rule) { + seenRace = true + } + if use.needsAnonymous { + if !seenAnonymousEvaluate { + if len(definedTagOrder) > 0 { + return nil, E.New("dns rule[", i, "]: response-based matching requires a preceding evaluate action without `tag`; use `match_response` with an evaluate tag to reference a tagged result") + } + return nil, E.New("dns rule[", i, "]: response-based matching requires a preceding evaluate action") + } + anonymousReadSinceLast = true + } + for _, tag := range use.referencedTags { + if !definedTags[tag] { + return nil, E.New("dns rule[", i, "]: undefined evaluate tag: ", tag) + } + referencedTags[tag] = true } if dnsRuleActionType(rule) == C.RuleActionTypeEvaluate { - seenEvaluate = true + tag := dnsRuleActionEvaluateTag(rule) + if tag == "" { + if lastAnonymousEvaluate >= 0 && !anonymousReadSinceLast { + warnings = append(warnings, F.ToString("dns rule[", lastAnonymousEvaluate, "]: evaluated response is overwritten by dns rule[", i, "] before any use")) + } + seenAnonymousEvaluate = true + lastAnonymousEvaluate = i + anonymousReadSinceLast = false + } else { + if definedTags[tag] { + return nil, E.New("dns rule[", i, "]: duplicate evaluate tag: ", tag) + } + definedTags[tag] = true + definedTagOrder = append(definedTagOrder, tag) + } } } - return nil + for _, tag := range definedTagOrder { + if !referencedTags[tag] { + warnings = append(warnings, F.ToString("evaluate tag is never referenced: ", tag)) + } + } + return warnings, nil } func validateEvaluateFakeIPRules(rules []option.DNSRule, transportManager adapter.DNSTransportManager) error { @@ -1146,54 +1556,75 @@ func validateEvaluateFakeIPRules(rules []option.DNSRule, transportManager adapte return nil } -func validateLegacyDNSModeDisabledRuleTree(router adapter.Router, rule option.DNSRule, metadataOverrides map[string]adapter.RuleSetMetadata) (bool, error) { +func validateLegacyDNSModeDisabledRuleTree(router adapter.Router, rule option.DNSRule, metadataOverrides map[string]adapter.RuleSetMetadata) (dnsRuleResponseUse, error) { switch rule.Type { case "", C.RuleTypeDefault: return validateLegacyDNSModeDisabledDefaultRule(router, rule.DefaultOptions, metadataOverrides) case C.RuleTypeLogical: - requiresPriorEvaluate := dnsRuleActionType(rule) == C.RuleActionTypeRespond + var use dnsRuleResponseUse for i, subRule := range rule.LogicalOptions.Rules { - subRequiresPriorEvaluate, err := validateLegacyDNSModeDisabledRuleTree(router, subRule, metadataOverrides) + subUse, err := validateLegacyDNSModeDisabledRuleTree(router, subRule, metadataOverrides) if err != nil { - return false, E.Cause(err, "sub rule[", i, "]") + return dnsRuleResponseUse{}, E.Cause(err, "sub rule[", i, "]") } - requiresPriorEvaluate = requiresPriorEvaluate || subRequiresPriorEvaluate + use.needsAnonymous = use.needsAnonymous || subUse.needsAnonymous + use.referencedTags = append(use.referencedTags, subUse.referencedTags...) } - return requiresPriorEvaluate, nil + if rule.LogicalOptions.Action == C.RuleActionTypeRespond { + if len(use.referencedTags) > 0 { + return dnsRuleResponseUse{}, E.New("respond on a logical rule cannot bind a `match_response` tag from its sub rules; use a non-logical rule") + } + use.needsAnonymous = true + } + return use, nil default: - return false, nil + return dnsRuleResponseUse{}, nil } } -func validateLegacyDNSModeDisabledDefaultRule(router adapter.Router, rule option.DefaultDNSRule, metadataOverrides map[string]adapter.RuleSetMetadata) (bool, error) { +func validateLegacyDNSModeDisabledDefaultRule(router adapter.Router, rule option.DefaultDNSRule, metadataOverrides map[string]adapter.RuleSetMetadata) (dnsRuleResponseUse, error) { hasResponseRecords := hasResponseMatchFields(rule) - if (hasResponseRecords || len(rule.IPCIDR) > 0 || rule.IPIsPrivate || rule.IPAcceptAny) && !rule.MatchResponse { - return false, E.New("Response Match Fields (ip_cidr, ip_is_private, ip_accept_any, response_rcode, response_answer, response_ns, response_extra) require match_response to be enabled") + if (hasResponseRecords || len(rule.IPCIDR) > 0 || rule.IPIsPrivate || rule.IPAcceptAny) && !rule.MatchResponse.IsEnabled() { + return dnsRuleResponseUse{}, E.New("Response Match Fields (ip_cidr, ip_is_private, ip_accept_any, response_rcode, response_answer, response_ns, response_extra) require match_response to be enabled") } // rule_set entries are only rejected when every referenced set is pure-IP; // mixed sets still fall through because their non-IP branches remain matchable // before a DNS response is available. - if !rule.MatchResponse && len(rule.RuleSet) > 0 { + if !rule.MatchResponse.IsEnabled() && len(rule.RuleSet) > 0 { for _, tag := range rule.RuleSet { metadata, err := lookupDNSRuleSetMetadata(router, tag, metadataOverrides) if err != nil { - return false, err + return dnsRuleResponseUse{}, err } if metadata.ContainsIPCIDRRule && !metadata.ContainsNonIPCIDRRule { - return false, E.New(deprecated.OptionLegacyDNSAddressFilter.MessageWithLink()) + return dnsRuleResponseUse{}, E.New(deprecated.OptionLegacyDNSAddressFilter.MessageWithLink()) } } } if rule.RuleSetIPCIDRAcceptEmpty { //nolint:staticcheck - return false, E.New(deprecated.OptionRuleSetIPCIDRAcceptEmpty.MessageWithLink()) + return dnsRuleResponseUse{}, E.New(deprecated.OptionRuleSetIPCIDRAcceptEmpty.MessageWithLink()) } - return rule.MatchResponse || rule.Action == C.RuleActionTypeRespond, nil + var use dnsRuleResponseUse + if rule.MatchResponse.IsEnabled() { + if responseTag := rule.MatchResponse.ResponseTag(); responseTag != "" { + use.referencedTags = append(use.referencedTags, responseTag) + } else { + use.needsAnonymous = true + } + } + if rule.Action == C.RuleActionTypeRespond && rule.MatchResponse.ResponseTag() == "" { + use.needsAnonymous = true + } + return use, nil } func dnsRuleActionDisablesLegacyDNSMode(action option.DNSRuleAction) bool { + if action.Race { + return true + } switch action.Action { case "", C.RuleActionTypeRoute, C.RuleActionTypeEvaluate: - return action.RouteOptions.DisableOptimisticCache + return action.RouteOptions.DisableOptimisticCache || action.RouteOptions.Speculative case C.RuleActionTypeRouteOptions: return action.RouteOptionsOptions.DisableOptimisticCache default: @@ -1239,3 +1670,36 @@ func dnsRuleActionServer(rule option.DNSRule) string { return "" } } + +func dnsRuleActionEvaluateTag(rule option.DNSRule) string { + switch rule.Type { + case "", C.RuleTypeDefault: + return rule.DefaultOptions.RouteOptions.Tag + case C.RuleTypeLogical: + return rule.LogicalOptions.RouteOptions.Tag + default: + return "" + } +} + +func dnsRuleActionSpeculative(rule option.DNSRule) bool { + switch rule.Type { + case "", C.RuleTypeDefault: + return rule.DefaultOptions.RouteOptions.Speculative + case C.RuleTypeLogical: + return rule.LogicalOptions.RouteOptions.Speculative + default: + return false + } +} + +func dnsRuleRace(rule option.DNSRule) bool { + switch rule.Type { + case "", C.RuleTypeDefault: + return rule.DefaultOptions.Race + case C.RuleTypeLogical: + return rule.LogicalOptions.Race + default: + return false + } +} diff --git a/dns/router_race_test.go b/dns/router_race_test.go new file mode 100644 index 000000000..22e72ff7b --- /dev/null +++ b/dns/router_race_test.go @@ -0,0 +1,527 @@ +package dns + +import ( + "context" + "net/netip" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/sagernet/sing-box/adapter" + C "github.com/sagernet/sing-box/constant" + "github.com/sagernet/sing-box/log" + "github.com/sagernet/sing-box/option" + R "github.com/sagernet/sing-box/route/rule" + + mDNS "github.com/miekg/dns" + "github.com/stretchr/testify/require" +) + +type fakeDNSTransport struct { + tag string + delay time.Duration + rcode int + address netip.Addr + exchangeErr error + + access sync.Mutex + queryCount atomic.Int32 + firstQueried time.Time +} + +func (t *fakeDNSTransport) Start(stage adapter.StartStage) error { + return nil +} + +func (t *fakeDNSTransport) Close() error { + return nil +} + +func (t *fakeDNSTransport) Type() string { + return "fake" +} + +func (t *fakeDNSTransport) Tag() string { + return t.tag +} + +func (t *fakeDNSTransport) Dependencies() []string { + return nil +} + +func (t *fakeDNSTransport) Reset() { +} + +func (t *fakeDNSTransport) Exchange(ctx context.Context, message *mDNS.Msg) (*mDNS.Msg, error) { + t.access.Lock() + if t.firstQueried.IsZero() { + t.firstQueried = time.Now() + } + t.access.Unlock() + t.queryCount.Add(1) + select { + case <-time.After(t.delay): + case <-ctx.Done(): + return nil, ctx.Err() + } + if t.exchangeErr != nil { + return nil, t.exchangeErr + } + if t.rcode != mDNS.RcodeSuccess { + return FixedResponseStatus(message, t.rcode), nil + } + return FixedResponse(message.Id, message.Question[0], []netip.Addr{t.address}, 300), nil +} + +func (t *fakeDNSTransport) ExchangeAsync(ctx context.Context, message *mDNS.Msg, callback func(response *mDNS.Msg, err error)) { + go func() { + callback(t.Exchange(ctx, message)) + }() +} + +type fakeDNSTransportManager struct { + transports map[string]adapter.DNSTransport + defaultTransport adapter.DNSTransport +} + +func (m *fakeDNSTransportManager) Start(stage adapter.StartStage) error { + return nil +} + +func (m *fakeDNSTransportManager) Close() error { + return nil +} + +func (m *fakeDNSTransportManager) Transports() []adapter.DNSTransport { + return nil +} + +func (m *fakeDNSTransportManager) Transport(tag string) (adapter.DNSTransport, bool) { + transport, loaded := m.transports[tag] + return transport, loaded +} + +func (m *fakeDNSTransportManager) Default() adapter.DNSTransport { + return m.defaultTransport +} + +func (m *fakeDNSTransportManager) FakeIP() adapter.FakeIPTransport { + return nil +} + +func (m *fakeDNSTransportManager) Remove(tag string) error { + return nil +} + +func (m *fakeDNSTransportManager) Create(ctx context.Context, logger log.ContextLogger, tag string, outboundType string, options any) error { + return nil +} + +func raceTestRouter(t *testing.T, transports ...*fakeDNSTransport) *Router { + transportMap := make(map[string]adapter.DNSTransport) + for _, transport := range transports { + transportMap[transport.tag] = transport + } + return &Router{ + ctx: context.Background(), + logger: log.NewNOPFactory().Logger(), + transport: &fakeDNSTransportManager{ + transports: transportMap, + defaultTransport: transportMap["final"], + }, + client: NewClient(ClientOptions{ + Context: context.Background(), + DisableCache: true, + Logger: log.NewNOPFactory().Logger(), + }), + } +} + +func raceTestRules(t *testing.T, rawRules []option.DNSRule) []adapter.DNSRule { + rules := make([]adapter.DNSRule, 0, len(rawRules)) + for _, rawRule := range rawRules { + rule, err := R.NewDNSRule(context.Background(), log.NewNOPFactory().Logger(), rawRule, true, false) + require.NoError(t, err) + rules = append(rules, rule) + } + return rules +} + +func raceTestExchange(router *Router, rules []adapter.DNSRule) exchangeWithRulesResult { + message := &mDNS.Msg{ + MsgHdr: mDNS.MsgHdr{ + Id: 1, + RecursionDesired: true, + }, + Question: []mDNS.Question{{ + Name: "race.example.org.", + Qtype: mDNS.TypeA, + Qclass: mDNS.ClassINET, + }}, + } + metadata := &adapter.InboundContext{ + Domain: "race.example.org", + QueryType: mDNS.TypeA, + } + ctx := adapter.WithContext(context.Background(), metadata) + return router.exchangeWithRules(ctx, rules, message, adapter.DNSQueryOptions{}, false) +} + +func evaluateRule(server string, tag string, speculative bool) option.DNSRule { + return option.DNSRule{ + Type: "", + DefaultOptions: option.DefaultDNSRule{ + DNSRuleAction: option.DNSRuleAction{ + Action: C.RuleActionTypeEvaluate, + RouteOptions: option.DNSRouteActionOptions{ + Server: server, + Tag: tag, + Speculative: speculative, + }, + }, + }, + } +} + +func respondRule(responseTag string, race bool, requireSuccess bool) option.DNSRule { + rule := option.DNSRule{ + Type: "", + DefaultOptions: option.DefaultDNSRule{ + RawDefaultDNSRule: option.RawDefaultDNSRule{ + MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: responseTag}, + }, + DNSRuleAction: option.DNSRuleAction{ + Action: C.RuleActionTypeRespond, + Race: race, + }, + }, + } + if requireSuccess { + successRcode := option.DNSRCode(mDNS.RcodeSuccess) + rule.DefaultOptions.ResponseRcode = &successRcode + } + return rule +} + +func routeRule(server string, speculative bool) option.DNSRule { + return option.DNSRule{ + Type: "", + DefaultOptions: option.DefaultDNSRule{ + DNSRuleAction: option.DNSRuleAction{ + Action: C.RuleActionTypeRoute, + RouteOptions: option.DNSRouteActionOptions{ + Server: server, + Speculative: speculative, + }, + }, + }, + } +} + +func responseAddress(t *testing.T, response *mDNS.Msg) netip.Addr { + require.NotNil(t, response) + require.Len(t, response.Answer, 1) + record, isA := response.Answer[0].(*mDNS.A) + require.True(t, isA) + address, _ := netip.AddrFromSlice(record.A) + return address.Unmap() +} + +// Both evaluate queries must launch in parallel, and a failed primary must +// fall through to the secondary instead of failing the request. +func TestDNSEvaluateParallelFallback(t *testing.T) { + t.Parallel() + transportX := &fakeDNSTransport{tag: "x", delay: 200 * time.Millisecond, exchangeErr: context.DeadlineExceeded} + transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + router := raceTestRouter(t, transportX, transportY) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", false, true), + respondRule("y", false, true), + }) + startTime := time.Now() + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response)) + require.Equal(t, int32(1), transportX.queryCount.Load()) + require.Equal(t, int32(1), transportY.queryCount.Load()) + require.Less(t, transportY.firstQueried.Sub(transportX.firstQueried), 100*time.Millisecond) + require.Less(t, time.Since(startTime), 350*time.Millisecond) +} + +// The first race rule whose response arrives and matches must commit +// immediately, without waiting for the slower rule written before it. +func TestDNSRaceFastestWins(t *testing.T) { + t.Parallel() + transportX := &fakeDNSTransport{tag: "x", delay: 500 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")} + transportY := &fakeDNSTransport{tag: "y", delay: 20 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + router := raceTestRouter(t, transportX, transportY) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", true, true), + respondRule("y", true, true), + }) + startTime := time.Now() + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response)) + require.Less(t, time.Since(startTime), 400*time.Millisecond) +} + +// Without race, rule order decides even when a later response arrives first. +func TestDNSOrderedReadsPreferEarlierRule(t *testing.T) { + t.Parallel() + transportX := &fakeDNSTransport{tag: "x", delay: 200 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")} + transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + router := raceTestRouter(t, transportX, transportY) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", false, true), + respondRule("y", false, true), + }) + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response)) +} + +// A pending race rule must hold back the default route: the default server +// is never queried when the race rule hits, and is queried only after the +// race decision resolved when it misses. +func TestDNSRaceBarrierProtectsDefaultRoute(t *testing.T) { + t.Parallel() + transportHit := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")} + transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")} + router := raceTestRouter(t, transportHit, transportFinal) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + respondRule("x", true, true), + }) + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response)) + require.Equal(t, int32(0), transportFinal.queryCount.Load()) + + transportMiss := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError} + transportFinal = &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")} + router = raceTestRouter(t, transportMiss, transportFinal) + rules = raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + respondRule("x", true, true), + }) + startTime := time.Now() + result = raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response)) + require.Equal(t, int32(1), transportFinal.queryCount.Load()) + require.GreaterOrEqual(t, transportFinal.firstQueried.Sub(startTime), 90*time.Millisecond) +} + +// A speculative route launches while the race decision is pending, but its +// response is only used after the race rule missed. +func TestDNSSpeculativeRoute(t *testing.T) { + t.Parallel() + transportMiss := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError} + transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")} + router := raceTestRouter(t, transportMiss, transportFinal) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + respondRule("x", true, true), + routeRule("final", true), + }) + startTime := time.Now() + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response)) + require.Equal(t, int32(1), transportFinal.queryCount.Load()) + require.Less(t, transportFinal.firstQueried.Sub(startTime), 90*time.Millisecond) + require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond) + + transportHit := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")} + transportFinal = &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")} + router = raceTestRouter(t, transportHit, transportFinal) + rules = raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + respondRule("x", true, true), + routeRule("final", true), + }) + result = raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response)) + require.Equal(t, int32(1), transportFinal.queryCount.Load()) +} + +// A matched rule without race must not take effect while a race rule is +// still pending: a race hit wins even when the other rule matched earlier, +// and on a race miss the other rule takes effect only after that decision. +func TestDNSNonRaceCommitWaitsForPendingRace(t *testing.T) { + t.Parallel() + transportHit := &fakeDNSTransport{tag: "x", delay: 150 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")} + transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + router := raceTestRouter(t, transportHit, transportY) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", true, true), + respondRule("y", false, true), + }) + startTime := time.Now() + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response)) + require.GreaterOrEqual(t, time.Since(startTime), 140*time.Millisecond) + + transportMiss := &fakeDNSTransport{tag: "x", delay: 150 * time.Millisecond, rcode: mDNS.RcodeNameError} + transportY = &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + router = raceTestRouter(t, transportMiss, transportY) + rules = raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", true, true), + respondRule("y", false, true), + }) + startTime = time.Now() + result = raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response)) + require.GreaterOrEqual(t, time.Since(startTime), 140*time.Millisecond) +} + +// speculative on a route rule with match_response launches the route query as +// soon as the rule matched, while its response is only used after the pending +// race rule missed. +func TestDNSSpeculativeRouteOnBindingRule(t *testing.T) { + t.Parallel() + transportMiss := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError} + transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + transportFinal := &fakeDNSTransport{tag: "final", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.9")} + router := raceTestRouter(t, transportMiss, transportY, transportFinal) + successRcode := option.DNSRCode(mDNS.RcodeSuccess) + boundRouteRule := option.DNSRule{ + Type: "", + DefaultOptions: option.DefaultDNSRule{ + RawDefaultDNSRule: option.RawDefaultDNSRule{ + MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: "y"}, + ResponseRcode: &successRcode, + }, + DNSRuleAction: option.DNSRuleAction{ + Action: C.RuleActionTypeRoute, + RouteOptions: option.DNSRouteActionOptions{ + Server: "final", + Speculative: true, + }, + }, + }, + } + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", true, true), + boundRouteRule, + }) + startTime := time.Now() + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.9"), responseAddress(t, result.response)) + require.Equal(t, int32(1), transportFinal.queryCount.Load()) + require.Less(t, transportFinal.firstQueried.Sub(startTime), 90*time.Millisecond) + require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond) +} + +// A race rule that rejects its response (NXDOMAIN vs required success) +// disarms and lets the other race rule win. +func TestDNSRaceSkipsRejectedResponse(t *testing.T) { + t.Parallel() + transportX := &fakeDNSTransport{tag: "x", delay: 10 * time.Millisecond, rcode: mDNS.RcodeNameError} + transportY := &fakeDNSTransport{tag: "y", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + router := raceTestRouter(t, transportX, transportY) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "x", false), + evaluateRule("y", "y", false), + respondRule("x", true, true), + respondRule("y", true, true), + }) + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.2"), responseAddress(t, result.response)) +} + +// A logical race rule is judged once all of its referenced responses arrived: +// it wins over a slower race rule when its sub-rules match, and on a miss the +// slower race rule takes over. +func TestDNSLogicalRace(t *testing.T) { + t.Parallel() + successRcode := option.DNSRCode(mDNS.RcodeSuccess) + logicalRule := func() option.DNSRule { + return option.DNSRule{ + Type: C.RuleTypeLogical, + LogicalOptions: option.LogicalDNSRule{ + RawLogicalDNSRule: option.RawLogicalDNSRule{ + Mode: C.LogicalTypeAnd, + Rules: []option.DNSRule{ + { + Type: C.RuleTypeDefault, + DefaultOptions: option.DefaultDNSRule{ + RawDefaultDNSRule: option.RawDefaultDNSRule{ + MatchResponse: &option.DNSRuleMatchResponse{Enabled: true}, + ResponseRcode: &successRcode, + }, + }, + }, + { + Type: C.RuleTypeDefault, + DefaultOptions: option.DefaultDNSRule{ + RawDefaultDNSRule: option.RawDefaultDNSRule{ + MatchResponse: &option.DNSRuleMatchResponse{Enabled: true, Tag: "y"}, + ResponseRcode: &successRcode, + }, + }, + }, + }, + }, + DNSRuleAction: option.DNSRuleAction{ + Action: C.RuleActionTypeRespond, + Race: true, + }, + }, + } + } + + transportX := &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.1")} + transportY := &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + transportZ := &fakeDNSTransport{tag: "z", delay: 250 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.3")} + router := raceTestRouter(t, transportX, transportY, transportZ) + rules := raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "", false), + evaluateRule("y", "y", false), + evaluateRule("z", "z", false), + logicalRule(), + respondRule("z", true, true), + }) + startTime := time.Now() + result := raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.1"), responseAddress(t, result.response)) + require.GreaterOrEqual(t, time.Since(startTime), 90*time.Millisecond) + require.Less(t, time.Since(startTime), 240*time.Millisecond) + + transportX = &fakeDNSTransport{tag: "x", delay: 100 * time.Millisecond, rcode: mDNS.RcodeNameError} + transportY = &fakeDNSTransport{tag: "y", delay: 10 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.2")} + transportZ = &fakeDNSTransport{tag: "z", delay: 250 * time.Millisecond, rcode: mDNS.RcodeSuccess, address: netip.MustParseAddr("192.0.2.3")} + router = raceTestRouter(t, transportX, transportY, transportZ) + rules = raceTestRules(t, []option.DNSRule{ + evaluateRule("x", "", false), + evaluateRule("y", "y", false), + evaluateRule("z", "z", false), + logicalRule(), + respondRule("z", true, true), + }) + startTime = time.Now() + result = raceTestExchange(router, rules) + require.NoError(t, result.err) + require.Equal(t, netip.MustParseAddr("192.0.2.3"), responseAddress(t, result.response)) + require.GreaterOrEqual(t, time.Since(startTime), 240*time.Millisecond) +} diff --git a/docs/configuration/dns/rule.md b/docs/configuration/dns/rule.md index 9a29cedf5..f87709d11 100644 --- a/docs/configuration/dns/rule.md +++ b/docs/configuration/dns/rule.md @@ -562,7 +562,11 @@ Enable response-based matching. When enabled, this rule matches against the eval (set by a preceding [`evaluate`](/configuration/dns/rule_action/#evaluate) action) instead of only matching the original query. -The evaluated response can also be returned directly by a later [`respond`](/configuration/dns/rule_action/#respond) action. +`true` or the `tag` of an `evaluate` action: `true` matches against the response of the latest +`evaluate` action without `tag`; a tag matches against the response of the `evaluate` action with the tag. + +The evaluated response can also be returned directly by a later [`respond`](/configuration/dns/rule_action/#respond) action; +in a rule with a `match_response` tag, `respond` returns the tagged response. Required for Response Match Fields (`response_rcode`, `response_answer`, `response_ns`, `response_extra`). Also required for `ip_cidr`, `ip_is_private`, and `ip_accept_any` when used with `evaluate` or Response Match Fields. diff --git a/docs/configuration/dns/rule.zh.md b/docs/configuration/dns/rule.zh.md index 9f07ca74f..af398506b 100644 --- a/docs/configuration/dns/rule.zh.md +++ b/docs/configuration/dns/rule.zh.md @@ -552,7 +552,9 @@ Available values: `wifi`, `cellular`, `ethernet` and `other`. 启用响应匹配。启用后,此规则将匹配已评估的响应(由前序 [`evaluate`](/zh/configuration/dns/rule_action/#evaluate) 动作设置),而不仅是匹配原始查询。 -该已评估的响应也可以被后续的 [`respond`](/zh/configuration/dns/rule_action/#respond) 动作直接返回。 +可以为 `true` 或 `evaluate` 动作的 `tag`:`true` 匹配最近一条无 `tag` 的 `evaluate` 动作的响应;标签则匹配对应 `evaluate` 动作的响应。 + +该已评估的响应也可以被后续的 [`respond`](/zh/configuration/dns/rule_action/#respond) 动作直接返回;在带 `match_response` 标签的规则中,`respond` 返回该标签的响应。 响应匹配字段(`response_rcode`、`response_answer`、`response_ns`、`response_extra`)需要此选项。 当与 `evaluate` 或响应匹配字段一起使用时,`ip_cidr`、`ip_is_private` 和 `ip_accept_any` 也需要此选项。 diff --git a/docs/configuration/dns/rule_action.md b/docs/configuration/dns/rule_action.md index 3555d6ede..c9f9e16dd 100644 --- a/docs/configuration/dns/rule_action.md +++ b/docs/configuration/dns/rule_action.md @@ -8,7 +8,9 @@ icon: material/new-box :material-plus: [evaluate](#evaluate) :material-plus: [respond](#respond) :material-plus: [disable_optimistic_cache](#disable_optimistic_cache) - :material-plus: [timeout](#timeout) + :material-plus: [timeout](#timeout) + :material-plus: [race](#race) + :material-plus: [speculative](#speculative) !!! quote "Changes in sing-box 1.12.0" @@ -17,12 +19,50 @@ icon: material/new-box !!! question "Since sing-box 1.11.0" +### Structure + +```json +{ + "action": "", + "race": false, + + ... // Action Fields +} +``` + +#### action + +The action to perform. `route` will be used by default. + +#### race + +!!! question "Since sing-box 1.14.0" + +Only available with `route`, `respond`, `reject` and `predefined` actions. + +Requires [`match_response`](/configuration/dns/rule/#match_response) (for logical rules, in sub-rules). +Conflict with `speculative`. + +By default, rules are matched one after another in listed order: a rule with `match_response` +waits for its referenced responses, and no later rule is matched until it has been judged. + +A rule with `race` enabled does not hold this order: rule matching continues past it while its +referenced responses are still pending, so the matching of race rules runs in parallel — with +each other and with the rules after them. Each race rule is judged once its referenced +responses are available, and the first race rule that matches terminates rule evaluation +immediately; the remaining queries are canceled. + +Rules without `race` still take effect strictly in listed order: while a preceding race rule is +not yet judged, the action of any other matched rule is held until none of the race rules +matched. The result may therefore depend on server speed only among race rules. + ### route ```json { "action": "route", // default "server": "", + "speculative": false, "strategy": "", "disable_cache": false, "disable_optimistic_cache": false, @@ -40,6 +80,19 @@ icon: material/new-box Tag of target server. +#### speculative + +!!! question "Since sing-box 1.14.0" + +Conflict with `race`. Has no effect without a preceding `race` rule. + +By default, no query is sent in parallel with pending race rules: a matched `route` action +holds its query until none of the race rules matched. + +When `speculative` is enabled, the query is sent as soon as the rule matches, in parallel with +the pending race rules, and may be wasted: its response is still used only after none of the +race rules matched. + #### strategy !!! question "Since sing-box 1.12.0" @@ -90,6 +143,8 @@ Will override `dns.client_subnet`. { "action": "evaluate", "server": "", + "tag": "", + "speculative": false, "disable_cache": false, "disable_optimistic_cache": false, "rewrite_ttl": null, @@ -113,6 +168,25 @@ does not satisfy this requirement, because matching happens before the action ru Tag of target server. +#### tag + +Tag of the evaluated response. + +A tagged response is only referenced via [`match_response`](/configuration/dns/rule/#match_response) with the tag; +`match_response: true` references the response of the latest `evaluate` action without `tag`. + +#### speculative + +!!! question "Since sing-box 1.14.0" + +Has no effect without a preceding `race` rule. + +By default, no query is sent in parallel with pending race rules: a matched `evaluate` action +holds its query, and rule matching stops there, until none of the race rules matched. + +When `speculative` is enabled, the query is sent as soon as the rule matches, in parallel with +the pending race rules, and may be wasted: rule matching continues without waiting for them. + #### disable_cache Disable cache and save cache in this query. @@ -155,7 +229,7 @@ Will override `dns.client_subnet`. `respond` terminates rule evaluation and returns the evaluated response from a preceding [`evaluate`](/configuration/dns/rule_action/#evaluate) action. -This action does not send a new DNS query and has no extra options. +This action does not send a new DNS query. Only allowed after a preceding top-level `evaluate` rule. If the action is reached without an evaluated response at runtime, the request fails with an error instead of falling through to later rules. diff --git a/docs/configuration/dns/rule_action.zh.md b/docs/configuration/dns/rule_action.zh.md index 756051dd9..c5d1f1b4f 100644 --- a/docs/configuration/dns/rule_action.zh.md +++ b/docs/configuration/dns/rule_action.zh.md @@ -8,7 +8,9 @@ icon: material/new-box :material-plus: [evaluate](#evaluate) :material-plus: [respond](#respond) :material-plus: [disable_optimistic_cache](#disable_optimistic_cache) - :material-plus: [timeout](#timeout) + :material-plus: [timeout](#timeout) + :material-plus: [race](#race) + :material-plus: [speculative](#speculative) !!! quote "sing-box 1.12.0 中的更改" @@ -17,12 +19,43 @@ icon: material/new-box !!! question "自 sing-box 1.11.0 起" +### 结构 + +```json +{ + "action": "", + "race": false, + + ... // 动作字段 +} +``` + +#### action + +要执行的动作。默认使用 `route`。 + +#### race + +!!! question "自 sing-box 1.14.0 起" + +仅可用于 `route`、`respond`、`reject` 和 `predefined` 动作。 + +需要 [`match_response`](/zh/configuration/dns/rule/#match_response)(对 logical 规则,位于子规则中)。 +与 `speculative` 冲突。 + +默认情况下,规则逐条按顺序匹配:带 `match_response` 的规则等待其引用的响应,在它被判定之前不会匹配任何后续规则。 + +启用 `race` 的规则成为竞态规则,不再保持这一顺序:其引用的响应尚未到达时,规则匹配会越过它继续进行,因此竞态规则的匹配相互并行、也与后续规则并行。每条竞态规则在其引用的响应可用时被判定,首个匹配的竞态规则立即终止规则评估,其余查询将被取消。 + +未启用 `race` 的规则仍严格按顺序生效:只要前面还有未判定的竞态规则,其他已匹配规则的动作就被扣住,直到所有竞态规则均未匹配。因此只有竞态规则之间的结果取决于服务器速度。 + ### route ```json { "action": "route", // 默认 "server": "", + "speculative": false, "strategy": "", "disable_cache": false, "disable_optimistic_cache": false, @@ -40,6 +73,16 @@ icon: material/new-box 目标 DNS 服务器的标签。 +#### speculative + +!!! question "自 sing-box 1.14.0 起" + +与 `race` 冲突。没有前序竞态规则时无效果。 + +默认情况下,查询决不与未判定的竞态规则并行发出:已匹配的 `route` 动作扣住其查询,直到所有竞态规则均未匹配后才发送。 + +启用 `speculative` 后,查询成为投机查询:在规则匹配时立即发出、与未判定的竞态规则并行,且可能被浪费;其响应仍仅在所有竞态规则均未匹配后才被使用。 + #### strategy !!! question "自 sing-box 1.12.0 起" @@ -90,6 +133,8 @@ icon: material/new-box { "action": "evaluate", "server": "", + "tag": "", + "speculative": false, "disable_cache": false, "disable_optimistic_cache": false, "rewrite_ttl": null, @@ -111,6 +156,23 @@ icon: material/new-box 目标 DNS 服务器的标签。 +#### tag + +已评估响应的标签。 + +带标签的响应仅能通过 [`match_response`](/zh/configuration/dns/rule/#match_response) 以标签引用; +`match_response: true` 引用最近一条无 `tag` 的 `evaluate` 动作的响应。 + +#### speculative + +!!! question "自 sing-box 1.14.0 起" + +没有前序竞态规则时无效果。 + +默认情况下,查询决不与未判定的竞态规则并行发出:已匹配的 `evaluate` 动作扣住其查询,规则匹配在此处停止,直到所有竞态规则均未匹配。 + +启用 `speculative` 后,查询成为投机查询:在规则匹配时立即发出、与未判定的竞态规则并行,且可能被浪费;规则匹配继续进行而不等待竞态规则。 + #### disable_cache 在此查询中禁用缓存。 @@ -153,7 +215,7 @@ icon: material/new-box `respond` 会终止规则评估,并直接返回前序 [`evaluate`](/zh/configuration/dns/rule_action/#evaluate) 动作保存的已评估的响应。 -此动作不会发起新的 DNS 查询,也没有额外选项。 +此动作不会发起新的 DNS 查询。 只能用于前面已有顶层 `evaluate` 规则的场景。如果运行时命中该动作时没有已评估的响应,则请求会直接返回错误,而不是继续匹配后续规则。 diff --git a/option/rule_action.go b/option/rule_action.go index a6f181f2d..75ea3910e 100644 --- a/option/rule_action.go +++ b/option/rule_action.go @@ -98,6 +98,7 @@ func (r *RuleAction) UnmarshalJSON(data []byte) error { type _DNSRuleAction struct { Action string `json:"action,omitempty"` + Race bool `json:"race,omitempty"` RouteOptions DNSRouteActionOptions `json:"-"` RouteOptionsOptions DNSRouteOptionsActionOptions `json:"-"` RejectOptions RejectActionOptions `json:"-"` @@ -160,7 +161,14 @@ func (r *DNSRuleAction) UnmarshalJSONContext(ctx context.Context, data []byte) e if v == nil { return json.UnmarshalDisallowUnknownFields(data, &_DNSRuleAction{}) } - return badjson.UnmarshallExcludedContext(ctx, data, (*_DNSRuleAction)(r), v) + err = badjson.UnmarshallExcludedContext(ctx, data, (*_DNSRuleAction)(r), v) + if err != nil { + return err + } + if r.Action == C.RuleActionTypeRoute && r.RouteOptions.Tag != "" { + return E.New("`tag` is only available in the `evaluate` action") + } + return nil } type RouteActionOptions struct { @@ -204,6 +212,8 @@ func (r *RouteOptionsActionOptions) UnmarshalJSON(data []byte) error { type DNSRouteActionOptions struct { Server string `json:"server,omitempty"` + Tag string `json:"tag,omitempty"` + Speculative bool `json:"speculative,omitempty"` Timeout badoption.Duration `json:"timeout,omitempty"` Strategy DomainStrategy `json:"strategy,omitempty"` DisableCache bool `json:"disable_cache,omitempty"` diff --git a/option/rule_dns.go b/option/rule_dns.go index ab1ddb24a..dc27bd6c2 100644 --- a/option/rule_dns.go +++ b/option/rule_dns.go @@ -67,6 +67,50 @@ func (r DNSRule) IsValid() bool { } } +type DNSRuleMatchResponse struct { + Enabled bool + Tag string +} + +func (m *DNSRuleMatchResponse) UnmarshalJSON(content []byte) error { + var boolValue bool + err := json.Unmarshal(content, &boolValue) + if err == nil { + m.Enabled = boolValue + m.Tag = "" + return nil + } + var stringValue string + err = json.Unmarshal(content, &stringValue) + if err != nil { + return E.New("invalid match_response value") + } + if stringValue == "" { + return E.New("empty match_response tag") + } + m.Enabled = true + m.Tag = stringValue + return nil +} + +func (m DNSRuleMatchResponse) MarshalJSON() ([]byte, error) { + if m.Tag != "" { + return json.Marshal(m.Tag) + } + return json.Marshal(m.Enabled) +} + +func (m *DNSRuleMatchResponse) IsEnabled() bool { + return m != nil && m.Enabled +} + +func (m *DNSRuleMatchResponse) ResponseTag() string { + if m == nil { + return "" + } + return m.Tag +} + type RawDefaultDNSRule struct { Inbound badoption.Listable[string] `json:"inbound,omitempty"` IPVersion int `json:"ip_version,omitempty"` @@ -106,7 +150,7 @@ type RawDefaultDNSRule struct { PreferredBy badoption.Listable[string] `json:"preferred_by,omitempty"` RuleSet badoption.Listable[string] `json:"rule_set,omitempty"` RuleSetIPCIDRMatchSource bool `json:"rule_set_ip_cidr_match_source,omitempty"` - MatchResponse bool `json:"match_response,omitempty"` + MatchResponse *DNSRuleMatchResponse `json:"match_response,omitempty"` IPCIDR badoption.Listable[string] `json:"ip_cidr,omitempty"` IPIsPrivate bool `json:"ip_is_private,omitempty"` IPAcceptAny bool `json:"ip_accept_any,omitempty"` diff --git a/route/rule/rule_action.go b/route/rule/rule_action.go index 5d64c12cb..7e38639e6 100644 --- a/route/rule/rule_action.go +++ b/route/rule/rule_action.go @@ -128,7 +128,8 @@ func NewDNSRuleAction(logger logger.ContextLogger, action option.DNSRuleAction) return nil case C.RuleActionTypeRoute: return &RuleActionDNSRoute{ - Server: action.RouteOptions.Server, + Server: action.RouteOptions.Server, + Speculative: action.RouteOptions.Speculative, RuleActionDNSRouteOptions: RuleActionDNSRouteOptions{ Strategy: C.DomainStrategy(action.RouteOptions.Strategy), Timeout: time.Duration(action.RouteOptions.Timeout), @@ -140,7 +141,9 @@ func NewDNSRuleAction(logger logger.ContextLogger, action option.DNSRuleAction) } case C.RuleActionTypeEvaluate: return &RuleActionEvaluate{ - Server: action.RouteOptions.Server, + Server: action.RouteOptions.Server, + Tag: action.RouteOptions.Tag, + Speculative: action.RouteOptions.Speculative, RuleActionDNSRouteOptions: RuleActionDNSRouteOptions{ Strategy: C.DomainStrategy(action.RouteOptions.Strategy), Timeout: time.Duration(action.RouteOptions.Timeout), @@ -285,7 +288,8 @@ func (r *RuleActionRouteOptions) Descriptions() []string { } type RuleActionDNSRoute struct { - Server string + Server string + Speculative bool RuleActionDNSRouteOptions } @@ -294,11 +298,13 @@ func (r *RuleActionDNSRoute) Type() string { } func (r *RuleActionDNSRoute) String() string { - return formatDNSRouteAction("route", r.Server, r.RuleActionDNSRouteOptions) + return formatDNSRouteAction("route", r.Server, r.Speculative, r.RuleActionDNSRouteOptions) } type RuleActionEvaluate struct { - Server string + Server string + Tag string + Speculative bool RuleActionDNSRouteOptions } @@ -307,7 +313,7 @@ func (r *RuleActionEvaluate) Type() string { } func (r *RuleActionEvaluate) String() string { - return formatDNSRouteAction("evaluate", r.Server, r.RuleActionDNSRouteOptions) + return formatDNSRouteAction("evaluate", r.Server, r.Speculative, r.RuleActionDNSRouteOptions) } type RuleActionRespond struct{} @@ -320,9 +326,12 @@ func (r *RuleActionRespond) String() string { return "respond" } -func formatDNSRouteAction(action string, server string, options RuleActionDNSRouteOptions) string { +func formatDNSRouteAction(action string, server string, speculative bool, options RuleActionDNSRouteOptions) string { var descriptions []string descriptions = append(descriptions, server) + if speculative { + descriptions = append(descriptions, "speculative") + } if options.DisableCache { descriptions = append(descriptions, "disable-cache") } diff --git a/route/rule/rule_dns.go b/route/rule/rule_dns.go index 2b7682828..facefa252 100644 --- a/route/rule/rule_dns.go +++ b/route/rule/rule_dns.go @@ -28,6 +28,9 @@ func NewDNSRule(ctx context.Context, logger log.ContextLogger, options option.DN if err != nil { return nil, err } + if options.DefaultOptions.Race && !options.DefaultOptions.MatchResponse.IsEnabled() { + return nil, E.New("`race` requires `match_response`") + } switch options.DefaultOptions.Action { case "", C.RuleActionTypeRoute, C.RuleActionTypeEvaluate: if options.DefaultOptions.RouteOptions.Server == "" && checkServer { @@ -62,6 +65,16 @@ func validateDNSRuleAction(action option.DNSRuleAction) error { if action.Action == C.RuleActionTypeReject && action.RejectOptions.Method == C.RuleActionRejectMethodReply { return E.New("reject method `reply` is not supported for DNS rules") } + if action.Race { + switch action.Action { + case "", C.RuleActionTypeRoute, C.RuleActionTypeRespond, C.RuleActionTypeReject, C.RuleActionTypePredefined: + default: + return E.New("`race` requires a final action") + } + if action.RouteOptions.Speculative { + return E.New("`race` and `speculative` cannot be combined on the same rule") + } + } return nil } @@ -69,7 +82,9 @@ var _ adapter.DNSRule = (*DefaultDNSRule)(nil) type DefaultDNSRule struct { abstractDefaultRule - matchResponse bool + matchResponse bool + matchResponseTag string + race bool } func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options option.DefaultDNSRule, legacyDNSMode bool) (*DefaultDNSRule, error) { @@ -78,7 +93,9 @@ func NewDefaultDNSRule(ctx context.Context, logger log.ContextLogger, options op invert: options.Invert, action: NewDNSRuleAction(logger, options.DNSRuleAction), }, - matchResponse: options.MatchResponse, + matchResponse: options.MatchResponse.IsEnabled(), + matchResponseTag: options.MatchResponse.ResponseTag(), + race: options.Race, } if len(options.Inbound) > 0 { item := NewInboundRule(options.Inbound) @@ -377,12 +394,36 @@ func (r *DefaultDNSRule) LegacyPreMatch(metadata *adapter.InboundContext) bool { return r.abstractDefaultRule.Match(metadata) } +func (r *DefaultDNSRule) MatchResponseTag() string { + return r.matchResponseTag +} + +func (r *DefaultDNSRule) MatchResponseTags() []string { + if r.matchResponseTag == "" { + return nil + } + return []string{r.matchResponseTag} +} + +func (r *DefaultDNSRule) MatchResponseAnonymous() bool { + return r.matchResponse && r.matchResponseTag == "" +} + +func (r *DefaultDNSRule) Race() bool { + return r.race +} + func (r *DefaultDNSRule) matchForMatch(metadata *adapter.InboundContext) bool { if r.matchResponse { - if metadata.DNSResponse == nil { + response := metadata.DNSResponse + if r.matchResponseTag != "" { + response = metadata.NamedDNSResponses[r.matchResponseTag] + } + if response == nil { return r.invert } matchMetadata := *metadata + matchMetadata.DNSResponse = response matchMetadata.DestinationAddressMatchFromResponse = true return r.abstractDefaultRule.Match(&matchMetadata) } @@ -400,17 +441,25 @@ var _ adapter.DNSRule = (*LogicalDNSRule)(nil) type LogicalDNSRule struct { abstractLogicalRule + matchResponseTags []string + matchResponseAnonymous bool + race bool } -func matchDNSHeadlessRuleForMatch(rule adapter.HeadlessRule, metadata *adapter.InboundContext) bool { - switch typedRule := rule.(type) { - case *DefaultDNSRule: - return typedRule.matchForMatch(metadata) - case *LogicalDNSRule: - return typedRule.matchForMatch(metadata) - default: - return typedRule.Match(metadata) - } +func (r *LogicalDNSRule) MatchResponseTag() string { + return "" +} + +func (r *LogicalDNSRule) MatchResponseTags() []string { + return r.matchResponseTags +} + +func (r *LogicalDNSRule) MatchResponseAnonymous() bool { + return r.matchResponseAnonymous +} + +func (r *LogicalDNSRule) Race() bool { + return r.race } func (r *LogicalDNSRule) matchForMatch(metadata *adapter.InboundContext) bool { @@ -420,7 +469,7 @@ func (r *LogicalDNSRule) matchForMatch(metadata *adapter.InboundContext) bool { for _, rule := range r.rules { nestedMetadata := *metadata nestedMetadata.ResetRuleCache() - if !matchDNSHeadlessRuleForMatch(rule, &nestedMetadata) { + if !rule.Match(&nestedMetadata) { matched = false break } @@ -429,7 +478,7 @@ func (r *LogicalDNSRule) matchForMatch(metadata *adapter.InboundContext) bool { for _, rule := range r.rules { nestedMetadata := *metadata nestedMetadata.ResetRuleCache() - if matchDNSHeadlessRuleForMatch(rule, &nestedMetadata) { + if rule.Match(&nestedMetadata) { matched = true break } @@ -448,6 +497,7 @@ func NewLogicalDNSRule(ctx context.Context, logger log.ContextLogger, options op invert: options.Invert, action: NewDNSRuleAction(logger, options.DNSRuleAction), }, + race: options.Race, } switch options.Mode { case C.LogicalTypeAnd: @@ -468,6 +518,16 @@ func NewLogicalDNSRule(ctx context.Context, logger log.ContextLogger, options op } r.rules[i] = rule } + for _, subRule := range r.rules { + if dnsRule, isDNSRule := subRule.(adapter.DNSRule); isDNSRule { + r.matchResponseTags = append(r.matchResponseTags, dnsRule.MatchResponseTags()...) + r.matchResponseAnonymous = r.matchResponseAnonymous || dnsRule.MatchResponseAnonymous() + } + } + r.matchResponseTags = common.Uniq(r.matchResponseTags) + if r.race && len(r.matchResponseTags) == 0 && !r.matchResponseAnonymous { + return nil, E.New("`race` requires `match_response` in sub-rules") + } return r, nil } @@ -477,15 +537,8 @@ func (r *LogicalDNSRule) Action() adapter.RuleAction { func (r *LogicalDNSRule) WithAddressLimit() bool { for _, rawRule := range r.rules { - switch rule := rawRule.(type) { - case *DefaultDNSRule: - if rule.WithAddressLimit() { - return true - } - case *LogicalDNSRule: - if rule.WithAddressLimit() { - return true - } + if dnsRule, isDNSRule := rawRule.(adapter.DNSRule); isDNSRule && dnsRule.WithAddressLimit() { + return true } } return false