Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
184 changes: 168 additions & 16 deletions dialtesting/netpath.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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")
}
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}

Expand Down Expand Up @@ -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
}
77 changes: 77 additions & 0 deletions dialtesting/netpath_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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}}"`)
Expand Down
Loading