diff --git a/HACKING.md b/HACKING.md index 550551d..6995b84 100644 --- a/HACKING.md +++ b/HACKING.md @@ -357,14 +357,14 @@ Built-in executors registered at startup: - `HTTPExecutor` - handles `http` steps - `LLMExecutor` - handles `llm` steps -## Run Registry +## Run Control Plane -The run registry tracks active workflow executions for cancellation support: +The run control plane tracks active workflow executions for cancellation support: ```go -// internal/executor/run_registry.go +// internal/executor/run_control_plane.go -type RunRegistry struct { +type RunControlPlane struct { mu sync.RWMutex runs map[string]*ActiveRun } @@ -377,13 +377,13 @@ type ActiveRun struct { } // Key operations -func (r *RunRegistry) Register(runUUID string, cancel context.CancelFunc) *ActiveRun -func (r *RunRegistry) Cancel(runUUID string) ([]int, error) // Returns killed PIDs -func (r *RunRegistry) AddPID(runUUID string, pid int) -func (r *RunRegistry) RemovePID(runUUID string, pid int) +func (r *RunControlPlane) Register(runUUID string, cancel context.CancelFunc) *ActiveRun +func (r *RunControlPlane) Cancel(runUUID string) ([]int, error) // Returns killed PIDs +func (r *RunControlPlane) AddPID(runUUID string, pid int) +func (r *RunControlPlane) RemovePID(runUUID string, pid int) ``` -The registry is accessed via `GetRunRegistry()` singleton. When a run is cancelled: +The control plane is accessed via `GetRunControlPlane()` singleton. When a run is cancelled: 1. Context is cancelled to stop new operations 2. All tracked PIDs are killed via SIGKILL (including process groups) @@ -991,6 +991,8 @@ type Run struct { TriggerName string TotalSteps int CompletedSteps int + RunPriority string // "low", "normal", "high", "critical" + RunMode string // "local", "distributed", "cloud" CreatedAt time.Time UpdatedAt time.Time } diff --git a/internal/client/types.go b/internal/client/types.go index 5af0c90..6acd384 100644 --- a/internal/client/types.go +++ b/internal/client/types.go @@ -55,6 +55,8 @@ type Run struct { TotalSteps int `json:"total_steps"` CompletedSteps int `json:"completed_steps"` Workspace string `json:"workspace"` + RunPriority string `json:"run_priority,omitempty"` + RunMode string `json:"run_mode,omitempty"` StartedAt *time.Time `json:"started_at,omitempty"` CompletedAt *time.Time `json:"completed_at,omitempty"` ErrorMessage string `json:"error_message,omitempty"` @@ -71,6 +73,7 @@ type CreateRunRequest struct { Params map[string]string `json:"params,omitempty"` Concurrency int `json:"concurrency,omitempty"` Priority string `json:"priority,omitempty"` + RunMode string `json:"run_mode,omitempty"` Timeout int `json:"timeout,omitempty"` RunnerType string `json:"runner_type,omitempty"` ThreadsHold int `json:"threads_hold,omitempty"` diff --git a/internal/core/step.go b/internal/core/step.go index 57e422e..94becf7 100644 --- a/internal/core/step.go +++ b/internal/core/step.go @@ -206,10 +206,11 @@ type Step struct { ParallelSteps []Step `yaml:"parallel_steps"` // Foreach step fields - Input string `yaml:"input"` - Variable string `yaml:"variable"` - Threads StepThreads `yaml:"threads,omitempty"` - Step *Step `yaml:"step"` + Input string `yaml:"input"` + Variable string `yaml:"variable"` + VariablePreProcess string `yaml:"variable_pre_process,omitempty"` // Transform each input line before storing in variable + Threads StepThreads `yaml:"threads,omitempty"` + Step *Step `yaml:"step"` // Remote-bash step fields StepRunnerConfig *StepRunnerConfig `yaml:"step_runner_config"` diff --git a/internal/database/models.go b/internal/database/models.go index ab9ad6a..1d5d7a9 100644 --- a/internal/database/models.go +++ b/internal/database/models.go @@ -39,6 +39,10 @@ type Run struct { // Process tracking - current running process ID for cancellation support CurrentPID int `bun:"current_pid" json:"current_pid,omitempty"` + // Priority and mode + RunPriority string `bun:"run_priority,notnull,default:'high'" json:"run_priority"` // low, normal, high, critical + RunMode string `bun:"run_mode,notnull,default:'local'" json:"run_mode"` // local, distributed, cloud + // Relations Steps []*StepResult `bun:"rel:has-many,join:id=run_id" json:"steps,omitempty"` Artifacts []*Artifact `bun:"rel:has-many,join:id=run_id" json:"artifacts,omitempty"` diff --git a/internal/database/seed.go b/internal/database/seed.go index 32b4191..93d5869 100644 --- a/internal/database/seed.go +++ b/internal/database/seed.go @@ -2639,9 +2639,9 @@ func ListTables(ctx context.Context) ([]TableInfo, error) { // tableSearchColumns defines which columns to search for each table var tableSearchColumns = map[string][]string{ - "runs": {"id", "run_uuid", "run_group_id", "workflow_name", "target", "status", "error_message"}, + "runs": {"id", "run_uuid", "run_group_id", "workflow_name", "target", "workspace", "status", "error_message"}, "step_results": {"id", "run_id", "step_name", "step_type", "status", "command", "output", "error_message"}, - "artifacts": {"id", "run_id", "name", "path", "type", "description"}, + "artifacts": {"id", "run_id", "workspace", "name", "artifact_path", "artifact_type", "description"}, "assets": {"workspace", "asset_value", "url", "title", "host_ip", "source", "labels"}, "event_logs": {"event_id", "topic", "name", "source", "workspace", "run_id", "workflow_name", "data"}, "schedules": {"id", "name", "workflow_name", "workflow_kind", "target", "trigger_name", "schedule"}, @@ -2653,9 +2653,9 @@ var tableSearchColumns = map[string][]string{ // tableDisplayColumns defines which columns to display by default for each table (ordered) var tableDisplayColumns = map[string][]string{ - "runs": {"run_uuid", "workflow_name", "target", "trigger_type", "status", "completed_steps", "total_steps", "started_at"}, + "runs": {"run_uuid", "workflow_name", "target", "workspace", "trigger_type", "status", "completed_steps", "total_steps", "started_at"}, "step_results": {"step_name", "step_type", "status", "duration_ms", "command"}, - "artifacts": {"name", "path", "type", "size_bytes", "line_count"}, + "artifacts": {"name", "artifact_path", "artifact_type", "content_type", "size_bytes", "line_count"}, "assets": {"asset_value", "host_ip", "title", "status_code", "last_seen_at", "url"}, "event_logs": {"topic", "name", "source", "workspace", "created_at"}, "schedules": {"name", "workflow_name", "workflow_kind", "target", "trigger_type", "schedule", "is_enabled"}, @@ -2668,19 +2668,20 @@ var tableDisplayColumns = map[string][]string{ // tableAllColumns defines ALL columns for each table (ordered, matching model structs) var tableAllColumns = map[string][]string{ "runs": {"id", "run_uuid", "run_group_id", "workflow_name", "workflow_kind", "target", "params", - "status", "workspace_path", "started_at", "completed_at", "error_message", + "status", "workspace", "started_at", "completed_at", "error_message", "schedule_id", "trigger_type", "trigger_name", "total_steps", - "completed_steps", "created_at", "updated_at"}, + "completed_steps", "current_pid", "run_priority", "run_mode", "created_at", "updated_at"}, "step_results": {"id", "run_id", "step_name", "step_type", "status", "command", "output", "error_message", "exports", "duration_ms", "log_file", "started_at", "completed_at", "created_at"}, - "artifacts": {"id", "run_id", "name", "path", "type", "size_bytes", - "line_count", "description", "created_at"}, + "artifacts": {"id", "run_id", "workspace", "name", "artifact_path", "artifact_type", + "content_type", "size_bytes", "line_count", "description", "created_at"}, "assets": {"id", "workspace", "asset_value", "url", "input", "scheme", "method", "path", "status_code", "content_type", "content_length", "title", "words", "lines", "host_ip", "dns_records", "tls", "asset_type", "technologies", - "response_time", "labels", "source", "last_seen_at", "created_at", "updated_at"}, - "event_logs": {"id", "topic", "event_id", "name", "source", "data_type", "data", + "response_time", "labels", "source", "raw_json_data", "raw_response", + "screenshot_base64_data", "last_seen_at", "created_at", "updated_at"}, + "event_logs": {"id", "topic", "event_id", "name", "source_type", "source", "data_type", "data", "workspace", "run_id", "workflow_name", "processed", "processed_at", "error", "created_at"}, "schedules": {"id", "name", "workflow_name", "workflow_kind", "target", "workspace", @@ -2690,7 +2691,8 @@ var tableAllColumns = map[string][]string{ "total_urls", "total_ips", "total_links", "total_content", "total_archive", "total_vulns", "vuln_critical", "vuln_high", "vuln_medium", "vuln_low", "vuln_potential", "risk_score", "tags", "last_run", "run_workflow", - "created_at", "updated_at"}, + "state_execution_log", "state_completed_file", "state_workflow_file", + "state_workflow_folder", "created_at", "updated_at"}, "vulnerabilities": {"id", "workspace", "vuln_info", "vuln_title", "vuln_desc", "vuln_poc", "severity", "confidence", "asset_type", "asset_value", "tags", "detail_http_request", "detail_http_response", "raw_vuln_json", diff --git a/internal/executor/dispatcher.go b/internal/executor/dispatcher.go index 379bdb3..789f1e1 100644 --- a/internal/executor/dispatcher.go +++ b/internal/executor/dispatcher.go @@ -106,7 +106,7 @@ func NewStepDispatcherWithConfig(cfg StepDispatcherConfig) *StepDispatcher { d.registry.Register(d.bashExecutor) d.registry.Register(NewFunctionExecutor(engine, d.functionRegistry)) d.registry.Register(NewParallelExecutor(d)) - d.registry.Register(NewForeachExecutor(d, engine)) + d.registry.Register(NewForeachExecutor(d, engine, d.functionRegistry)) d.registry.Register(NewRemoteBashExecutor(engine)) d.registry.Register(NewHTTPExecutor(engine)) d.registry.Register(d.llmExecutor) @@ -247,6 +247,7 @@ func collectRenderRequests(step *core.Step) []template.RenderRequest { add("StdFile", step.StdFile) add("Function", step.Function) add("Input", step.Input) + add("VariablePreProcess", step.VariablePreProcess) add("Log", step.Log) add("Timeout", string(step.Timeout)) add("Threads", string(step.Threads)) @@ -360,6 +361,9 @@ func (d *StepDispatcher) renderStepBatch(step *core.Step, vars map[string]any) ( if v := get("Input"); v != "" { rendered.Input = v } + if v := get("VariablePreProcess"); v != "" { + rendered.VariablePreProcess = v + } if v := get("Log"); v != "" { rendered.Log = v } @@ -641,6 +645,13 @@ func (d *StepDispatcher) renderStepSequential(step *core.Step, vars map[string]a } rendered.Input = input } + if step.VariablePreProcess != "" { + vpp, err := d.templateEngine.Render(step.VariablePreProcess, vars) + if err != nil { + return nil, fmt.Errorf("error rendering variable_pre_process: %w", err) + } + rendered.VariablePreProcess = vpp + } // Render log message if step.Log != "" { diff --git a/internal/executor/executor.go b/internal/executor/executor.go index 0efdf77..02a4809 100644 --- a/internal/executor/executor.go +++ b/internal/executor/executor.go @@ -717,16 +717,16 @@ func (e *Executor) ExecuteModule(ctx context.Context, module *core.Workflow, par return nil, fmt.Errorf("workflow is not a module") } - // Create cancellable context for run registry support + // Create cancellable context for run control plane support ctx, cancel := context.WithCancel(ctx) defer cancel() - // Register with run registry if we have a database run UUID + // Register with run control plane if we have a database run UUID // This enables API-based cancellation of the run if e.dbRunUUID != "" { - activeRun := GetRunRegistry().Register(e.dbRunUUID, cancel) - defer GetRunRegistry().Unregister(e.dbRunUUID) - e.logger.Debug("Registered run with registry", + activeRun := GetRunControlPlane().Register(e.dbRunUUID, cancel) + defer GetRunControlPlane().Unregister(e.dbRunUUID) + e.logger.Debug("Registered run with control plane", zap.String("run_uuid", e.dbRunUUID), zap.Time("started_at", activeRun.StartedAt), ) @@ -787,12 +787,12 @@ func (e *Executor) ExecuteModule(ctx context.Context, module *core.Workflow, par runUUID := e.dbRunUUID r.SetPIDCallbacks( func(pid int) { - GetRunRegistry().AddPID(runUUID, pid) + GetRunControlPlane().AddPID(runUUID, pid) // Update database with current PID for API visibility _ = database.UpdateRunPID(ctx, runUUID, pid) }, func(pid int) { - GetRunRegistry().RemovePID(runUUID, pid) + GetRunControlPlane().RemovePID(runUUID, pid) // Clear PID from database when process ends _ = database.ClearRunPID(ctx, runUUID) }, @@ -1385,12 +1385,12 @@ func (e *Executor) ExecuteFlow(ctx context.Context, flow *core.Workflow, params ctx, cancel := context.WithCancel(ctx) defer cancel() - // Register with run registry if we have a database run UUID + // Register with run control plane if we have a database run UUID // This enables API-based cancellation of the run if e.dbRunUUID != "" { - activeRun := GetRunRegistry().Register(e.dbRunUUID, cancel) - defer GetRunRegistry().Unregister(e.dbRunUUID) - e.logger.Debug("Registered flow with registry", + activeRun := GetRunControlPlane().Register(e.dbRunUUID, cancel) + defer GetRunControlPlane().Unregister(e.dbRunUUID) + e.logger.Debug("Registered flow with control plane", zap.String("run_uuid", e.dbRunUUID), zap.Time("started_at", activeRun.StartedAt), ) diff --git a/internal/executor/foreach_executor.go b/internal/executor/foreach_executor.go index 3c7f4c1..c2ac8ff 100644 --- a/internal/executor/foreach_executor.go +++ b/internal/executor/foreach_executor.go @@ -10,20 +10,25 @@ import ( "time" "github.com/j3ssie/osmedeus/v5/internal/core" + "github.com/j3ssie/osmedeus/v5/internal/functions" + "github.com/j3ssie/osmedeus/v5/internal/logger" "github.com/j3ssie/osmedeus/v5/internal/template" + "go.uber.org/zap" ) // ForeachExecutor executes foreach steps type ForeachExecutor struct { - dispatcher *StepDispatcher - templateEngine template.TemplateEngine + dispatcher *StepDispatcher + templateEngine template.TemplateEngine + functionRegistry *functions.Registry } // NewForeachExecutor creates a new foreach executor -func NewForeachExecutor(dispatcher *StepDispatcher, engine template.TemplateEngine) *ForeachExecutor { +func NewForeachExecutor(dispatcher *StepDispatcher, engine template.TemplateEngine, registry *functions.Registry) *ForeachExecutor { return &ForeachExecutor{ - dispatcher: dispatcher, - templateEngine: engine, + dispatcher: dispatcher, + templateEngine: engine, + functionRegistry: registry, } } @@ -240,8 +245,23 @@ func (e *ForeachExecutor) executeWithWorkerPool(ctx context.Context, step *core. continue } + // Apply variable pre-processing if configured + loopValue := work.value + if step.VariablePreProcess != "" { + processedValue, err := e.preProcessVariable(step.VariablePreProcess, step.Variable, work.value, execCtx) + if err != nil { + // Log warning but continue with original value (fail-safe) + logger.Get().Warn("variable pre-process failed, using original value", + zap.String("expression", step.VariablePreProcess), + zap.String("original_value", work.value), + zap.Error(err)) + } else { + loopValue = processedValue + } + } + // Create optimized child context with loop variables pre-set - childCtx := execCtx.CloneForLoop(step.Variable, work.value, work.index+1) + childCtx := execCtx.CloneForLoop(step.Variable, loopValue, work.index+1) // Clone inner step and render secondary templates [[ ]] innerStep := e.renderSecondaryTemplates(step.Step, childCtx) @@ -317,3 +337,39 @@ func (e *ForeachExecutor) executeWithWorkerPool(ctx context.Context, step *core. func (e *ForeachExecutor) CanHandle(stepType core.StepType) bool { return stepType == core.StepTypeForeach } + +// autoQuoteForJS wraps string values in single quotes for JS function calls, +// escaping any internal single quotes. +func autoQuoteForJS(value string) string { + // Escape single quotes: ' -> \' + escaped := strings.ReplaceAll(value, "'", "\\'") + return "'" + escaped + "'" +} + +// preProcessVariable evaluates the pre-process expression with the loop variable set. +// The expression can use [[variable]] syntax to reference the current value. +// For example: "get_parent_url([[url]])" with url="http://example.com/path" +// becomes: "get_parent_url('http://example.com/path')" +func (e *ForeachExecutor) preProcessVariable(expr, varName, varValue string, execCtx *core.ExecutionContext) (string, error) { + // Build context with parent variables + ctx := execCtx.GetVariables() + + // For pre-process expressions, auto-quote the loop variable value + // This allows clean syntax: get_parent_url([[url]]) instead of get_parent_url('[[url]]') + ctx[varName] = autoQuoteForJS(varValue) + + // Render secondary templates [[var]] -> 'quoted_value' + renderedExpr, err := e.templateEngine.RenderSecondary(expr, ctx) + if err != nil { + return "", fmt.Errorf("failed to render pre-process expression: %w", err) + } + + // Execute the function expression + result, err := e.functionRegistry.Execute(renderedExpr, ctx) + if err != nil { + return "", fmt.Errorf("failed to execute pre-process function: %w", err) + } + + // Convert result to string + return fmt.Sprintf("%v", result), nil +} diff --git a/internal/executor/foreach_executor_test.go b/internal/executor/foreach_executor_test.go new file mode 100644 index 0000000..b89c5b5 --- /dev/null +++ b/internal/executor/foreach_executor_test.go @@ -0,0 +1,344 @@ +package executor + +import ( + "context" + "os" + "path/filepath" + "testing" + + "github.com/j3ssie/osmedeus/v5/internal/core" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestForeachExecutor_VariablePreProcess(t *testing.T) { + // Create temp directory for test files + tmpDir := t.TempDir() + + // Create test input file with URLs + inputFile := filepath.Join(tmpDir, "urls.txt") + err := os.WriteFile(inputFile, []byte(`https://example.com/path/file.php?id=1 +https://example.com/api/v1/users +https://example.com/deep/nested/path/resource +`), 0644) + require.NoError(t, err) + + // Create output file path + outputFile := filepath.Join(tmpDir, "output.txt") + + // Create dispatcher + dispatcher := NewStepDispatcher() + + // Create execution context + execCtx := core.NewExecutionContext("test-workflow", core.KindModule, "test-run-uuid", "test-target") + execCtx.SetVariable("Output", tmpDir) + + // Test case 1: Basic pre-processing with get_parent_url + t.Run("get_parent_url pre-processing", func(t *testing.T) { + step := &core.Step{ + Name: "test-preprocess", + Type: core.StepTypeForeach, + Input: inputFile, + Variable: "url", + VariablePreProcess: "get_parent_url([[url]])", + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[url]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Verify output contains parent URLs + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + assert.Contains(t, outputStr, "https://example.com/path/") + assert.Contains(t, outputStr, "https://example.com/api/v1/") + assert.Contains(t, outputStr, "https://example.com/deep/nested/path/") + }) + + // Test case 2: No pre-processing (existing behavior unchanged) + t.Run("no pre-processing", func(t *testing.T) { + step := &core.Step{ + Name: "test-no-preprocess", + Type: core.StepTypeForeach, + Input: inputFile, + Variable: "url", + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[url]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Verify output contains original URLs + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + assert.Contains(t, outputStr, "https://example.com/path/file.php?id=1") + assert.Contains(t, outputStr, "https://example.com/api/v1/users") + }) + + // Test case 3: Pre-processing with trim + t.Run("trim pre-processing", func(t *testing.T) { + // Create input with whitespace + trimInputFile := filepath.Join(tmpDir, "trim-input.txt") + err := os.WriteFile(trimInputFile, []byte(` hello + world +`), 0644) + require.NoError(t, err) + + step := &core.Step{ + Name: "test-trim-preprocess", + Type: core.StepTypeForeach, + Input: trimInputFile, + Variable: "line", + VariablePreProcess: "trim([[line]])", + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[line]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Verify output contains trimmed values + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + // The file iterator already trims lines, but the pre-process should work anyway + assert.Contains(t, outputStr, "hello") + assert.Contains(t, outputStr, "world") + }) + + // Test case 4: Pre-processing with parse_url + t.Run("parse_url pre-processing", func(t *testing.T) { + step := &core.Step{ + Name: "test-parse-url-preprocess", + Type: core.StepTypeForeach, + Input: inputFile, + Variable: "url", + VariablePreProcess: "parse_url([[url]], '%d')", // Extract domain only + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[url]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Verify output contains just domains + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + assert.Contains(t, outputStr, "example.com") + }) + + // Test case 5: Pre-processing with chained functions + t.Run("chained functions pre-processing", func(t *testing.T) { + step := &core.Step{ + Name: "test-chained-preprocess", + Type: core.StepTypeForeach, + Input: inputFile, + Variable: "url", + VariablePreProcess: "to_lower_case(parse_url([[url]], '%d'))", + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[url]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Verify output + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + assert.Contains(t, outputStr, "example.com") + }) +} + +func TestForeachExecutor_VariablePreProcess_ErrorHandling(t *testing.T) { + // Create temp directory for test files + tmpDir := t.TempDir() + + // Create test input file + inputFile := filepath.Join(tmpDir, "input.txt") + err := os.WriteFile(inputFile, []byte(`line1 +line2 +`), 0644) + require.NoError(t, err) + + outputFile := filepath.Join(tmpDir, "output.txt") + + // Create dispatcher + dispatcher := NewStepDispatcher() + + // Create execution context + execCtx := core.NewExecutionContext("test-workflow", core.KindModule, "test-run-uuid", "test-target") + + // Test case: Invalid function should fall back to original value + t.Run("invalid function fallback", func(t *testing.T) { + step := &core.Step{ + Name: "test-invalid-func", + Type: core.StepTypeForeach, + Input: inputFile, + Variable: "line", + VariablePreProcess: "invalid_nonexistent_function([[line]])", + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[line]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Should still produce output with original values since invalid function falls back + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + assert.Contains(t, outputStr, "line1") + assert.Contains(t, outputStr, "line2") + }) +} + +func TestForeachExecutor_VariablePreProcess_SpecialChars(t *testing.T) { + // Create temp directory for test files + tmpDir := t.TempDir() + + // Create test input file with URLs containing special characters + inputFile := filepath.Join(tmpDir, "special-urls.txt") + err := os.WriteFile(inputFile, []byte(`https://example.com/path?foo='bar'&baz=1 +https://example.com/it's/a/test +`), 0644) + require.NoError(t, err) + + outputFile := filepath.Join(tmpDir, "output.txt") + + // Create dispatcher + dispatcher := NewStepDispatcher() + + // Create execution context + execCtx := core.NewExecutionContext("test-workflow", core.KindModule, "test-run-uuid", "test-target") + + t.Run("special characters in URL", func(t *testing.T) { + step := &core.Step{ + Name: "test-special-chars", + Type: core.StepTypeForeach, + Input: inputFile, + Variable: "url", + VariablePreProcess: "get_parent_url([[url]])", + Threads: "1", + Step: &core.Step{ + Type: core.StepTypeBash, + Command: "echo '[[url]]' >> " + outputFile, + }, + } + + // Clean output file + _ = os.Remove(outputFile) + + ctx := context.Background() + result, err := dispatcher.Dispatch(ctx, step, execCtx) + require.NoError(t, err) + assert.Equal(t, core.StepStatusSuccess, result.Status) + + // Verify output contains parent URLs + output, err := os.ReadFile(outputFile) + require.NoError(t, err) + outputStr := string(output) + + assert.Contains(t, outputStr, "https://example.com/") + }) +} + +func TestAutoQuoteForJS(t *testing.T) { + tests := []struct { + name string + input string + expected string + }{ + { + name: "simple string", + input: "hello", + expected: "'hello'", + }, + { + name: "string with single quote", + input: "it's a test", + expected: "'it\\'s a test'", + }, + { + name: "URL", + input: "https://example.com/path?foo=bar", + expected: "'https://example.com/path?foo=bar'", + }, + { + name: "URL with quotes", + input: "https://example.com/path?foo='bar'", + expected: "'https://example.com/path?foo=\\'bar\\''", + }, + { + name: "empty string", + input: "", + expected: "''", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := autoQuoteForJS(tt.input) + assert.Equal(t, tt.expected, result) + }) + } +} diff --git a/internal/executor/run_registry.go b/internal/executor/run_control_plane.go similarity index 73% rename from internal/executor/run_registry.go rename to internal/executor/run_control_plane.go index f48759b..6800759 100644 --- a/internal/executor/run_registry.go +++ b/internal/executor/run_control_plane.go @@ -57,27 +57,27 @@ func (a *ActiveRun) KillAllPIDs() []int { return killed } -// RunRegistry tracks active workflow runs for cancellation support -type RunRegistry struct { +// RunControlPlane tracks active workflow runs for cancellation support +type RunControlPlane struct { mu sync.RWMutex runs map[string]*ActiveRun } -var globalRegistry *RunRegistry -var registryOnce sync.Once +var globalControlPlane *RunControlPlane +var controlPlaneOnce sync.Once -// GetRunRegistry returns the singleton run registry -func GetRunRegistry() *RunRegistry { - registryOnce.Do(func() { - globalRegistry = &RunRegistry{ +// GetRunControlPlane returns the singleton run control plane +func GetRunControlPlane() *RunControlPlane { + controlPlaneOnce.Do(func() { + globalControlPlane = &RunControlPlane{ runs: make(map[string]*ActiveRun), } }) - return globalRegistry + return globalControlPlane } -// Register adds a new run to the registry -func (r *RunRegistry) Register(runUUID string, cancel context.CancelFunc) *ActiveRun { +// Register adds a new run to the control plane +func (r *RunControlPlane) Register(runUUID string, cancel context.CancelFunc) *ActiveRun { r.mu.Lock() defer r.mu.Unlock() @@ -91,15 +91,15 @@ func (r *RunRegistry) Register(runUUID string, cancel context.CancelFunc) *Activ return activeRun } -// Unregister removes a run from the registry -func (r *RunRegistry) Unregister(runUUID string) { +// Unregister removes a run from the control plane +func (r *RunControlPlane) Unregister(runUUID string) { r.mu.Lock() defer r.mu.Unlock() delete(r.runs, runUUID) } // Get retrieves an active run by its UUID -func (r *RunRegistry) Get(runUUID string) *ActiveRun { +func (r *RunControlPlane) Get(runUUID string) *ActiveRun { r.mu.RLock() defer r.mu.RUnlock() return r.runs[runUUID] @@ -107,13 +107,13 @@ func (r *RunRegistry) Get(runUUID string) *ActiveRun { // Cancel cancels a run by calling its cancel function and killing all tracked PIDs. // Returns the list of killed PIDs and any error. -func (r *RunRegistry) Cancel(runUUID string) ([]int, error) { +func (r *RunControlPlane) Cancel(runUUID string) ([]int, error) { r.mu.Lock() activeRun, exists := r.runs[runUUID] r.mu.Unlock() if !exists { - return nil, fmt.Errorf("run %s not found in registry", runUUID) + return nil, fmt.Errorf("run %s not found in control plane", runUUID) } // First, cancel the context to stop any new operations @@ -128,7 +128,7 @@ func (r *RunRegistry) Cancel(runUUID string) ([]int, error) { } // AddPID adds a PID to a run's tracked processes -func (r *RunRegistry) AddPID(runUUID string, pid int) { +func (r *RunControlPlane) AddPID(runUUID string, pid int) { r.mu.RLock() activeRun := r.runs[runUUID] r.mu.RUnlock() @@ -139,7 +139,7 @@ func (r *RunRegistry) AddPID(runUUID string, pid int) { } // RemovePID removes a PID from a run's tracked processes -func (r *RunRegistry) RemovePID(runUUID string, pid int) { +func (r *RunControlPlane) RemovePID(runUUID string, pid int) { r.mu.RLock() activeRun := r.runs[runUUID] r.mu.RUnlock() @@ -150,7 +150,7 @@ func (r *RunRegistry) RemovePID(runUUID string, pid int) { } // ListActive returns a list of all active run UUIDs -func (r *RunRegistry) ListActive() []string { +func (r *RunControlPlane) ListActive() []string { r.mu.RLock() defer r.mu.RUnlock() @@ -162,7 +162,7 @@ func (r *RunRegistry) ListActive() []string { } // Count returns the number of active runs -func (r *RunRegistry) Count() int { +func (r *RunControlPlane) Count() int { r.mu.RLock() defer r.mu.RUnlock() return len(r.runs) diff --git a/internal/executor/run_registry_test.go b/internal/executor/run_control_plane_test.go similarity index 61% rename from internal/executor/run_registry_test.go rename to internal/executor/run_control_plane_test.go index ba24bfd..31117ac 100644 --- a/internal/executor/run_registry_test.go +++ b/internal/executor/run_control_plane_test.go @@ -7,8 +7,8 @@ import ( "time" ) -func TestRunRegistryRegisterAndGet(t *testing.T) { - registry := &RunRegistry{ +func TestRunControlPlaneRegisterAndGet(t *testing.T) { + controlPlane := &RunControlPlane{ runs: make(map[string]*ActiveRun), } @@ -16,7 +16,7 @@ func TestRunRegistryRegisterAndGet(t *testing.T) { defer cancel() runUUID := "test-run-123" - activeRun := registry.Register(runUUID, cancel) + activeRun := controlPlane.Register(runUUID, cancel) if activeRun == nil { t.Fatal("Register returned nil") @@ -27,20 +27,20 @@ func TestRunRegistryRegisterAndGet(t *testing.T) { } // Test Get - retrieved := registry.Get(runUUID) + retrieved := controlPlane.Get(runUUID) if retrieved != activeRun { t.Error("Get returned different instance") } // Test Get for non-existent run - nonExistent := registry.Get("non-existent") + nonExistent := controlPlane.Get("non-existent") if nonExistent != nil { t.Error("Get should return nil for non-existent run") } } -func TestRunRegistryUnregister(t *testing.T) { - registry := &RunRegistry{ +func TestRunControlPlaneUnregister(t *testing.T) { + controlPlane := &RunControlPlane{ runs: make(map[string]*ActiveRun), } @@ -48,33 +48,33 @@ func TestRunRegistryUnregister(t *testing.T) { defer cancel() runUUID := "test-run-456" - registry.Register(runUUID, cancel) + controlPlane.Register(runUUID, cancel) // Verify registered - if registry.Get(runUUID) == nil { + if controlPlane.Get(runUUID) == nil { t.Fatal("Run should be registered") } // Unregister - registry.Unregister(runUUID) + controlPlane.Unregister(runUUID) // Verify unregistered - if registry.Get(runUUID) != nil { + if controlPlane.Get(runUUID) != nil { t.Error("Run should be unregistered") } } -func TestRunRegistryCancel(t *testing.T) { - registry := &RunRegistry{ +func TestRunControlPlaneCancel(t *testing.T) { + controlPlane := &RunControlPlane{ runs: make(map[string]*ActiveRun), } ctx, cancel := context.WithCancel(context.Background()) runUUID := "test-run-789" - registry.Register(runUUID, cancel) + controlPlane.Register(runUUID, cancel) // Cancel should call the cancel function - _, err := registry.Cancel(runUUID) + _, err := controlPlane.Cancel(runUUID) if err != nil { t.Errorf("Cancel returned error: %v", err) } @@ -88,14 +88,14 @@ func TestRunRegistryCancel(t *testing.T) { } // Cancel non-existent run - _, err = registry.Cancel("non-existent") + _, err = controlPlane.Cancel("non-existent") if err == nil { t.Error("Cancel should return error for non-existent run") } } -func TestRunRegistryPIDTracking(t *testing.T) { - registry := &RunRegistry{ +func TestRunControlPlanePIDTracking(t *testing.T) { + controlPlane := &RunControlPlane{ runs: make(map[string]*ActiveRun), } @@ -103,14 +103,14 @@ func TestRunRegistryPIDTracking(t *testing.T) { defer cancel() runUUID := "test-run-pid" - registry.Register(runUUID, cancel) + controlPlane.Register(runUUID, cancel) // Add PIDs - registry.AddPID(runUUID, 1234) - registry.AddPID(runUUID, 5678) + controlPlane.AddPID(runUUID, 1234) + controlPlane.AddPID(runUUID, 5678) // Verify PIDs are tracked - activeRun := registry.Get(runUUID) + activeRun := controlPlane.Get(runUUID) count := 0 activeRun.PIDs.Range(func(_, _ any) bool { count++ @@ -121,7 +121,7 @@ func TestRunRegistryPIDTracking(t *testing.T) { } // Remove a PID - registry.RemovePID(runUUID, 1234) + controlPlane.RemovePID(runUUID, 1234) // Verify PID removed count = 0 @@ -134,8 +134,8 @@ func TestRunRegistryPIDTracking(t *testing.T) { } } -func TestRunRegistryKillAllPIDs(t *testing.T) { - // Create a fresh ActiveRun (not using the registry to avoid syscall issues) +func TestRunControlPlaneKillAllPIDs(t *testing.T) { + // Create a fresh ActiveRun (not using the control plane to avoid syscall issues) activeRun := &ActiveRun{ RunUUID: "test-kill", PIDs: &sync.Map{}, @@ -165,8 +165,8 @@ func TestRunRegistryKillAllPIDs(t *testing.T) { } } -func TestRunRegistryConcurrency(t *testing.T) { - registry := &RunRegistry{ +func TestRunControlPlaneConcurrency(t *testing.T) { + controlPlane := &RunControlPlane{ runs: make(map[string]*ActiveRun), } @@ -181,51 +181,51 @@ func TestRunRegistryConcurrency(t *testing.T) { defer cancel() runUUID := "concurrent-run-" + string(rune('a'+i%26)) - registry.Register(runUUID, cancel) - registry.AddPID(runUUID, i+1000) - registry.RemovePID(runUUID, i+1000) - registry.Get(runUUID) - registry.Unregister(runUUID) + controlPlane.Register(runUUID, cancel) + controlPlane.AddPID(runUUID, i+1000) + controlPlane.RemovePID(runUUID, i+1000) + controlPlane.Get(runUUID) + controlPlane.Unregister(runUUID) }(i) } wg.Wait() // Should have no leftover runs - if registry.Count() != 0 { - t.Errorf("Expected 0 runs, got %d", registry.Count()) + if controlPlane.Count() != 0 { + t.Errorf("Expected 0 runs, got %d", controlPlane.Count()) } } -func TestRunRegistryListActive(t *testing.T) { - registry := &RunRegistry{ +func TestRunControlPlaneListActive(t *testing.T) { + controlPlane := &RunControlPlane{ runs: make(map[string]*ActiveRun), } _, cancel := context.WithCancel(context.Background()) defer cancel() - registry.Register("run-1", cancel) - registry.Register("run-2", cancel) - registry.Register("run-3", cancel) + controlPlane.Register("run-1", cancel) + controlPlane.Register("run-2", cancel) + controlPlane.Register("run-3", cancel) - active := registry.ListActive() + active := controlPlane.ListActive() if len(active) != 3 { t.Errorf("Expected 3 active runs, got %d", len(active)) } - if registry.Count() != 3 { - t.Errorf("Expected count 3, got %d", registry.Count()) + if controlPlane.Count() != 3 { + t.Errorf("Expected count 3, got %d", controlPlane.Count()) } } -func TestGetRunRegistrySingleton(t *testing.T) { +func TestGetRunControlPlaneSingleton(t *testing.T) { // Get the singleton twice - reg1 := GetRunRegistry() - reg2 := GetRunRegistry() + cp1 := GetRunControlPlane() + cp2 := GetRunControlPlane() // Should be the same instance - if reg1 != reg2 { - t.Error("GetRunRegistry should return the same singleton instance") + if cp1 != cp2 { + t.Error("GetRunControlPlane should return the same singleton instance") } } diff --git a/internal/functions/constants.go b/internal/functions/constants.go index 07424e2..296926a 100644 --- a/internal/functions/constants.go +++ b/internal/functions/constants.go @@ -215,9 +215,13 @@ const ( FnJSONLFilter = "jsonl_filter" ) -// URL Processing Functions - URL deduplication and filtering +// URL Processing Functions - URL deduplication, filtering, and parsing const ( FnInterestingUrls = "interesting_urls" // interesting_urls(src, dest, json_field?) -> bool + FnGetParentURL = "get_parent_url" // get_parent_url(url) -> string (strips last path component) + FnParseURL = "parse_url" // parse_url(url, format) -> string (format directives like unfurl) + FnQueryReplace = "query_replace" // query_replace(url, value, mode?) -> string (replace all query param values) + FnPathReplace = "path_replace" // path_replace(url, value, position?) -> string (replace path segment at position) ) // Markdown Functions - Markdown rendering and conversion @@ -450,6 +454,10 @@ func AllFunctions() []string { // URL Processing Functions FnInterestingUrls, + FnGetParentURL, + FnParseURL, + FnQueryReplace, + FnPathReplace, // Markdown Functions FnRenderMarkdownFromFile, @@ -758,6 +766,10 @@ func FunctionRegistry() map[string][]FunctionInfo { }, CategoryURLProcessing: { {FnInterestingUrls, "interesting_urls(src, dest, json_field?)", "Deduplicate URLs by hostname+path+params, filter static files and noise patterns", "bool", "interesting_urls('{{Output}}/all-urls.txt', '{{Output}}/interesting-urls.txt', 'url')"}, + {FnGetParentURL, "get_parent_url(url)", "Strip last path component and return parent directory URL", "string", "get_parent_url('https://example.com/path/file.php')"}, + {FnParseURL, "parse_url(url, format)", "Format URL using directives: %s(scheme) %d(domain) %S(subdomain) %r(root) %t(tld) %P(port) %p(path) %e(ext) %q(query) %f(fragment) %a(authority)", "string", "parse_url('https://sub.example.com/path', '%S.%r')"}, + {FnQueryReplace, "query_replace(url, value, mode?)", "Replace all query param values; mode: 'replace' (default) or 'append'", "string", "query_replace('https://example.com?a=1&b=2', 'test')"}, + {FnPathReplace, "path_replace(url, value, position?)", "Replace path segment at position (1-indexed); 0 replaces all", "string", "path_replace('https://example.com/a/b/c', 'new', 2)"}, }, CategoryMarkdown: { {FnRenderMarkdownFromFile, "render_markdown_from_file(path)", "Render markdown with terminal styling", "string", "render_markdown_from_file('{{Output}}/report.md')"}, diff --git a/internal/functions/goja_runtime.go b/internal/functions/goja_runtime.go index c1874f2..3c35274 100644 --- a/internal/functions/goja_runtime.go +++ b/internal/functions/goja_runtime.go @@ -205,6 +205,10 @@ func (r *GojaRuntime) registerFunctionsOnVM(vm *goja.Runtime) { // URL processing functions _ = vm.Set(FnInterestingUrls, vf.interestingUrls) + _ = vm.Set(FnGetParentURL, vf.getParentURL) + _ = vm.Set(FnParseURL, vf.parseURL) + _ = vm.Set(FnQueryReplace, vf.queryReplace) + _ = vm.Set(FnPathReplace, vf.pathReplace) // Markdown functions _ = vm.Set(FnRenderMarkdownFromFile, vf.renderMarkdownFromFile) diff --git a/internal/functions/url_functions.go b/internal/functions/url_functions.go index ecc39f8..c58f2c0 100644 --- a/internal/functions/url_functions.go +++ b/internal/functions/url_functions.go @@ -10,6 +10,7 @@ import ( "path/filepath" "regexp" "sort" + "strconv" "strings" "github.com/dop251/goja" @@ -325,3 +326,449 @@ func resolveToIP(hostname string) string { return "" } + +// getParentURL strips the last path component from a URL and returns the parent directory. +// Usage: get_parent_url(url) -> string +// - url: the URL to process +// +// Examples: +// +// get_parent_url("https://example.com/j3ssie/sample.php?query=123") -> "https://example.com/j3ssie/" +// get_parent_url("https://example.com/a/b/c/") -> "https://example.com/a/b/" +// get_parent_url("https://example.com/file.txt") -> "https://example.com/" +// get_parent_url("https://example.com/") -> "https://example.com/" +func (vf *vmFunc) getParentURL(call goja.FunctionCall) goja.Value { + input := call.Argument(0).String() + logger.Get().Debug("Calling "+terminal.HiGreen("get_parent_url"), + zap.String("input", input)) + + if input == "undefined" || input == "" { + logger.Get().Warn("get_parent_url: empty input provided") + return vf.vm.ToValue("") + } + + result := getParentURLImpl(input) + + logger.Get().Debug(terminal.HiGreen("get_parent_url")+" result", + zap.String("input", input), + zap.String("result", result)) + + return vf.vm.ToValue(result) +} + +// getParentURLImpl extracts the parent directory URL. +// This is the implementation that can be tested independently. +func getParentURLImpl(input string) string { + // Parse the URL + u, err := url.Parse(input) + if err != nil { + return input + } + + // Get the path and strip query/fragment + path := u.Path + + // Handle empty path + if path == "" || path == "/" { + // Ensure trailing slash + u.Path = "/" + u.RawQuery = "" + u.Fragment = "" + return u.String() + } + + // Remove trailing slash for uniform processing + path = strings.TrimSuffix(path, "/") + + // Find the last slash + lastSlash := strings.LastIndex(path, "/") + if lastSlash == -1 { + // No slash found, return root + u.Path = "/" + } else { + // Keep everything up to and including the last slash + u.Path = path[:lastSlash+1] + } + + // Clear query and fragment + u.RawQuery = "" + u.Fragment = "" + + return u.String() +} + +// parseURL formats a URL using format directives similar to unfurl. +// Usage: parse_url(url, format) -> string +// - url: the URL to parse +// - format: format string with directives +// +// Format directives: +// +// %% - Literal percent character +// %s - Request scheme (http, https) +// %u - User info (username:password) +// %d - Full domain (sub.example.com) +// %S - Subdomain (sub) +// %r - Root domain (example) +// %t - TLD (com) +// %P - Port (8080) +// %p - Path (/users/list) +// %e - File extension (jpg) +// %q - Raw query string (a=1&b=2) +// %f - Fragment (section) +// %@ - @ if user info exists, empty otherwise +// %: - : if port exists, empty otherwise +// %? - ? if query exists, empty otherwise +// %# - # if fragment exists, empty otherwise +// %a - Authority (user:pass@domain:port) +func (vf *vmFunc) parseURL(call goja.FunctionCall) goja.Value { + urlStr := call.Argument(0).String() + format := call.Argument(1).String() + logger.Get().Debug("Calling "+terminal.HiGreen("parse_url"), + zap.String("url", urlStr), zap.String("format", format)) + + if urlStr == "undefined" || urlStr == "" { + logger.Get().Warn("parse_url: empty URL provided") + return vf.vm.ToValue("") + } + + if format == "undefined" || format == "" { + logger.Get().Warn("parse_url: empty format provided") + return vf.vm.ToValue("") + } + + result := parseURLImpl(urlStr, format) + + logger.Get().Debug(terminal.HiGreen("parse_url")+" result", + zap.String("url", urlStr), + zap.String("format", format), + zap.String("result", result)) + + return vf.vm.ToValue(result) +} + +// parseURLImpl formats a URL using format directives. +// This is the implementation that can be tested independently. +func parseURLImpl(urlStr, format string) string { + u, err := url.Parse(urlStr) + if err != nil { + return "" + } + + // Extract domain parts + domain := u.Hostname() + subdomain, root, tld := extractDomainParts(domain) + + // Extract file extension from path + ext := "" + if u.Path != "" { + base := filepath.Base(u.Path) + if dotIdx := strings.LastIndex(base, "."); dotIdx != -1 && dotIdx < len(base)-1 { + ext = base[dotIdx+1:] + } + } + + // Extract user info + userInfo := "" + if u.User != nil { + userInfo = u.User.String() + } + + // Extract port + port := u.Port() + + // Build result by parsing format string + var result strings.Builder + i := 0 + for i < len(format) { + if format[i] == '%' && i+1 < len(format) { + switch format[i+1] { + case '%': + result.WriteByte('%') + case 's': + result.WriteString(u.Scheme) + case 'u': + result.WriteString(userInfo) + case 'd': + result.WriteString(domain) + case 'S': + result.WriteString(subdomain) + case 'r': + result.WriteString(root) + case 't': + result.WriteString(tld) + case 'P': + result.WriteString(port) + case 'p': + result.WriteString(u.Path) + case 'e': + result.WriteString(ext) + case 'q': + result.WriteString(u.RawQuery) + case 'f': + result.WriteString(u.Fragment) + case '@': + if userInfo != "" { + result.WriteByte('@') + } + case ':': + if port != "" { + result.WriteByte(':') + } + case '?': + if u.RawQuery != "" { + result.WriteByte('?') + } + case '#': + if u.Fragment != "" { + result.WriteByte('#') + } + case 'a': + // Authority: user:pass@domain:port + if userInfo != "" { + result.WriteString(userInfo) + result.WriteByte('@') + } + result.WriteString(domain) + if port != "" { + result.WriteByte(':') + result.WriteString(port) + } + default: + // Unknown directive, output as-is + result.WriteByte('%') + result.WriteByte(format[i+1]) + } + i += 2 + } else { + result.WriteByte(format[i]) + i++ + } + } + + return result.String() +} + +// knownMultiPartTLDs contains common multi-part TLDs +var knownMultiPartTLDs = map[string]bool{ + "co.uk": true, + "com.au": true, + "co.jp": true, + "co.nz": true, + "co.za": true, + "com.br": true, + "com.cn": true, + "com.mx": true, + "com.tw": true, + "com.hk": true, + "com.sg": true, + "org.uk": true, + "net.au": true, + "gov.uk": true, + "ac.uk": true, + "edu.au": true, + "co.in": true, + "com.ar": true, + "com.co": true, + "co.kr": true, + "or.jp": true, + "ne.jp": true, + "ac.jp": true, + "go.jp": true, +} + +// extractDomainParts splits a domain into subdomain, root, and TLD. +// Examples: +// +// sub.example.com -> (sub, example, com) +// example.com -> ("", example, com) +// api.sub.example.co.uk -> (api.sub, example, co.uk) +func extractDomainParts(domain string) (subdomain, root, tld string) { + if domain == "" { + return "", "", "" + } + + parts := strings.Split(domain, ".") + if len(parts) == 1 { + // No dots, treat as root + return "", parts[0], "" + } + + // Check for multi-part TLDs + if len(parts) >= 2 { + potentialMultiTLD := parts[len(parts)-2] + "." + parts[len(parts)-1] + if knownMultiPartTLDs[potentialMultiTLD] { + tld = potentialMultiTLD + if len(parts) == 2 { + // e.g., "co.uk" - no root domain + return "", "", tld + } + root = parts[len(parts)-3] + if len(parts) > 3 { + subdomain = strings.Join(parts[:len(parts)-3], ".") + } + return subdomain, root, tld + } + } + + // Standard single-part TLD + tld = parts[len(parts)-1] + if len(parts) == 2 { + // e.g., "example.com" + return "", parts[0], tld + } + + // e.g., "sub.example.com" or "api.sub.example.com" + root = parts[len(parts)-2] + subdomain = strings.Join(parts[:len(parts)-2], ".") + return subdomain, root, tld +} + +// queryReplace replaces all query parameter values in a URL. +// Usage: query_replace(url, value, mode?) -> string +// - url: the URL to modify +// - value: the replacement value +// - mode: "replace" (default) or "append" +// +// Examples: +// +// query_replace("https://example.com?a=1&b=2", "new") -> "https://example.com?a=new&b=new" +// query_replace("https://example.com?a=1&b=2", "FUZZ", "append") -> "https://example.com?a=1FUZZ&b=2FUZZ" +func (vf *vmFunc) queryReplace(call goja.FunctionCall) goja.Value { + urlStr := call.Argument(0).String() + value := call.Argument(1).String() + mode := call.Argument(2).String() + + logger.Get().Debug("Calling "+terminal.HiGreen("query_replace"), + zap.String("url", urlStr), zap.String("value", value), zap.String("mode", mode)) + + if urlStr == "undefined" || urlStr == "" { + logger.Get().Warn("query_replace: empty URL provided") + return vf.vm.ToValue("") + } + if value == "undefined" { + value = "" + } + if mode == "undefined" || mode == "" { + mode = "replace" + } + + result := queryReplaceImpl(urlStr, value, mode) + + logger.Get().Debug(terminal.HiGreen("query_replace")+" result", + zap.String("url", urlStr), zap.String("value", value), + zap.String("mode", mode), zap.String("result", result)) + + return vf.vm.ToValue(result) +} + +// queryReplaceImpl is the implementation for testing. +func queryReplaceImpl(urlStr, value, mode string) string { + u, err := url.Parse(urlStr) + if err != nil { + return urlStr + } + + q := u.Query() + if len(q) == 0 { + return urlStr + } + + newQuery := url.Values{} + for key, values := range q { + for _, v := range values { + switch mode { + case "append": + newQuery.Add(key, v+value) + default: // "replace" + newQuery.Add(key, value) + } + } + } + u.RawQuery = newQuery.Encode() + return u.String() +} + +// pathReplace replaces a path segment at a specific position. +// Usage: path_replace(url, value, position?) -> string +// - url: the URL to modify +// - value: the replacement value +// - position: 1-indexed position (default 1), 0 or negative replaces all segments +// +// Examples: +// +// path_replace("https://example.com/a/b/c", "new") -> "https://example.com/new/b/c" +// path_replace("https://example.com/a/b/c", "new", 2) -> "https://example.com/a/new/c" +// path_replace("https://example.com/a/b/c", "new", 0) -> "https://example.com/new/new/new" +func (vf *vmFunc) pathReplace(call goja.FunctionCall) goja.Value { + urlStr := call.Argument(0).String() + value := call.Argument(1).String() + posArg := call.Argument(2) + + logger.Get().Debug("Calling "+terminal.HiGreen("path_replace"), + zap.String("url", urlStr), zap.String("value", value)) + + if urlStr == "undefined" || urlStr == "" { + logger.Get().Warn("path_replace: empty URL provided") + return vf.vm.ToValue("") + } + if value == "undefined" { + value = "" + } + + // Parse position (default 1) + position := 1 + if !goja.IsUndefined(posArg) && !goja.IsNull(posArg) { + if p, ok := posArg.Export().(int64); ok { + position = int(p) + } else if p, ok := posArg.Export().(float64); ok { + position = int(p) + } else if s := posArg.String(); s != "undefined" && s != "" { + if p, err := strconv.Atoi(s); err == nil { + position = p + } + } + } + + result := pathReplaceImpl(urlStr, value, position) + + logger.Get().Debug(terminal.HiGreen("path_replace")+" result", + zap.String("url", urlStr), zap.String("value", value), + zap.Int("position", position), zap.String("result", result)) + + return vf.vm.ToValue(result) +} + +// pathReplaceImpl is the implementation for testing. +func pathReplaceImpl(urlStr, value string, position int) string { + u, err := url.Parse(urlStr) + if err != nil { + return urlStr + } + + // Split path into segments (skip empty segments from leading slash) + path := strings.TrimPrefix(u.Path, "/") + if path == "" { + return urlStr + } + + segments := strings.Split(path, "/") + if len(segments) == 0 { + return urlStr + } + + // Replace based on position + if position <= 0 { + // Replace all segments + for i := range segments { + segments[i] = value + } + } else if position <= len(segments) { + // Replace specific segment (1-indexed) + segments[position-1] = value + } + // If position > len(segments), return unchanged + + u.Path = "/" + strings.Join(segments, "/") + return u.String() +} diff --git a/internal/functions/url_functions_test.go b/internal/functions/url_functions_test.go index 29ab2fd..d3b7006 100644 --- a/internal/functions/url_functions_test.go +++ b/internal/functions/url_functions_test.go @@ -349,3 +349,515 @@ func TestResolveToIP(t *testing.T) { // We just check it returns something assert.NotEmpty(t, result, "localhost should resolve") } + +func TestGetParentURLImpl(t *testing.T) { + tests := []struct { + name string + input string + expected string + }{ + { + name: "URL with file and query", + input: "https://example.com/j3ssie/sample.php?query=123", + expected: "https://example.com/j3ssie/", + }, + { + name: "URL with trailing slash", + input: "https://example.com/a/b/c/", + expected: "https://example.com/a/b/", + }, + { + name: "URL with single path", + input: "https://example.com/file.txt", + expected: "https://example.com/", + }, + { + name: "URL with root path", + input: "https://example.com/", + expected: "https://example.com/", + }, + { + name: "URL without path", + input: "https://example.com", + expected: "https://example.com/", + }, + { + name: "URL with port and path", + input: "http://example.com:8080/path/to/file", + expected: "http://example.com:8080/path/to/", + }, + { + name: "URL with fragment", + input: "https://example.com/path/file.html#section", + expected: "https://example.com/path/", + }, + { + name: "URL with query and fragment", + input: "https://example.com/api/data?id=1#results", + expected: "https://example.com/api/", + }, + { + name: "deep nested path", + input: "https://example.com/a/b/c/d/e/f.js", + expected: "https://example.com/a/b/c/d/e/", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := getParentURLImpl(tt.input) + assert.Equal(t, tt.expected, result) + }) + } +} + +func TestGetParentURL_WithRegistry(t *testing.T) { + registry := NewRegistry() + + // Test basic URL + result, err := registry.Execute(`get_parent_url("https://example.com/path/file.txt")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "https://example.com/path/", result) + + // Test with query parameters + result, err = registry.Execute(`get_parent_url("https://example.com/api/endpoint?foo=bar")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "https://example.com/api/", result) + + // Test with empty input + result, err = registry.Execute(`get_parent_url("")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "", result) +} + +func TestParseURLImpl(t *testing.T) { + tests := []struct { + name string + url string + format string + expected string + }{ + // Basic directives + { + name: "scheme", + url: "https://example.com", + format: "%s", + expected: "https", + }, + { + name: "full domain", + url: "https://sub.example.com", + format: "%d", + expected: "sub.example.com", + }, + { + name: "subdomain", + url: "https://sub.example.com", + format: "%S", + expected: "sub", + }, + { + name: "root domain", + url: "https://sub.example.com", + format: "%r", + expected: "example", + }, + { + name: "tld", + url: "https://sub.example.com", + format: "%t", + expected: "com", + }, + { + name: "port", + url: "https://example.com:8080/path", + format: "%P", + expected: "8080", + }, + { + name: "path", + url: "https://example.com/path/file.jpg", + format: "%p", + expected: "/path/file.jpg", + }, + { + name: "extension", + url: "https://example.com/path/file.jpg", + format: "%e", + expected: "jpg", + }, + { + name: "query string", + url: "https://example.com?a=1&b=2", + format: "%q", + expected: "a=1&b=2", + }, + { + name: "fragment", + url: "https://example.com#section", + format: "%f", + expected: "section", + }, + { + name: "user info", + url: "https://user:pass@example.com", + format: "%u", + expected: "user:pass", + }, + + // Conditional directives + { + name: "at sign with user info", + url: "https://user:pass@example.com", + format: "%u%@%d", + expected: "user:pass@example.com", + }, + { + name: "at sign without user info", + url: "https://example.com", + format: "%u%@%d", + expected: "example.com", + }, + { + name: "colon with port", + url: "https://example.com:8080", + format: "%d%:%P", + expected: "example.com:8080", + }, + { + name: "colon without port", + url: "https://example.com", + format: "%d%:%P", + expected: "example.com", + }, + { + name: "question mark with query", + url: "https://example.com?q=1", + format: "%d%?%q", + expected: "example.com?q=1", + }, + { + name: "question mark without query", + url: "https://example.com", + format: "%d%?%q", + expected: "example.com", + }, + { + name: "hash with fragment", + url: "https://example.com#section", + format: "%d%#%f", + expected: "example.com#section", + }, + { + name: "hash without fragment", + url: "https://example.com", + format: "%d%#%f", + expected: "example.com", + }, + + // Authority directive + { + name: "authority full", + url: "https://user:pass@example.com:8080", + format: "%a", + expected: "user:pass@example.com:8080", + }, + { + name: "authority no port", + url: "https://user:pass@example.com", + format: "%a", + expected: "user:pass@example.com", + }, + { + name: "authority no user", + url: "https://example.com:8080", + format: "%a", + expected: "example.com:8080", + }, + { + name: "authority simple", + url: "https://example.com", + format: "%a", + expected: "example.com", + }, + + // Literal percent + { + name: "literal percent", + url: "https://example.com", + format: "100%%", + expected: "100%", + }, + + // Combined formats + { + name: "subdomain and root", + url: "https://api.sub.example.com", + format: "%S.%r.%t", + expected: "api.sub.example.com", + }, + { + name: "scheme and domain", + url: "https://example.com", + format: "%s://%d", + expected: "https://example.com", + }, + { + name: "full URL reconstruction", + url: "https://example.com:8080/path?q=1#sec", + format: "%s://%d%:%P%p%?%q%#%f", + expected: "https://example.com:8080/path?q=1#sec", + }, + + // Multi-part TLDs + { + name: "co.uk domain", + url: "https://api.example.co.uk", + format: "%S|%r|%t", + expected: "api|example|co.uk", + }, + { + name: "com.au domain", + url: "https://sub.site.com.au", + format: "%S|%r|%t", + expected: "sub|site|com.au", + }, + + // Edge cases + { + name: "no extension", + url: "https://example.com/path/file", + format: "%e", + expected: "", + }, + { + name: "no subdomain", + url: "https://example.com", + format: "%S", + expected: "", + }, + { + name: "empty path", + url: "https://example.com", + format: "%p", + expected: "", + }, + { + name: "unknown directive", + url: "https://example.com", + format: "%z", + expected: "%z", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := parseURLImpl(tt.url, tt.format) + assert.Equal(t, tt.expected, result) + }) + } +} + +func TestParseURL_WithRegistry(t *testing.T) { + registry := NewRegistry() + + // Test scheme extraction + result, err := registry.Execute(`parse_url("https://example.com", "%s")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "https", result) + + // Test domain extraction + result, err = registry.Execute(`parse_url("https://sub.example.com", "%d")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "sub.example.com", result) + + // Test extension extraction + result, err = registry.Execute(`parse_url("https://example.com/path/image.jpg", "%e")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "jpg", result) + + // Test with empty URL + result, err = registry.Execute(`parse_url("", "%s")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "", result) + + // Test with empty format + result, err = registry.Execute(`parse_url("https://example.com", "")`, map[string]interface{}{}) + require.NoError(t, err) + assert.Equal(t, "", result) +} + +func TestExtractDomainParts(t *testing.T) { + tests := []struct { + name string + domain string + expectedSubdomain string + expectedRoot string + expectedTLD string + }{ + { + name: "simple domain", + domain: "example.com", + expectedSubdomain: "", + expectedRoot: "example", + expectedTLD: "com", + }, + { + name: "subdomain", + domain: "sub.example.com", + expectedSubdomain: "sub", + expectedRoot: "example", + expectedTLD: "com", + }, + { + name: "multiple subdomains", + domain: "api.sub.example.com", + expectedSubdomain: "api.sub", + expectedRoot: "example", + expectedTLD: "com", + }, + { + name: "co.uk domain", + domain: "example.co.uk", + expectedSubdomain: "", + expectedRoot: "example", + expectedTLD: "co.uk", + }, + { + name: "co.uk with subdomain", + domain: "api.example.co.uk", + expectedSubdomain: "api", + expectedRoot: "example", + expectedTLD: "co.uk", + }, + { + name: "com.au domain", + domain: "site.com.au", + expectedSubdomain: "", + expectedRoot: "site", + expectedTLD: "com.au", + }, + { + name: "single part", + domain: "localhost", + expectedSubdomain: "", + expectedRoot: "localhost", + expectedTLD: "", + }, + { + name: "empty domain", + domain: "", + expectedSubdomain: "", + expectedRoot: "", + expectedTLD: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + subdomain, root, tld := extractDomainParts(tt.domain) + assert.Equal(t, tt.expectedSubdomain, subdomain, "subdomain mismatch") + assert.Equal(t, tt.expectedRoot, root, "root mismatch") + assert.Equal(t, tt.expectedTLD, tld, "tld mismatch") + }) + } +} + +func TestQueryReplaceImpl(t *testing.T) { + tests := []struct { + name string + url string + value string + mode string + expected string + }{ + {"single param replace", "https://example.com?id=123", "new", "replace", "https://example.com?id=new"}, + {"multiple params replace", "https://example.com/path?one=1&two=2", "newval", "replace", "https://example.com/path?one=newval&two=newval"}, + {"append mode", "https://example.com?a=1&b=2", "FUZZ", "append", "https://example.com?a=1FUZZ&b=2FUZZ"}, + {"no query params", "https://example.com/path", "new", "replace", "https://example.com/path"}, + {"empty value", "https://example.com?a=1", "", "replace", "https://example.com?a="}, + {"with fragment", "https://example.com?a=1#section", "new", "replace", "https://example.com?a=new#section"}, + {"default mode", "https://example.com?x=old", "test", "", "https://example.com?x=test"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := queryReplaceImpl(tt.url, tt.value, tt.mode) + assert.Equal(t, tt.expected, result) + }) + } +} + +func TestPathReplaceImpl(t *testing.T) { + tests := []struct { + name string + url string + value string + position int + expected string + }{ + {"replace first", "https://example.com/first/second/third", "new", 1, "https://example.com/new/second/third"}, + {"replace second", "https://example.com/first/second/third", "new", 2, "https://example.com/first/new/third"}, + {"replace third", "https://example.com/first/second/third", "new", 3, "https://example.com/first/second/new"}, + {"replace all", "https://example.com/a/b/c", "x", 0, "https://example.com/x/x/x"}, + {"replace all negative", "https://example.com/a/b/c", "x", -1, "https://example.com/x/x/x"}, + {"position out of range", "https://example.com/a/b", "new", 5, "https://example.com/a/b"}, + {"with query", "https://example.com/a/b?q=1", "new", 1, "https://example.com/new/b?q=1"}, + {"single segment", "https://example.com/only", "new", 1, "https://example.com/new"}, + {"no path", "https://example.com", "new", 1, "https://example.com"}, + {"root path only", "https://example.com/", "new", 1, "https://example.com/"}, + {"with fragment", "https://example.com/a/b#sec", "new", 2, "https://example.com/a/new#sec"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := pathReplaceImpl(tt.url, tt.value, tt.position) + assert.Equal(t, tt.expected, result) + }) + } +} + +func TestQueryReplace_WithRegistry(t *testing.T) { + registry := NewRegistry() + + // Test replace mode (default) + result, err := registry.Execute(`query_replace("https://example.com?a=1&b=2", "test")`, nil) + require.NoError(t, err) + resultStr := result.(string) + assert.Contains(t, resultStr, "a=test") + assert.Contains(t, resultStr, "b=test") + + // Test append mode + result, err = registry.Execute(`query_replace("https://example.com?x=old", "FUZZ", "append")`, nil) + require.NoError(t, err) + assert.Equal(t, "https://example.com?x=oldFUZZ", result) + + // Test with empty URL + result, err = registry.Execute(`query_replace("", "test")`, nil) + require.NoError(t, err) + assert.Equal(t, "", result) +} + +func TestPathReplace_WithRegistry(t *testing.T) { + registry := NewRegistry() + + // Test default position (1) + result, err := registry.Execute(`path_replace("https://example.com/a/b/c", "new")`, nil) + require.NoError(t, err) + assert.Equal(t, "https://example.com/new/b/c", result) + + // Test specific position + result, err = registry.Execute(`path_replace("https://example.com/a/b/c", "new", 2)`, nil) + require.NoError(t, err) + assert.Equal(t, "https://example.com/a/new/c", result) + + // Test replace all (position 0) + result, err = registry.Execute(`path_replace("https://example.com/a/b/c", "x", 0)`, nil) + require.NoError(t, err) + assert.Equal(t, "https://example.com/x/x/x", result) + + // Test with empty URL + result, err = registry.Execute(`path_replace("", "test")`, nil) + require.NoError(t, err) + assert.Equal(t, "", result) +} diff --git a/pkg/cli/api_client.go b/pkg/cli/api_client.go index d39a1a6..4e8f8e7 100644 --- a/pkg/cli/api_client.go +++ b/pkg/cli/api_client.go @@ -121,3 +121,129 @@ func (c *ScheduleClient) RegisterCronTrigger(ctx context.Context, workflow *core body, _ := io.ReadAll(resp.Body) return fmt.Errorf("server returned %d: %s", resp.StatusCode, string(body)) } + +// CreateRunRequest represents a request to create a new run via the API +type CreateRunRequest struct { + Flow string `json:"flow,omitempty"` + Module string `json:"module,omitempty"` + Target string `json:"target,omitempty"` + Targets []string `json:"targets,omitempty"` + Params map[string]string `json:"params,omitempty"` + Concurrency int `json:"concurrency,omitempty"` + Priority string `json:"priority,omitempty"` + RunMode string `json:"run_mode,omitempty"` + ThreadsHold int `json:"threads_hold,omitempty"` + HeuristicsCheck string `json:"heuristics_check,omitempty"` + EmptyTarget bool `json:"empty_target,omitempty"` +} + +// CreateRunResponse represents the response from creating a run +type CreateRunResponse struct { + Message string `json:"message"` + Workflow string `json:"workflow"` + Kind string `json:"kind"` + TargetCount int `json:"target_count"` + Priority string `json:"priority"` + JobID string `json:"job_id"` + Status string `json:"status"` + PollURL string `json:"poll_url"` + RunUUID string `json:"run_uuid,omitempty"` +} + +// RunClient handles run submission to the osmedeus server +type RunClient struct { + baseURL string + apiKey string + client *http.Client +} + +// NewRunClient creates a new RunClient from config +func NewRunClient(cfg *config.Config) *RunClient { + return &RunClient{ + baseURL: cfg.Server.GetServerURL(), + apiKey: cfg.Server.AuthAPIKey, + client: &http.Client{ + Timeout: 30 * time.Second, + }, + } +} + +// SetBaseURL overrides the base URL (for --server-url flag) +func (c *RunClient) SetBaseURL(url string) { + c.baseURL = url +} + +// IsServerAvailable checks if the server is reachable via GET /server-info +func (c *RunClient) IsServerAvailable() bool { + if c.baseURL == "" { + return false + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+"/server-info", nil) + if err != nil { + return false + } + + resp, err := c.client.Do(req) + if err != nil { + return false + } + defer func() { _ = resp.Body.Close() }() + + return resp.StatusCode == http.StatusOK +} + +// CreateRun POSTs to /osm/api/runs to create a new run +func (c *RunClient) CreateRun(ctx context.Context, req *CreateRunRequest) (*CreateRunResponse, error) { + if c.baseURL == "" { + return nil, fmt.Errorf("server URL not configured") + } + + jsonBody, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("failed to marshal request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+"/osm/api/runs", bytes.NewReader(jsonBody)) + if err != nil { + return nil, fmt.Errorf("failed to create request: %w", err) + } + + httpReq.Header.Set("Content-Type", "application/json") + if c.apiKey != "" { + httpReq.Header.Set("x-osm-api-key", c.apiKey) + } + + resp, err := c.client.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("request failed: %w", err) + } + defer func() { _ = resp.Body.Close() }() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("failed to read response: %w", err) + } + + // Check for error responses + if resp.StatusCode >= 400 { + var errResp struct { + Error bool `json:"error"` + Message string `json:"message"` + } + if json.Unmarshal(body, &errResp) == nil && errResp.Message != "" { + return nil, fmt.Errorf("server error: %s", errResp.Message) + } + return nil, fmt.Errorf("server returned %d: %s", resp.StatusCode, string(body)) + } + + var result CreateRunResponse + if err := json.Unmarshal(body, &result); err != nil { + return nil, fmt.Errorf("failed to parse response: %w", err) + } + + return &result, nil +} diff --git a/pkg/cli/api_client_test.go b/pkg/cli/api_client_test.go index 13ab0c6..4eb265b 100644 --- a/pkg/cli/api_client_test.go +++ b/pkg/cli/api_client_test.go @@ -159,3 +159,221 @@ func TestScheduleClient_RegisterCronTrigger(t *testing.T) { assert.Contains(t, err.Error(), "server URL not configured") }) } + +func TestNewRunClient(t *testing.T) { + cfg := &config.Config{ + Server: config.ServerConfig{ + Host: "localhost", + Port: 8002, + AuthAPIKey: "test-api-key", + }, + } + + client := NewRunClient(cfg) + assert.NotNil(t, client) + assert.Equal(t, "http://localhost:8002", client.baseURL) + assert.Equal(t, "test-api-key", client.apiKey) +} + +func TestRunClient_SetBaseURL(t *testing.T) { + client := &RunClient{} + client.SetBaseURL("http://custom:9000") + assert.Equal(t, "http://custom:9000", client.baseURL) +} + +func TestRunClient_IsServerAvailable(t *testing.T) { + t.Run("server available", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/server-info", r.URL.Path) + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"version": "5.0.0"}`)) + })) + defer server.Close() + + client := &RunClient{ + baseURL: server.URL, + client: http.DefaultClient, + } + assert.True(t, client.IsServerAvailable()) + }) + + t.Run("server unavailable", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + defer server.Close() + + client := &RunClient{ + baseURL: server.URL, + client: http.DefaultClient, + } + assert.False(t, client.IsServerAvailable()) + }) + + t.Run("empty baseURL", func(t *testing.T) { + client := &RunClient{ + baseURL: "", + client: http.DefaultClient, + } + assert.False(t, client.IsServerAvailable()) + }) +} + +func TestRunClient_CreateRun(t *testing.T) { + t.Run("success - 202 accepted", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "/osm/api/runs", r.URL.Path) + assert.Equal(t, "POST", r.Method) + assert.Equal(t, "application/json", r.Header.Get("Content-Type")) + assert.Equal(t, "test-key", r.Header.Get("x-osm-api-key")) + + var req CreateRunRequest + err := json.NewDecoder(r.Body).Decode(&req) + require.NoError(t, err) + + assert.Equal(t, "test-module", req.Module) + assert.Equal(t, "example.com", req.Target) + assert.Equal(t, "high", req.Priority) + + w.WriteHeader(http.StatusAccepted) + resp := CreateRunResponse{ + Message: "Run started", + Workflow: "test-module", + Kind: "module", + TargetCount: 1, + Priority: "high", + JobID: "abc123", + Status: "queued", + PollURL: "/osm/api/jobs/abc123", + RunUUID: "run-uuid-123", + } + _ = json.NewEncoder(w).Encode(resp) + })) + defer server.Close() + + client := &RunClient{ + baseURL: server.URL, + apiKey: "test-key", + client: http.DefaultClient, + } + + req := &CreateRunRequest{ + Module: "test-module", + Target: "example.com", + Priority: "high", + } + + resp, err := client.CreateRun(context.Background(), req) + require.NoError(t, err) + assert.Equal(t, "Run started", resp.Message) + assert.Equal(t, "test-module", resp.Workflow) + assert.Equal(t, "high", resp.Priority) + assert.Equal(t, "abc123", resp.JobID) + assert.Equal(t, "run-uuid-123", resp.RunUUID) + }) + + t.Run("failure - 400 bad request", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusBadRequest) + _, _ = w.Write([]byte(`{"error": true, "message": "Invalid priority"}`)) + })) + defer server.Close() + + client := &RunClient{ + baseURL: server.URL, + client: http.DefaultClient, + } + + req := &CreateRunRequest{ + Module: "test-module", + Target: "example.com", + Priority: "invalid", + } + + resp, err := client.CreateRun(context.Background(), req) + assert.Nil(t, resp) + assert.Error(t, err) + assert.Contains(t, err.Error(), "Invalid priority") + }) + + t.Run("failure - 404 workflow not found", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + _, _ = w.Write([]byte(`{"error": true, "message": "Workflow not found"}`)) + })) + defer server.Close() + + client := &RunClient{ + baseURL: server.URL, + client: http.DefaultClient, + } + + req := &CreateRunRequest{ + Module: "nonexistent-module", + Target: "example.com", + } + + resp, err := client.CreateRun(context.Background(), req) + assert.Nil(t, resp) + assert.Error(t, err) + assert.Contains(t, err.Error(), "Workflow not found") + }) + + t.Run("failure - empty baseURL", func(t *testing.T) { + client := &RunClient{ + baseURL: "", + client: http.DefaultClient, + } + + req := &CreateRunRequest{ + Module: "test-module", + Target: "example.com", + } + + resp, err := client.CreateRun(context.Background(), req) + assert.Nil(t, resp) + assert.Error(t, err) + assert.Contains(t, err.Error(), "server URL not configured") + }) + + t.Run("success - multiple targets", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req CreateRunRequest + err := json.NewDecoder(r.Body).Decode(&req) + require.NoError(t, err) + + assert.Equal(t, 3, len(req.Targets)) + assert.Equal(t, 2, req.Concurrency) + + w.WriteHeader(http.StatusAccepted) + resp := CreateRunResponse{ + Message: "Run started", + Workflow: "test-module", + Kind: "module", + TargetCount: 3, + Priority: "normal", + JobID: "multi123", + Status: "queued", + PollURL: "/osm/api/jobs/multi123", + } + _ = json.NewEncoder(w).Encode(resp) + })) + defer server.Close() + + client := &RunClient{ + baseURL: server.URL, + client: http.DefaultClient, + } + + req := &CreateRunRequest{ + Module: "test-module", + Targets: []string{"t1.com", "t2.com", "t3.com"}, + Concurrency: 2, + Priority: "normal", + } + + resp, err := client.CreateRun(context.Background(), req) + require.NoError(t, err) + assert.Equal(t, 3, resp.TargetCount) + }) +} diff --git a/pkg/cli/run.go b/pkg/cli/run.go index dff72ba..bdda41a 100644 --- a/pkg/cli/run.go +++ b/pkg/cli/run.go @@ -75,6 +75,9 @@ var ( // Server registration flag serverURL string + // Run priority flag (for server submission mode) + runPriority string + // activeChunkInfo holds chunk info during execution (nil when not chunking) activeChunkInfo *ChunkInfo @@ -129,6 +132,9 @@ func init() { // Server registration flag runCmd.Flags().StringVar(&serverURL, "server-url", "", "Server URL for cron trigger registration (e.g., http://localhost:8002)") + + // Run priority flag (for server submission mode) + runCmd.Flags().StringVar(&runPriority, "run-priority", "", "Run priority: low, normal, high, critical (requires --server-url to submit to server)") } // captureExplicitFlags records which CLI flags were explicitly set by the user @@ -292,6 +298,17 @@ func runRun(cmd *cobra.Command, args []string) error { return fmt.Errorf("--module-url cannot be combined with --flow, --module, or --std-module") } + // Validate --run-priority flag + if runPriority != "" { + validPriorities := map[string]bool{"low": true, "normal": true, "high": true, "critical": true} + if !validPriorities[runPriority] { + return fmt.Errorf("invalid --run-priority value '%s'. Must be one of: low, normal, high, critical", runPriority) + } + if serverURL == "" { + return fmt.Errorf("--run-priority requires --server-url to be specified") + } + } + // Parse timeout duration var timeoutDuration time.Duration if runTimeout != "" { @@ -396,6 +413,11 @@ func runRun(cmd *cobra.Command, args []string) error { return runDistributedRun(cfg, allTargets, printer) } + // Handle server submission mode (--run-priority with --server-url) + if runPriority != "" && serverURL != "" { + return runServerSubmission(cfg, allTargets, printer) + } + loader := parser.NewLoader(cfg.WorkflowsPath) // Execute workflow for each target (with concurrency) @@ -1010,6 +1032,8 @@ func createCLIRunRecord(ctx context.Context, cfg *config.Config, workflow *core. StartedAt: &now, TotalSteps: calculateTotalSteps(workflow, loader), Workspace: workspace, + RunPriority: "critical", // CLI runs execute immediately + RunMode: "local", } if err := database.CreateRun(ctx, run); err != nil { @@ -1683,6 +1707,108 @@ func runDistributedRun(cfg *config.Config, allTargets []string, printer *termina return nil } +// runServerSubmission submits run tasks to the server API with the specified priority +func runServerSubmission(cfg *config.Config, allTargets []string, printer *terminal.Printer) error { + // Determine workflow name and kind + workflowName := flowName + workflowKind := "flow" + if workflowName == "" && len(moduleNames) > 0 { + workflowName = moduleNames[0] + workflowKind = "module" + } + + if workflowName == "" { + return fmt.Errorf("workflow name required (use -f or -m)") + } + + // Create run client and set server URL + client := NewRunClient(cfg) + client.SetBaseURL(serverURL) + + // Check server availability + printer.Info("Connecting to server at %s...", serverURL) + if !client.IsServerAvailable() { + return fmt.Errorf("server not available at %s", serverURL) + } + printer.Success("Server is available") + + // Parse additional params + params := make(map[string]string) + + // Load params from file if specified + if paramsFile != "" { + fileParams, err := loadParamsFromFile(paramsFile) + if err != nil { + return fmt.Errorf("failed to load params file: %w", err) + } + for k, v := range fileParams { + params[k] = v + } + } + + // CLI params (-p) override file params + for k, v := range parseParams(paramFlags) { + params[k] = v + } + + // Build the request + req := &CreateRunRequest{ + Params: params, + Priority: runPriority, + Concurrency: concurrency, + } + + // Set workflow type + if workflowKind == "flow" { + req.Flow = workflowName + } else { + req.Module = workflowName + } + + // Set target(s) + if len(allTargets) == 1 { + req.Target = allTargets[0] + } else { + req.Targets = allTargets + } + + // Add optional parameters + if threadsHold > 0 { + req.ThreadsHold = threadsHold + } + if heuristicsCheck != "" { + req.HeuristicsCheck = heuristicsCheck + } + if emptyTarget { + req.EmptyTarget = true + } + + // Submit the run + printer.Section("Submitting Run to Server") + ctx := context.Background() + resp, err := client.CreateRun(ctx, req) + if err != nil { + return fmt.Errorf("failed to submit run: %w", err) + } + + // Display response + fmt.Println() + printer.Success("Run submitted successfully!") + printer.KeyValue("Job ID", resp.JobID) + if resp.RunUUID != "" { + printer.KeyValue("Run UUID", resp.RunUUID) + } + printer.KeyValue("Workflow", resp.Workflow) + printer.KeyValue("Kind", resp.Kind) + printer.KeyValue("Priority", resp.Priority) + printer.KeyValue("Status", resp.Status) + printer.KeyValue("Target Count", fmt.Sprintf("%d", resp.TargetCount)) + printer.KeyValue("Poll URL", serverURL+resp.PollURL) + fmt.Println() + + return nil +} + // registerCronTriggersWithServer registers workflow cron triggers with the server. // Best-effort: failures are logged but don't block execution. func registerCronTriggersWithServer(ctx context.Context, workflow *core.Workflow, target string, params map[string]string, cfg *config.Config, printer *terminal.Printer, log *zap.Logger) { diff --git a/pkg/server/handlers/common.go b/pkg/server/handlers/common.go index e2678ba..abbda29 100644 --- a/pkg/server/handlers/common.go +++ b/pkg/server/handlers/common.go @@ -28,7 +28,8 @@ type CreateRunRequest struct { Concurrency int `json:"concurrency,omitempty"` // Number of concurrent runs (default: 1) // Priority and timeout - Priority string `json:"priority,omitempty"` // low, medium, high (default: medium) + Priority string `json:"priority,omitempty"` // low, normal, high, critical (default: high) + RunMode string `json:"run_mode,omitempty"` // local, distributed, cloud (default: local) Timeout int `json:"timeout,omitempty"` // Timeout in minutes (0 = no timeout) // Runner configuration diff --git a/pkg/server/handlers/runs.go b/pkg/server/handlers/runs.go index 8038bad..009b292 100644 --- a/pkg/server/handlers/runs.go +++ b/pkg/server/handlers/runs.go @@ -88,7 +88,7 @@ func sanitizeTargetForWorkspace(target string) string { } // createRunRecord creates a database record for a run -func createRunRecord(ctx context.Context, _ *config.Config, workflow *core.Workflow, target string, params map[string]string, triggerType, jobID string) (*database.Run, error) { +func createRunRecord(ctx context.Context, _ *config.Config, workflow *core.Workflow, target string, params map[string]string, triggerType, jobID, priority, runMode string) (*database.Run, error) { now := time.Now() runID := uuid.New().String() @@ -112,6 +112,8 @@ func createRunRecord(ctx context.Context, _ *config.Config, workflow *core.Workf StartedAt: &now, TotalSteps: calculateTotalSteps(workflow), Workspace: workspace, + RunPriority: priority, + RunMode: runMode, } if err := database.CreateRun(ctx, run); err != nil { @@ -155,6 +157,8 @@ func executeRunsConcurrently( maxConcurrency int, isFlow bool, jobID string, + priority string, + runMode string, ) { if maxConcurrency <= 0 { maxConcurrency = 1 @@ -184,7 +188,7 @@ func executeRunsConcurrently( ctx := context.Background() // Create run record in database - run, err := createRunRecord(ctx, cfg, workflow, t, targetParams, "api", jobID) + run, err := createRunRecord(ctx, cfg, workflow, t, targetParams, "api", jobID, priority, runMode) var runUUID string var runID int64 if err == nil && run != nil { @@ -316,7 +320,29 @@ func CreateRun(cfg *config.Config) fiber.Handler { // Set default priority if not specified priority := req.Priority if priority == "" { - priority = "medium" + priority = "high" + } + // Validate priority + validPriorities := map[string]bool{"low": true, "normal": true, "high": true, "critical": true} + if !validPriorities[priority] { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": true, + "message": "Invalid priority. Must be one of: low, normal, high, critical", + }) + } + + // Set default run_mode if not specified + runMode := req.RunMode + if runMode == "" { + runMode = "local" + } + // Validate run_mode + validModes := map[string]bool{"local": true, "distributed": true, "cloud": true} + if !validModes[runMode] { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": true, + "message": "Invalid run_mode. Must be one of: local, distributed, cloud", + }) } // Add runner configuration to params if specified @@ -366,7 +392,7 @@ func CreateRun(cfg *config.Config) fiber.Handler { // Create run record in database ctx := context.Background() - run, _ := createRunRecord(ctx, cfgCopy, workflow, targets[0], params, "api", jobID) + run, _ := createRunRecord(ctx, cfgCopy, workflow, targets[0], params, "api", jobID, priority, runMode) if run != nil { runIDs = append(runIDs, run.RunUUID) } @@ -408,7 +434,7 @@ func CreateRun(cfg *config.Config) fiber.Handler { }()) } else { // Multiple targets - concurrent execution - go executeRunsConcurrently(workflow, targets, params, cfgCopy, concurrency, isFlow, jobID) + go executeRunsConcurrently(workflow, targets, params, cfgCopy, concurrency, isFlow, jobID, priority, runMode) } // Build response @@ -418,6 +444,7 @@ func CreateRun(cfg *config.Config) fiber.Handler { "kind": workflow.Kind, "target_count": len(targets), "priority": priority, + "run_mode": runMode, "job_id": jobID, "status": "queued", "poll_url": fmt.Sprintf("/osm/api/jobs/%s", jobID), @@ -583,16 +610,16 @@ func CancelRun(cfg *config.Config) fiber.Handler { var killedPIDs []int var killMethod string - // Try to cancel via registry first (kills running processes tracked in memory) - registry := executor.GetRunRegistry() - registryPIDs, registryErr := registry.Cancel(run.RunUUID) + // Try to cancel via control plane first (kills running processes tracked in memory) + controlPlane := executor.GetRunControlPlane() + controlPlanePIDs, controlPlaneErr := controlPlane.Cancel(run.RunUUID) - if registryErr == nil && len(registryPIDs) > 0 { - // Registry had the run and killed processes - killedPIDs = registryPIDs - killMethod = "registry" + if controlPlaneErr == nil && len(controlPlanePIDs) > 0 { + // Control plane had the run and killed processes + killedPIDs = controlPlanePIDs + killMethod = "control_plane" } else if run.CurrentPID > 0 { - // Run not in registry, but we have a PID from database - kill it directly + // Run not in control plane, but we have a PID from database - kill it directly killed := killProcessAndChildren(run.CurrentPID) if killed { killedPIDs = []int{run.CurrentPID} diff --git a/test/e2e/api_test.go b/test/e2e/api_test.go index 239c45b..bd34f51 100644 --- a/test/e2e/api_test.go +++ b/test/e2e/api_test.go @@ -565,64 +565,238 @@ func testVulnerabilityEndpoints(t *testing.T, log *TestLogger) { log.Success("Vulnerability endpoints OK") } -// testRunEndpoints tests run management endpoints +// testRunEndpoints tests run management endpoints comprehensively func testRunEndpoints(t *testing.T, log *TestLogger) { log.Info("Testing run endpoints") - // GET /osm/api/runs - // Note: Current implementation is a stub that returns empty data + // ===== LIST RUNS ===== + log.Info("Testing GET /osm/api/runs") resp := apiGet(t, "/osm/api/runs") assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs should return 200") body := parseJSONResponse(t, resp) assert.Contains(t, body, "data", "Should contain data array") + assert.Contains(t, body, "pagination", "Should contain pagination") data, ok := body["data"].([]interface{}) assert.True(t, ok, "Data should be an array") - // Use a test run ID for endpoint testing (handlers are stubs) - testRunID := "test-run-123" + // Test with pagination parameters + resp = apiGet(t, "/osm/api/runs?offset=0&limit=5") + assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs with pagination should return 200") - // If we have seeded runs, use the first one - if len(data) > 0 { + // Test with status filter + resp = apiGet(t, "/osm/api/runs?status=completed") + assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs?status=completed should return 200") + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "data", "Should contain data array") + + // Test with workflow filter + resp = apiGet(t, "/osm/api/runs?workflow=test-bash") + assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs?workflow=test-bash should return 200") + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "data", "Should contain data array") + + // Test with target filter + resp = apiGet(t, "/osm/api/runs?target=example") + assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs?target=example should return 200") + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "data", "Should contain data array") + + // Test with workspace filter + resp = apiGet(t, "/osm/api/runs?workspace=example.com") + assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs?workspace=example.com should return 200") + + // ===== CREATE RUN - VALIDATION ===== + log.Info("Testing POST /osm/api/runs validation") + + // Test missing workflow + invalidRun := map[string]interface{}{ + "target": "test.example.com", + } + resp = apiPost(t, "/osm/api/runs", invalidRun) + assert.Equal(t, 400, resp.StatusCode, "POST /osm/api/runs without workflow should return 400") + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "error", "Should contain error field") + + // Test missing target + invalidRun = map[string]interface{}{ + "module": "test-bash", + } + resp = apiPost(t, "/osm/api/runs", invalidRun) + assert.Equal(t, 400, resp.StatusCode, "POST /osm/api/runs without target should return 400") + + // Test invalid priority + invalidRun = map[string]interface{}{ + "module": "test-bash", + "target": "test.example.com", + "priority": "invalid-priority", + } + resp = apiPost(t, "/osm/api/runs", invalidRun) + assert.Equal(t, 400, resp.StatusCode, "POST /osm/api/runs with invalid priority should return 400") + body = parseJSONResponse(t, resp) + assert.Contains(t, body["message"], "priority", "Error message should mention priority") + + // Test invalid run_mode + invalidRun = map[string]interface{}{ + "module": "test-bash", + "target": "test.example.com", + "run_mode": "invalid-mode", + } + resp = apiPost(t, "/osm/api/runs", invalidRun) + assert.Equal(t, 400, resp.StatusCode, "POST /osm/api/runs with invalid run_mode should return 400") + body = parseJSONResponse(t, resp) + assert.Contains(t, body["message"], "run_mode", "Error message should mention run_mode") + + // ===== CREATE RUN - ALL PRIORITIES ===== + log.Info("Testing POST /osm/api/runs with all priority levels") + priorities := []string{"low", "normal", "high", "critical"} + for _, priority := range priorities { + runReq := map[string]interface{}{ + "module": "test-bash", + "target": fmt.Sprintf("priority-%s.example.com", priority), + "priority": priority, + } + resp = apiPost(t, "/osm/api/runs", runReq) + // 202 (accepted) or 404 (workflow not found) are valid + assert.True(t, resp.StatusCode == 202 || resp.StatusCode == 404, + "POST /osm/api/runs with priority=%s should return 202 or 404, got %d", priority, resp.StatusCode) + if resp.StatusCode == 202 { + body = parseJSONResponse(t, resp) + assert.Equal(t, priority, body["priority"], "Response priority should match request") + } + } + + // ===== CREATE RUN - VALID REQUEST ===== + log.Info("Testing POST /osm/api/runs with valid request") + validRun := map[string]interface{}{ + "module": "test-bash", + "target": "run-test.example.com", + "priority": "high", + "params": map[string]string{ + "custom_param": "test_value", + }, + } + resp = apiPost(t, "/osm/api/runs", validRun) + // Accept 202 (accepted) or 404 (workflow not found) + assert.True(t, resp.StatusCode == 202 || resp.StatusCode == 404, + "POST /osm/api/runs should return 202 or 404") + + var createdRunUUID string + if resp.StatusCode == 202 { + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "job_id", "Response should contain job_id") + assert.Contains(t, body, "workflow", "Response should contain workflow") + assert.Contains(t, body, "priority", "Response should contain priority") + assert.Contains(t, body, "status", "Response should contain status") + assert.Contains(t, body, "poll_url", "Response should contain poll_url") + if runUUID, ok := body["run_uuid"].(string); ok { + createdRunUUID = runUUID + } + } + + // ===== CREATE RUN - MULTIPLE TARGETS ===== + log.Info("Testing POST /osm/api/runs with multiple targets") + multiTargetRun := map[string]interface{}{ + "module": "test-bash", + "targets": []string{"target1.example.com", "target2.example.com", "target3.example.com"}, + "concurrency": 2, + "priority": "normal", + } + resp = apiPost(t, "/osm/api/runs", multiTargetRun) + assert.True(t, resp.StatusCode == 202 || resp.StatusCode == 404, + "POST /osm/api/runs with multiple targets should return 202 or 404") + if resp.StatusCode == 202 { + body = parseJSONResponse(t, resp) + targetCount, _ := body["target_count"].(float64) + assert.Equal(t, float64(3), targetCount, "Target count should be 3") + assert.Contains(t, body, "concurrency", "Response should contain concurrency") + } + + // ===== CREATE RUN - EMPTY TARGET ===== + log.Info("Testing POST /osm/api/runs with empty_target") + emptyTargetRun := map[string]interface{}{ + "module": "test-bash", + "empty_target": true, + "priority": "low", + } + resp = apiPost(t, "/osm/api/runs", emptyTargetRun) + assert.True(t, resp.StatusCode == 202 || resp.StatusCode == 404, + "POST /osm/api/runs with empty_target should return 202 or 404") + + // ===== GET RUN DETAILS ===== + log.Info("Testing GET /osm/api/runs/:id") + + // Use a run UUID from earlier if we created one, otherwise use test ID + testRunID := "test-run-123" + if createdRunUUID != "" { + testRunID = createdRunUUID + } else if len(data) > 0 { + // Try to get a seeded run UUID if firstRun, ok := data[0].(map[string]interface{}); ok { - if id, ok := firstRun["id"].(string); ok { - testRunID = id + if uuid, ok := firstRun["run_uuid"].(string); ok { + testRunID = uuid } } } - // GET /osm/api/runs/:id + // GET run by ID - may be 200 (found) or 404 (not found) resp = apiGet(t, "/osm/api/runs/"+testRunID) - assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs/:id should return 200") + assert.True(t, resp.StatusCode == 200 || resp.StatusCode == 404, + "GET /osm/api/runs/:id should return 200 or 404") - // GET /osm/api/runs/:id/steps + // Test with include_steps query param + resp = apiGet(t, "/osm/api/runs/"+testRunID+"?include_steps=true") + assert.True(t, resp.StatusCode == 200 || resp.StatusCode == 404, + "GET /osm/api/runs/:id?include_steps=true should return 200 or 404") + + // Test with include_artifacts query param + resp = apiGet(t, "/osm/api/runs/"+testRunID+"?include_artifacts=true") + assert.True(t, resp.StatusCode == 200 || resp.StatusCode == 404, + "GET /osm/api/runs/:id?include_artifacts=true should return 200 or 404") + + // ===== GET RUN STEPS ===== + log.Info("Testing GET /osm/api/runs/:id/steps") resp = apiGet(t, "/osm/api/runs/"+testRunID+"/steps") assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs/:id/steps should return 200") body = parseJSONResponse(t, resp) assert.Contains(t, body, "data", "Should contain steps data") - // GET /osm/api/runs/:id/artifacts + // ===== GET RUN ARTIFACTS ===== + log.Info("Testing GET /osm/api/runs/:id/artifacts") resp = apiGet(t, "/osm/api/runs/"+testRunID+"/artifacts") assert.Equal(t, 200, resp.StatusCode, "GET /osm/api/runs/:id/artifacts should return 200") body = parseJSONResponse(t, resp) assert.Contains(t, body, "data", "Should contain artifacts data") - // POST /osm/api/runs - Create new run (dry-run mode) - newRun := map[string]interface{}{ - "workflow_name": "test-bash", - "target": "test-run.example.com", - "dry_run": true, + // ===== DUPLICATE RUN ===== + log.Info("Testing POST /osm/api/runs/:id/duplicate") + resp = apiPost(t, "/osm/api/runs/"+testRunID+"/duplicate", nil) + // May return 201 (created) or 404 (run not found) + assert.True(t, resp.StatusCode == 201 || resp.StatusCode == 404, + "POST /osm/api/runs/:id/duplicate should return 201 or 404") + if resp.StatusCode == 201 { + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "run_uuid", "Should contain new run_uuid") + assert.Contains(t, body, "original_run_uuid", "Should contain original_run_uuid") + assert.Equal(t, "pending", body["status"], "Duplicated run should be pending") } - resp = apiPost(t, "/osm/api/runs", newRun) - // May return 201 (created) or 202 (accepted) or 400 (if workflow not found) - // Accept 201, 202, or 400 as valid responses - assert.True(t, resp.StatusCode == 201 || resp.StatusCode == 202 || resp.StatusCode == 400, - "POST /osm/api/runs should return 201, 202, or 400 (workflow may not exist)") - // DELETE /osm/api/runs/:id (cancel) - test with a test run ID + // ===== START RUN ===== + log.Info("Testing POST /osm/api/runs/:id/start") + resp = apiPost(t, "/osm/api/runs/"+testRunID+"/start", nil) + // May return 202 (started), 400 (not pending), or 404 (not found) + assert.True(t, resp.StatusCode == 202 || resp.StatusCode == 400 || resp.StatusCode == 404, + "POST /osm/api/runs/:id/start should return 202, 400, or 404") + + // ===== CANCEL RUN ===== + log.Info("Testing DELETE /osm/api/runs/:id (cancel)") resp = apiDelete(t, "/osm/api/runs/"+testRunID) - // May return 200 (cancelled) or 400 (already completed/failed) - assert.True(t, resp.StatusCode == 200 || resp.StatusCode == 400, - "DELETE /osm/api/runs/:id should return 200 or 400") + // May return 200 (cancelled), 400 (cannot cancel), or 404 (not found) + assert.True(t, resp.StatusCode == 200 || resp.StatusCode == 400 || resp.StatusCode == 404, + "DELETE /osm/api/runs/:id should return 200, 400, or 404") + if resp.StatusCode == 200 { + body = parseJSONResponse(t, resp) + assert.Contains(t, body, "message", "Should contain message") + } log.Success("Run endpoints OK") } diff --git a/test/e2e/foreach_preprocess_test.go b/test/e2e/foreach_preprocess_test.go new file mode 100644 index 0000000..a0c6d4e --- /dev/null +++ b/test/e2e/foreach_preprocess_test.go @@ -0,0 +1,144 @@ +package e2e + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// TestForeachPreprocess_DryRun tests that the workflow validates correctly in dry-run mode +func TestForeachPreprocess_DryRun(t *testing.T) { + log := NewTestLogger(t) + log.Step("Testing foreach-preprocess workflow in dry-run mode") + + workflowPath := getTestdataPath(t) + log.Info("Using workflow path: %s", workflowPath) + + stdout, _, err := runCLIWithLog(t, log, "run", "-m", "test-foreach-preprocess", "-t", "example.com", "--dry-run", "-F", workflowPath) + require.NoError(t, err) + + log.Info("Asserting dry-run mode") + assert.Contains(t, stdout, "DRY-RUN") + + log.Info("Asserting workflow name is displayed") + assert.Contains(t, stdout, "test-foreach-preprocess") + + log.Success("foreach-preprocess workflow dry-run validates correctly") +} + +// TestForeachPreprocess_GetParentURL tests that get_parent_url pre-processing works +func TestForeachPreprocess_GetParentURL(t *testing.T) { + log := NewTestLogger(t) + log.Step("Testing get_parent_url() pre-processing in foreach loop") + + workflowPath := getTestdataPath(t) + log.Info("Using workflow path: %s", workflowPath) + + stdout, _, err := runCLIWithLog(t, log, "run", "-m", "test-foreach-preprocess", "-t", "example.com", "-F", workflowPath) + require.NoError(t, err) + + log.Info("Asserting workflow completed") + assert.Contains(t, stdout, "completed") + + log.Info("Asserting parent URLs were extracted correctly") + assert.Contains(t, stdout, "VERIFY:path_parent_ok") + assert.Contains(t, stdout, "VERIFY:api_parent_ok") + assert.Contains(t, stdout, "VERIFY:nested_parent_ok") + + log.Success("get_parent_url() pre-processing works correctly") +} + +// TestForeachPreprocess_ParseURLDomain tests that parse_url pre-processing extracts domains +func TestForeachPreprocess_ParseURLDomain(t *testing.T) { + log := NewTestLogger(t) + log.Step("Testing parse_url() pre-processing to extract domains") + + workflowPath := getTestdataPath(t) + log.Info("Using workflow path: %s", workflowPath) + + stdout, _, err := runCLIWithLog(t, log, "run", "-m", "test-foreach-preprocess", "-t", "example.com", "-F", workflowPath) + require.NoError(t, err) + + log.Info("Asserting domains were extracted") + assert.Contains(t, stdout, "VERIFY:domain1_ok") + assert.Contains(t, stdout, "VERIFY:domain2_ok") + assert.Contains(t, stdout, "VERIFY:domain3_ok") + + log.Success("parse_url() pre-processing extracts domains correctly") +} + +// TestForeachPreprocess_ChainedFunctions tests chained function calls in pre-processing +func TestForeachPreprocess_ChainedFunctions(t *testing.T) { + log := NewTestLogger(t) + log.Step("Testing chained functions in variable_pre_process") + + workflowPath := getTestdataPath(t) + log.Info("Using workflow path: %s", workflowPath) + + stdout, _, err := runCLIWithLog(t, log, "run", "-m", "test-foreach-preprocess", "-t", "example.com", "-F", workflowPath) + require.NoError(t, err) + + log.Info("Asserting chained functions work (to_lower_case + parse_url)") + assert.Contains(t, stdout, "VERIFY:chained1_ok") + assert.Contains(t, stdout, "VERIFY:chained2_ok") + assert.Contains(t, stdout, "VERIFY:chained3_ok") + + log.Success("chained functions in pre-processing work correctly") +} + +// TestForeachPreprocess_NoPreprocess tests that foreach without pre-processing still works +func TestForeachPreprocess_NoPreprocess(t *testing.T) { + log := NewTestLogger(t) + log.Step("Testing foreach without variable_pre_process (original behavior)") + + workflowPath := getTestdataPath(t) + log.Info("Using workflow path: %s", workflowPath) + + stdout, _, err := runCLIWithLog(t, log, "run", "-m", "test-foreach-preprocess", "-t", "example.com", "-F", workflowPath) + require.NoError(t, err) + + log.Info("Asserting original URLs are preserved when no pre-processing") + assert.Contains(t, stdout, "VERIFY:original1_ok") + assert.Contains(t, stdout, "VERIFY:original2_ok") + + log.Success("foreach without pre-processing preserves original values") +} + +// TestForeachPreprocess_FullWorkflow tests the complete workflow execution +func TestForeachPreprocess_FullWorkflow(t *testing.T) { + log := NewTestLogger(t) + log.Step("Testing complete foreach-preprocess workflow execution") + + workflowPath := getTestdataPath(t) + log.Info("Using workflow path: %s", workflowPath) + + stdout, _, err := runCLIWithLog(t, log, "run", "-m", "test-foreach-preprocess", "-t", "example.com", "-F", workflowPath) + require.NoError(t, err) + + log.Info("Asserting workflow completed successfully") + assert.Contains(t, stdout, "completed") + + log.Info("Asserting final summary was reached") + assert.Contains(t, stdout, "=== Foreach Pre-process Test Summary ===") + assert.Contains(t, stdout, "=== Test Complete ===") + + log.Info("Asserting all verification checks passed") + // Parent URL checks + assert.Contains(t, stdout, "VERIFY:path_parent_ok") + assert.Contains(t, stdout, "VERIFY:api_parent_ok") + assert.Contains(t, stdout, "VERIFY:nested_parent_ok") + // Domain extraction checks + assert.Contains(t, stdout, "VERIFY:domain1_ok") + assert.Contains(t, stdout, "VERIFY:domain2_ok") + assert.Contains(t, stdout, "VERIFY:domain3_ok") + // Chained function checks + assert.Contains(t, stdout, "VERIFY:chained1_ok") + assert.Contains(t, stdout, "VERIFY:chained2_ok") + assert.Contains(t, stdout, "VERIFY:chained3_ok") + // Original value checks + assert.Contains(t, stdout, "VERIFY:original1_ok") + assert.Contains(t, stdout, "VERIFY:original2_ok") + + log.Success("complete foreach-preprocess workflow executed successfully") +} diff --git a/test/testdata/workflows/test-foreach-preprocess.yaml b/test/testdata/workflows/test-foreach-preprocess.yaml new file mode 100644 index 0000000..1fa5dda --- /dev/null +++ b/test/testdata/workflows/test-foreach-preprocess.yaml @@ -0,0 +1,126 @@ +name: test-foreach-preprocess +kind: module +description: Test foreach loop with variable_pre_process +tags: test,foreach,loop,preprocess + +params: + - name: target + required: true + +steps: + - name: create-input-urls + type: bash + commands: + - mkdir -p {{Output}}/osm-test + - | + cat > {{Output}}/osm-test/urls.txt << 'EOF' + https://example.com/path/file.php?id=1 + https://example.com/api/v1/users + https://example.com/deep/nested/path/resource + EOF + + - name: test-get-parent-url + log: "Testing get_parent_url pre-processing" + type: foreach + input: "{{Output}}/osm-test/urls.txt" + variable: url + variable_pre_process: "get_parent_url([[url]])" + threads: 1 + step: + name: output-parent-url + type: bash + command: echo "PARENT:[[url]]" >> {{Output}}/osm-test/parent-urls.txt + + - name: verify-parent-urls + type: bash + command: | + echo "=== Parent URL Results ===" + cat {{Output}}/osm-test/parent-urls.txt + # Verify the parent URLs were extracted correctly + grep -q "PARENT:https://example.com/path/" {{Output}}/osm-test/parent-urls.txt && echo "VERIFY:path_parent_ok" + grep -q "PARENT:https://example.com/api/v1/" {{Output}}/osm-test/parent-urls.txt && echo "VERIFY:api_parent_ok" + grep -q "PARENT:https://example.com/deep/nested/path/" {{Output}}/osm-test/parent-urls.txt && echo "VERIFY:nested_parent_ok" + + - name: create-domain-input + type: bash + command: | + cat > {{Output}}/osm-test/mixed-urls.txt << 'EOF' + https://SUB.EXAMPLE.COM/path + https://API.TEST.ORG/endpoint + https://WWW.SAMPLE.NET/page + EOF + + - name: test-parse-url-domain + log: "Testing parse_url pre-processing to extract domain" + type: foreach + input: "{{Output}}/osm-test/mixed-urls.txt" + variable: url + variable_pre_process: "parse_url([[url]], '%d')" + threads: 1 + step: + name: output-domain + type: bash + command: echo "DOMAIN:[[url]]" >> {{Output}}/osm-test/domains.txt + + - name: verify-domains + type: bash + command: | + echo "=== Domain Extraction Results ===" + cat {{Output}}/osm-test/domains.txt + # Verify domains were extracted (parse_url returns lowercase) + grep -qi "DOMAIN:sub.example.com" {{Output}}/osm-test/domains.txt && echo "VERIFY:domain1_ok" + grep -qi "DOMAIN:api.test.org" {{Output}}/osm-test/domains.txt && echo "VERIFY:domain2_ok" + grep -qi "DOMAIN:www.sample.net" {{Output}}/osm-test/domains.txt && echo "VERIFY:domain3_ok" + + - name: test-chained-functions + log: "Testing chained functions in pre-processing" + type: foreach + input: "{{Output}}/osm-test/mixed-urls.txt" + variable: url + variable_pre_process: "to_lower_case(parse_url([[url]], '%d'))" + threads: 1 + step: + name: output-lowercase-domain + type: bash + command: echo "LOWER:[[url]]" >> {{Output}}/osm-test/lowercase-domains.txt + + - name: verify-chained + type: bash + command: | + echo "=== Chained Functions Results ===" + cat {{Output}}/osm-test/lowercase-domains.txt + # Verify chained functions work (lowercase domains) + grep -q "LOWER:sub.example.com" {{Output}}/osm-test/lowercase-domains.txt && echo "VERIFY:chained1_ok" + grep -q "LOWER:api.test.org" {{Output}}/osm-test/lowercase-domains.txt && echo "VERIFY:chained2_ok" + grep -q "LOWER:www.sample.net" {{Output}}/osm-test/lowercase-domains.txt && echo "VERIFY:chained3_ok" + + - name: test-no-preprocess + log: "Testing foreach without pre-processing (original behavior)" + type: foreach + input: "{{Output}}/osm-test/urls.txt" + variable: url + threads: 1 + step: + name: output-original + type: bash + command: echo "ORIGINAL:[[url]]" >> {{Output}}/osm-test/original-urls.txt + + - name: verify-no-preprocess + type: bash + command: | + echo "=== No Pre-process Results ===" + cat {{Output}}/osm-test/original-urls.txt + # Verify original URLs are unchanged + grep -q "ORIGINAL:https://example.com/path/file.php?id=1" {{Output}}/osm-test/original-urls.txt && echo "VERIFY:original1_ok" + grep -q "ORIGINAL:https://example.com/api/v1/users" {{Output}}/osm-test/original-urls.txt && echo "VERIFY:original2_ok" + + - name: final-summary + type: bash + command: | + echo "=== Foreach Pre-process Test Summary ===" + echo "All pre-process tests completed" + echo "=== Test Complete ===" + + - name: cleanup + type: bash + command: rm -rf {{Output}}/osm-test