From 7690e10ee8652097dbd6d3b01c0e0a25e60c5545 Mon Sep 17 00:00:00 2001 From: vircoys Date: Wed, 29 Jul 2026 15:54:30 +0800 Subject: [PATCH] fix(dialtesting): validate rendered netpath tasks --- dialtesting/netpath.go | 184 ++++++++++++++++++++++++++++++++---- dialtesting/netpath_test.go | 77 +++++++++++++++ 2 files changed, 245 insertions(+), 16 deletions(-) diff --git a/dialtesting/netpath.go b/dialtesting/netpath.go index 4e824ba8..3be6f5f2 100644 --- a/dialtesting/netpath.go +++ b/dialtesting/netpath.go @@ -119,6 +119,16 @@ func (t *NetPathTask) init() error { return nil } + // CheckTask runs before variables are rendered by the debug API. Defer + // value validation until RenderTemplateAndInit replaces every template. + if t.rawTask == nil && t.hasUnrenderedTemplate() { + return nil + } + + if err := t.check(); err != nil { + return err + } + config, err := t.probeConfig(time.Time{}) if err != nil { return err @@ -130,6 +140,10 @@ func (t *NetPathTask) init() error { } func (t *NetPathTask) check() error { + if t.rawTask == nil && t.hasUnrenderedTemplate() { + return nil + } + if strings.TrimSpace(t.Host) == "" { return errors.New("host should not be empty") } @@ -159,6 +173,25 @@ func (t *NetPathTask) check() error { return nil } +func (t *NetPathTask) hasUnrenderedTemplate() bool { + data, err := json.Marshal(struct { + Protocol string `json:"protocol"` + Host string `json:"host"` + Port string `json:"port"` + AdvanceOptions NetPathAdvanceOptions `json:"advance_options"` + SuccessWhen []*NetPathSuccess `json:"success_when"` + SuccessWhenLogic string `json:"success_when_logic"` + }{ + Protocol: t.Protocol, + Host: t.Host, + Port: t.Port, + AdvanceOptions: t.AdvanceOptions, + SuccessWhen: t.SuccessWhen, + SuccessWhenLogic: t.SuccessWhenLogic, + }) + return err == nil && hasTemplateTags(string(data)) +} + func (t *NetPathTask) run() error { executor := t.getExecutor() if executor == nil { @@ -271,25 +304,144 @@ func (t *NetPathTask) renderTemplate(fm template.FuncMap) error { t.rawTask = rawTask } - raw, err := json.Marshal(t.rawTask) + rawTask := t.rawTask + if rawTask == nil { + return errors.New("raw NETPATH task is nil") + } + + var err error + if t.Protocol, err = t.renderNetPathString("protocol", rawTask.Protocol, fm); err != nil { + return err + } + if t.Host, err = t.renderNetPathString("host", rawTask.Host, fm); err != nil { + return err + } + if t.Port, err = t.renderNetPathString("port", rawTask.Port, fm); err != nil { + return err + } + if t.SuccessWhenLogic, err = t.renderNetPathString( + "success_when_logic", + rawTask.SuccessWhenLogic, + fm, + ); err != nil { + return err + } + + t.AdvanceOptions = rawTask.AdvanceOptions + if t.AdvanceOptions.Timeout, err = t.renderNetPathString( + "advance_options.timeout", + rawTask.AdvanceOptions.Timeout, + fm, + ); err != nil { + return err + } + if t.AdvanceOptions.SourceName, err = t.renderNetPathString( + "advance_options.source_name", + rawTask.AdvanceOptions.SourceName, + fm, + ); err != nil { + return err + } + if t.AdvanceOptions.TargetName, err = t.renderNetPathString( + "advance_options.target_name", + rawTask.AdvanceOptions.TargetName, + fm, + ); err != nil { + return err + } + + t.SuccessWhen, err = cloneNetPathSuccessWhen(rawTask.SuccessWhen) + if err != nil { + return err + } + if err := t.renderNetPathSuccessWhen(fm); err != nil { + return err + } + return nil +} + +func (t *NetPathTask) renderNetPathString( + field string, + value string, + fm template.FuncMap, +) (string, error) { + rendered, err := t.GetParsedString(value, fm) if err != nil { - return fmt.Errorf("marshal raw NETPATH task failed: %w", err) + return "", fmt.Errorf("render %s failed: %w", field, err) + } + return rendered, nil +} + +func cloneNetPathSuccessWhen(successWhen []*NetPathSuccess) ([]*NetPathSuccess, error) { + data, err := json.Marshal(successWhen) + if err != nil { + return nil, fmt.Errorf("marshal raw NETPATH success_when failed: %w", err) + } + var cloned []*NetPathSuccess + if err := json.Unmarshal(data, &cloned); err != nil { + return nil, fmt.Errorf("unmarshal raw NETPATH success_when failed: %w", err) + } + return cloned, nil +} + +func (t *NetPathTask) renderNetPathSuccessWhen(fm template.FuncMap) error { + for _, success := range t.SuccessWhen { + if success == nil { + continue + } + for _, assertions := range success.assertions() { + for _, condition := range assertions.conditions { + if condition == nil { + continue + } + op, err := t.renderNetPathString(assertions.field+" operator", condition.Op, fm) + if err != nil { + return err + } + condition.Op = op + if err := t.renderNetPathConditionTarget(assertions.field, condition, fm); err != nil { + return err + } + } + } + } + return nil +} + +func (t *NetPathTask) renderNetPathConditionTarget( + field string, + condition *NetPathCondition, + fm template.FuncMap, +) error { + var target string + if err := json.Unmarshal(condition.Target, &target); err != nil { + return nil } - rendered, err := t.GetParsedString(string(raw), fm) + + rendered, err := t.renderNetPathString(field+" target", target, fm) if err != nil { - return fmt.Errorf("render NETPATH task failed: %w", err) + return err } - task := &NetPathTask{} - if err := json.Unmarshal([]byte(rendered), task); err != nil { - return fmt.Errorf("unmarshal rendered NETPATH task failed: %w", err) + // Config variables are strings. For numeric assertions, convert a + // templated string such as "{{loss}}" back to a JSON number after + // rendering; ordinary string targets remain strings and validation + // reports their type mismatch as before. + if !isNetPathStatusField(field) && + !isNetPathDurationField(field) && + hasTemplateTags(target) { + var number *float64 + if err := json.Unmarshal([]byte(rendered), &number); err == nil && number != nil { + condition.Target = json.RawMessage(rendered) + return nil + } } - t.Protocol = task.Protocol - t.Host = task.Host - t.Port = task.Port - t.AdvanceOptions = task.AdvanceOptions - t.SuccessWhen = task.SuccessWhen - t.SuccessWhenLogic = task.SuccessWhenLogic + + data, err := json.Marshal(rendered) + if err != nil { + return fmt.Errorf("marshal rendered %s target failed: %w", field, err) + } + condition.Target = data return nil } @@ -479,9 +631,9 @@ func (c *NetPathCondition) durationTarget() (time.Duration, error) { } func (c *NetPathCondition) numberTarget() (float64, error) { - var target float64 - if err := json.Unmarshal(c.Target, &target); err != nil { + var target *float64 + if err := json.Unmarshal(c.Target, &target); err != nil || target == nil { return 0, errors.New("target must be numeric") } - return target, nil + return *target, nil } diff --git a/dialtesting/netpath_test.go b/dialtesting/netpath_test.go index 3e4bea28..89c65336 100644 --- a/dialtesting/netpath_test.go +++ b/dialtesting/netpath_test.go @@ -143,6 +143,15 @@ func TestNetPathTaskValidation(t *testing.T) { }, want: "success_when is required", }, + { + name: "numeric assertion rejects null", + replace: func(task map[string]any) { + success := task["success_when"].([]any)[0].(map[string]any) + condition := success["e2e_probe_loss_percent"].([]any)[0].(map[string]any) + condition["target"] = nil + }, + want: "target must be numeric", + }, } for _, test := range tests { @@ -162,6 +171,74 @@ func TestNetPathTaskValidation(t *testing.T) { } } +func TestNetPathTaskRendersAndValidatesTemplates(t *testing.T) { + var raw map[string]any + require.NoError(t, json.Unmarshal([]byte(validNetPathTaskJSON()), &raw)) + raw["protocol"] = "{{protocol}}" + raw["host"] = "{{host}}" + raw["port"] = "{{port}}" + raw["success_when_logic"] = "{{logic}}" + advanceOptions := raw["advance_options"].(map[string]any) + advanceOptions["timeout"] = "{{timeout}}" + advanceOptions["source_name"] = "{{source}}" + success := raw["success_when"].([]any)[0].(map[string]any) + loss := success["e2e_probe_loss_percent"].([]any)[0].(map[string]any) + loss["target"] = "{{loss}}" + status := success["e2e_status"].([]any)[0].(map[string]any) + status["target"] = "{{status}}" + raw["config_vars"] = []any{ + map[string]any{"name": "protocol", "value": " TCP "}, + map[string]any{"name": "host", "value": "rendered.example.com"}, + map[string]any{"name": "port", "value": "8443"}, + map[string]any{"name": "logic", "value": "or"}, + map[string]any{"name": "timeout", "value": "5s"}, + map[string]any{"name": "source", "value": `source "quoted"`}, + map[string]any{"name": "loss", "value": "2.5"}, + map[string]any{"name": "status", "value": "reached"}, + } + data, err := json.Marshal(raw) + require.NoError(t, err) + + child, err := CreateTaskChild(ClassNetPath) + require.NoError(t, err) + task, err := NewTask(string(data), child) + require.NoError(t, err) + + // The debug path validates before rendering. Template placeholders must + // not be parsed as concrete NetPath values at this stage. + require.NoError(t, task.CheckTask()) + require.NoError(t, task.RenderTemplateAndInit(nil)) + + netPathTask := task.(*NetPathTask) + assert.Equal(t, " TCP ", netPathTask.Protocol) + assert.Equal(t, "rendered.example.com", netPathTask.Host) + assert.Equal(t, "8443", netPathTask.Port) + assert.Equal(t, "5s", netPathTask.AdvanceOptions.Timeout) + assert.Equal(t, `source "quoted"`, netPathTask.AdvanceOptions.SourceName) + assert.Equal(t, "or", netPathTask.SuccessWhenLogic) + + var number float64 + require.NoError(t, json.Unmarshal( + netPathTask.SuccessWhen[0].E2EProbeLossPercent[0].Target, + &number, + )) + assert.Equal(t, 2.5, number) + assert.JSONEq(t, `"reached"`, string(netPathTask.SuccessWhen[0].E2EStatus[0].Target)) +} + +func TestNetPathTaskRejectsInvalidRenderedEndpoint(t *testing.T) { + raw := strings.Replace(validNetPathTaskJSON(), `"tcp"`, `"{{protocol}}"`, 1) + raw = strings.Replace(raw, `"config_vars": [`, `"config_vars": [ + {"name": "protocol", "value": "http"},`, 1) + + child, err := CreateTaskChild(ClassNetPath) + require.NoError(t, err) + task, err := NewTask(raw, child) + require.NoError(t, err) + require.NoError(t, task.CheckTask()) + require.ErrorContains(t, task.RenderTemplateAndInit(nil), "unsupported protocol") +} + func TestNetPathTaskTemplateAndCancellation(t *testing.T) { raw := strings.ReplaceAll(validNetPathTaskJSON(), "example.com", "{{host}}") raw = strings.ReplaceAll(raw, `"443"`, `"{{port}}"`)