mirror of
https://github.com/j3ssie/osmedeus.git
synced 2026-10-01 14:05:00 +02:00
Complete rewrite and re-architecture Osmedeus Engine in v5
This commit is contained in:
@@ -0,0 +1,194 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
// DownloadWorkspaceArtifact handles downloading a file artifact from a workspace
|
||||
// @Summary Download workspace artifact
|
||||
// @Description Download a single file under the given workspace by relative artifact path
|
||||
// @Tags Artifacts
|
||||
// @Produce application/octet-stream
|
||||
// @Param workspace_name path string true "Workspace name"
|
||||
// @Param artifact_path query string true "Relative path to artifact under workspace"
|
||||
// @Success 200 {file} binary "Artifact file"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 403 {object} map[string]interface{} "Forbidden"
|
||||
// @Failure 404 {object} map[string]interface{} "Artifact not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to download artifact"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/artifacts/{workspace_name} [get]
|
||||
func DownloadWorkspaceArtifact(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
workspaceName := c.Params("workspace_name")
|
||||
artifactPath := c.Query("artifact_path")
|
||||
|
||||
if !isValidWorkspaceName(workspaceName) {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid workspace name",
|
||||
})
|
||||
}
|
||||
|
||||
if artifactPath == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "artifact_path query parameter is required",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
workspace, err := database.GetWorkspaceByName(ctx, workspaceName)
|
||||
var workspaceDir string
|
||||
if err == nil && workspace != nil && workspace.LocalPath != "" {
|
||||
workspaceDir = workspace.LocalPath
|
||||
} else {
|
||||
workspaceDir = filepath.Join(cfg.GetWorkspacesDir(), workspaceName)
|
||||
}
|
||||
|
||||
cleanRel := filepath.Clean(artifactPath)
|
||||
if cleanRel == "." || filepath.IsAbs(cleanRel) || strings.HasPrefix(cleanRel, ".."+string(filepath.Separator)) || cleanRel == ".." {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid artifact_path",
|
||||
})
|
||||
}
|
||||
|
||||
fullPath := filepath.Join(workspaceDir, cleanRel)
|
||||
if !isPathUnderWorkspace(fullPath, workspaceDir) {
|
||||
return c.Status(fiber.StatusForbidden).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Path traversal attempt detected",
|
||||
})
|
||||
}
|
||||
|
||||
info, err := os.Stat(fullPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Artifact not found",
|
||||
})
|
||||
}
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to read artifact: " + err.Error(),
|
||||
})
|
||||
}
|
||||
if info.IsDir() {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "artifact_path must point to a file",
|
||||
})
|
||||
}
|
||||
|
||||
return c.SendFile(fullPath)
|
||||
}
|
||||
}
|
||||
|
||||
// ListArtifacts handles listing artifacts with pagination and filtering
|
||||
// @Summary List artifacts
|
||||
// @Description Get a paginated list of artifacts with optional filtering and existence checks
|
||||
// @Tags Artifacts
|
||||
// @Produce json
|
||||
// @Param workspace query string false "Filter by workspace name"
|
||||
// @Param search query string false "Search in artifact name/path"
|
||||
// @Param status_code query int false "Filter by HTTP status code (also accepts statusCode)"
|
||||
// @Param verify_exist query bool false "Annotate results with path_exists and path_is_dir" default(false)
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return" default(20)
|
||||
// @Success 200 {object} map[string]interface{} "List of artifacts with pagination"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to fetch artifacts"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/artifacts [get]
|
||||
func ListArtifacts(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
workspace := c.Query("workspace")
|
||||
search := c.Query("search")
|
||||
statusCode, _ := strconv.Atoi(c.Query("statusCode", c.Query("status_code", "0")))
|
||||
verifyExist := c.Query("verify_exist", "false") == "true"
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
result, err := database.ListArtifacts(ctx, database.ArtifactQuery{
|
||||
Workspace: workspace,
|
||||
Search: search,
|
||||
StatusCode: statusCode,
|
||||
Offset: offset,
|
||||
Limit: limit,
|
||||
})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
data := any(result.Data)
|
||||
if verifyExist {
|
||||
annotated := make([]fiber.Map, 0, len(result.Data))
|
||||
for _, a := range result.Data {
|
||||
exists := false
|
||||
isDir := false
|
||||
if a.ArtifactPath != "" {
|
||||
info, statErr := os.Stat(a.ArtifactPath)
|
||||
if statErr == nil {
|
||||
exists = true
|
||||
isDir = info.IsDir()
|
||||
}
|
||||
}
|
||||
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
|
||||
annotated = append(annotated, fiber.Map{
|
||||
"id": a.ID,
|
||||
"run_id": a.RunID,
|
||||
"workspace": a.Workspace,
|
||||
"name": a.Name,
|
||||
"artifact_path": a.ArtifactPath,
|
||||
"artifact_type": a.ArtifactType,
|
||||
"content_type": a.ContentType,
|
||||
"size_bytes": a.SizeBytes,
|
||||
"line_count": a.LineCount,
|
||||
"description": a.Description,
|
||||
"created_at": a.CreatedAt,
|
||||
"path_exists": exists,
|
||||
"path_is_dir": isDir,
|
||||
})
|
||||
}
|
||||
data = annotated
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
// ListAssets handles listing assets with pagination and filtering
|
||||
// @Summary List assets
|
||||
// @Description Get a paginated list of assets with optional filtering
|
||||
// @Tags Assets
|
||||
// @Produce json
|
||||
// @Param workspace query string false "Filter by workspace name"
|
||||
// @Param search query string false "Search in asset_value, url, title, host_ip"
|
||||
// @Param status_code query int false "Filter by HTTP status code"
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return" default(20)
|
||||
// @Success 200 {object} map[string]interface{} "List of assets with pagination"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to fetch assets"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/assets [get]
|
||||
func ListAssets(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse query parameters
|
||||
workspace := c.Query("workspace")
|
||||
search := c.Query("search")
|
||||
statusCode, _ := strconv.Atoi(c.Query("status_code", "0"))
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
|
||||
// Validate pagination
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Get assets from database
|
||||
result, err := database.ListAssets(ctx, database.AssetQuery{
|
||||
Workspace: workspace,
|
||||
Search: search,
|
||||
StatusCode: statusCode,
|
||||
Offset: offset,
|
||||
Limit: limit,
|
||||
})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/pkg/server/middleware"
|
||||
)
|
||||
|
||||
// Login handles user authentication
|
||||
// @Summary User login
|
||||
// @Description Authenticate user and get JWT token
|
||||
// @Tags Auth
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param credentials body LoginRequest true "Login credentials"
|
||||
// @Success 200 {object} map[string]string "JWT token"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 401 {object} map[string]interface{} "Invalid credentials"
|
||||
// @Router /osm/api/login [post]
|
||||
func Login(cfg *config.Config, noAuth bool) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req LoginRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
// If running in no-auth mode, accept any credentials
|
||||
if noAuth {
|
||||
// Use a default username if none provided
|
||||
username := req.Username
|
||||
if username == "" {
|
||||
username = "anonymous"
|
||||
}
|
||||
token, err := middleware.GenerateToken(username, cfg)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to generate token",
|
||||
})
|
||||
}
|
||||
return c.JSON(fiber.Map{
|
||||
"token": token,
|
||||
})
|
||||
}
|
||||
|
||||
// Validate credentials against user map
|
||||
expectedPassword, userExists := cfg.Server.SimpleUserMapKey[req.Username]
|
||||
if !userExists || expectedPassword != req.Password {
|
||||
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid credentials",
|
||||
})
|
||||
}
|
||||
|
||||
// Generate token
|
||||
token, err := middleware.GenerateToken(req.Username, cfg)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to generate token",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"token": token,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// RefreshToken handles token refresh
|
||||
func RefreshToken(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
claims := middleware.GetUser(c)
|
||||
if claims == nil {
|
||||
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid token",
|
||||
})
|
||||
}
|
||||
|
||||
token, err := middleware.GenerateToken(claims.Username, cfg)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to generate token",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"token": token,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// LoginRequest represents login credentials
|
||||
type LoginRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
// CreateRunRequest represents a run creation request
|
||||
type CreateRunRequest struct {
|
||||
// Workflow identification
|
||||
Flow string `json:"flow"` // Flow workflow name
|
||||
Module string `json:"module"` // Module workflow name
|
||||
Target string `json:"target,omitempty"`
|
||||
Params map[string]string `json:"params"`
|
||||
|
||||
// Multi-target support
|
||||
Targets []string `json:"targets,omitempty"` // Array of targets to run against
|
||||
TargetFile string `json:"target_file,omitempty"` // Path to file containing targets (one per line)
|
||||
|
||||
// Concurrency control
|
||||
Concurrency int `json:"concurrency,omitempty"` // Number of concurrent runs (default: 1)
|
||||
|
||||
// Priority and timeout
|
||||
Priority string `json:"priority,omitempty"` // low, medium, high (default: medium)
|
||||
Timeout int `json:"timeout,omitempty"` // Timeout in minutes (0 = no timeout)
|
||||
|
||||
// Runner configuration
|
||||
RunnerType string `json:"runner_type,omitempty"` // host, docker, ssh (default: host)
|
||||
DockerImage string `json:"docker_image,omitempty"` // Docker image to use when runner_type=docker
|
||||
SSHHost string `json:"ssh_host,omitempty"` // SSH host when runner_type=ssh
|
||||
|
||||
// Scheduling options
|
||||
Schedule string `json:"schedule,omitempty"` // Cron expression for scheduled scans
|
||||
ScheduleEnabled bool `json:"schedule_enabled,omitempty"` // Enable scheduled execution
|
||||
NotifyOnComplete bool `json:"notify_on_complete,omitempty"` // Send notification when run completes
|
||||
|
||||
// Execution options (mirrors CLI flags)
|
||||
ThreadsHold int `json:"threads_hold,omitempty"` // Override thread count (0 = use tactic default)
|
||||
EmptyTarget bool `json:"empty_target,omitempty"` // Run without target (generates placeholder target)
|
||||
Repeat bool `json:"repeat,omitempty"` // Repeat run after completion
|
||||
RepeatWaitTime string `json:"repeat_wait_time,omitempty"` // Wait time between repeats (e.g., 30s, 20m, 10h, 1d)
|
||||
HeuristicsCheck string `json:"heuristics_check,omitempty"` // Heuristics check level: none, basic, advanced
|
||||
}
|
||||
|
||||
// CreateScheduleRequest represents a schedule creation request
|
||||
type CreateScheduleRequest struct {
|
||||
Name string `json:"name"`
|
||||
WorkflowName string `json:"workflow_name"`
|
||||
WorkflowKind string `json:"workflow_kind"` // flow or module
|
||||
Target string `json:"target"`
|
||||
Schedule string `json:"schedule"` // cron expression
|
||||
Params map[string]string `json:"params,omitempty"`
|
||||
Enabled bool `json:"enabled"`
|
||||
RunnerType string `json:"runner_type,omitempty"`
|
||||
}
|
||||
|
||||
// UpdateScheduleRequest represents a schedule update request
|
||||
type UpdateScheduleRequest struct {
|
||||
Name string `json:"name,omitempty"`
|
||||
Target string `json:"target,omitempty"`
|
||||
Schedule string `json:"schedule,omitempty"`
|
||||
Params map[string]string `json:"params,omitempty"`
|
||||
Enabled *bool `json:"enabled,omitempty"`
|
||||
}
|
||||
|
||||
// BinaryStatusEntry represents a binary with its registry info and installation status
|
||||
type BinaryStatusEntry struct {
|
||||
Desc string `json:"desc,omitempty"`
|
||||
RepoLink string `json:"repo_link,omitempty"`
|
||||
Version string `json:"version,omitempty"`
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
ValidateCommand string `json:"valide-command,omitempty"`
|
||||
Linux map[string]string `json:"linux,omitempty"`
|
||||
Darwin map[string]string `json:"darwin,omitempty"`
|
||||
Windows map[string]string `json:"windows,omitempty"`
|
||||
CommandLinux map[string]string `json:"command-linux,omitempty"`
|
||||
CommandDarwin map[string]string `json:"command-darwin,omitempty"`
|
||||
CommandDual map[string]string `json:"command-dual,omitempty"`
|
||||
MultiCommandsLinux []string `json:"multi-commands-linux,omitempty"`
|
||||
MultiCommandsDarwin []string `json:"multi-commands-darwin,omitempty"`
|
||||
Installed bool `json:"installed"`
|
||||
Path string `json:"path,omitempty"`
|
||||
}
|
||||
|
||||
// InstallRequest represents an installation request
|
||||
type InstallRequest struct {
|
||||
Type string `json:"type"` // "binary" or "workflow"
|
||||
Names []string `json:"names,omitempty"` // Binary names to install (for type=binary)
|
||||
Source string `json:"source,omitempty"` // Git URL, zip URL, or file path (for type=workflow)
|
||||
RegistryURL string `json:"registry_url,omitempty"` // Custom registry URL (optional, for type=binary)
|
||||
InstallAll bool `json:"install_all,omitempty"` // Install all binaries from registry (for type=binary)
|
||||
RegistryMode string `json:"registry_mode,omitempty"` // "direct-fetch" or "nix-build" (default: direct-fetch)
|
||||
}
|
||||
|
||||
// FunctionEvalRequest represents a function evaluation request
|
||||
type FunctionEvalRequest struct {
|
||||
Script string `json:"script"`
|
||||
Target string `json:"target,omitempty"`
|
||||
Params map[string]string `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
// FunctionListResponse represents a function in the list response
|
||||
type FunctionListResponse struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
ReturnType string `json:"return_type"`
|
||||
Example string `json:"example,omitempty"`
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
}
|
||||
|
||||
// readTargetsFromFile reads targets from a file (one per line)
|
||||
func readTargetsFromFile(filePath string) ([]string, error) {
|
||||
file, err := os.Open(filePath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
var result []string
|
||||
scanner := bufio.NewScanner(file)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
// Skip empty lines and comments
|
||||
if line != "" && !strings.HasPrefix(line, "#") {
|
||||
result = append(result, line)
|
||||
}
|
||||
}
|
||||
return result, scanner.Err()
|
||||
}
|
||||
|
||||
// deduplicateTargets removes duplicates while preserving order
|
||||
func deduplicateTargets(targets []string) []string {
|
||||
seen := make(map[string]bool)
|
||||
var result []string
|
||||
for _, t := range targets {
|
||||
t = strings.TrimSpace(t)
|
||||
if t != "" && !seen[t] {
|
||||
seen[t] = true
|
||||
result = append(result, t)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,264 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/distributed"
|
||||
)
|
||||
|
||||
// ListWorkers returns all registered workers with their current status
|
||||
// @Summary List all workers
|
||||
// @Description Get a list of all registered workers in the distributed pool
|
||||
// @Tags Distributed
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]interface{} "List of workers"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to list workers"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workers [get]
|
||||
func ListWorkers(master *distributed.Master) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
ctx := c.Context()
|
||||
workers, err := master.ListWorkers(ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
data := make([]fiber.Map, 0, len(workers))
|
||||
for _, w := range workers {
|
||||
data = append(data, fiber.Map{
|
||||
"id": w.ID,
|
||||
"hostname": w.Hostname,
|
||||
"status": w.Status,
|
||||
"current_task": w.CurrentTaskID,
|
||||
"joined_at": w.JoinedAt,
|
||||
"last_heartbeat": w.LastHeartbeat,
|
||||
"tasks_complete": w.TasksComplete,
|
||||
"tasks_failed": w.TasksFailed,
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": data,
|
||||
"count": len(data),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetWorker returns details for a specific worker
|
||||
// @Summary Get worker details
|
||||
// @Description Get details for a specific worker by ID
|
||||
// @Tags Distributed
|
||||
// @Produce json
|
||||
// @Param id path string true "Worker ID"
|
||||
// @Success 200 {object} map[string]interface{} "Worker details"
|
||||
// @Failure 404 {object} map[string]interface{} "Worker not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to get worker"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workers/{id} [get]
|
||||
func GetWorker(master *distributed.Master) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
workerID := c.Params("id")
|
||||
ctx := c.Context()
|
||||
|
||||
workers, err := master.ListWorkers(ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
for _, w := range workers {
|
||||
if w.ID == workerID {
|
||||
return c.JSON(fiber.Map{
|
||||
"id": w.ID,
|
||||
"hostname": w.Hostname,
|
||||
"status": w.Status,
|
||||
"current_task": w.CurrentTaskID,
|
||||
"joined_at": w.JoinedAt,
|
||||
"last_heartbeat": w.LastHeartbeat,
|
||||
"tasks_complete": w.TasksComplete,
|
||||
"tasks_failed": w.TasksFailed,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Worker not found",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ListTasks returns all tasks (running and completed)
|
||||
// @Summary List all tasks
|
||||
// @Description Get a list of all running and completed tasks
|
||||
// @Tags Distributed
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]interface{} "List of running and completed tasks"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to list tasks"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/tasks [get]
|
||||
func ListTasks(master *distributed.Master) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
ctx := c.Context()
|
||||
running, completed, err := master.ListTasks(ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
runningData := make([]fiber.Map, 0, len(running))
|
||||
for _, t := range running {
|
||||
runningData = append(runningData, taskToMap(t))
|
||||
}
|
||||
|
||||
completedData := make([]fiber.Map, 0)
|
||||
for _, r := range completed {
|
||||
completedData = append(completedData, taskResultToMap(r))
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"running": runningData,
|
||||
"completed": completedData,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetTask returns details for a specific task
|
||||
// @Summary Get task details
|
||||
// @Description Get details for a specific task by ID
|
||||
// @Tags Distributed
|
||||
// @Produce json
|
||||
// @Param id path string true "Task ID"
|
||||
// @Success 200 {object} map[string]interface{} "Task details"
|
||||
// @Failure 404 {object} map[string]interface{} "Task not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/tasks/{id} [get]
|
||||
func GetTask(master *distributed.Master) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
taskID := c.Params("id")
|
||||
ctx := c.Context()
|
||||
|
||||
task, result, err := master.GetTaskStatus(ctx, taskID)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
if task != nil {
|
||||
return c.JSON(taskToMap(task))
|
||||
}
|
||||
|
||||
if result != nil {
|
||||
return c.JSON(taskResultToMap(result))
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Task not found",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// SubmitTaskRequest represents a task submission request
|
||||
type SubmitTaskRequest struct {
|
||||
WorkflowName string `json:"workflow_name"`
|
||||
WorkflowKind string `json:"workflow_kind"`
|
||||
Target string `json:"target"`
|
||||
Params map[string]interface{} `json:"params"`
|
||||
}
|
||||
|
||||
// SubmitTask submits a new task to the distributed queue
|
||||
// @Summary Submit a new task
|
||||
// @Description Submit a new task to the distributed worker queue
|
||||
// @Tags Distributed
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param task body SubmitTaskRequest true "Task configuration"
|
||||
// @Success 202 {object} map[string]interface{} "Task submitted"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to submit task"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/tasks [post]
|
||||
func SubmitTask(master *distributed.Master) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req SubmitTaskRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
if req.WorkflowName == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "workflow_name is required",
|
||||
})
|
||||
}
|
||||
|
||||
if req.Target == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "target is required",
|
||||
})
|
||||
}
|
||||
|
||||
task := &distributed.Task{
|
||||
WorkflowName: req.WorkflowName,
|
||||
WorkflowKind: req.WorkflowKind,
|
||||
Target: req.Target,
|
||||
Params: req.Params,
|
||||
}
|
||||
|
||||
ctx := c.Context()
|
||||
if err := master.SubmitTask(ctx, task); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusAccepted).JSON(fiber.Map{
|
||||
"message": "Task submitted",
|
||||
"task_id": task.ID,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// taskToMap converts a Task to a fiber.Map
|
||||
func taskToMap(t *distributed.Task) fiber.Map {
|
||||
m := fiber.Map{
|
||||
"id": t.ID,
|
||||
"scan_id": t.ScanID,
|
||||
"workflow_name": t.WorkflowName,
|
||||
"workflow_kind": t.WorkflowKind,
|
||||
"target": t.Target,
|
||||
"status": t.Status,
|
||||
"worker_id": t.WorkerID,
|
||||
"created_at": t.CreatedAt,
|
||||
}
|
||||
if t.StartedAt != nil {
|
||||
m["started_at"] = t.StartedAt
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// taskResultToMap converts a TaskResult to a fiber.Map
|
||||
func taskResultToMap(r *distributed.TaskResult) fiber.Map {
|
||||
return fiber.Map{
|
||||
"task_id": r.TaskID,
|
||||
"status": r.Status,
|
||||
"output": r.Output,
|
||||
"error": r.Error,
|
||||
"exports": r.Exports,
|
||||
"completed_at": r.CompletedAt,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
// ListEventLogs handles listing event logs with filtering and pagination
|
||||
// @Summary List event logs
|
||||
// @Description Get a paginated list of event logs with optional filtering
|
||||
// @Tags EventLogs
|
||||
// @Produce json
|
||||
// @Param topic query string false "Filter by event topic (e.g., run.started, run.completed)"
|
||||
// @Param name query string false "Filter by event name"
|
||||
// @Param source query string false "Filter by source (scheduler, api, webhook)"
|
||||
// @Param workspace query string false "Filter by workspace"
|
||||
// @Param run_id query string false "Filter by run ID"
|
||||
// @Param workflow_name query string false "Filter by workflow name"
|
||||
// @Param processed query string false "Filter by processed status (true/false)"
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return" default(20)
|
||||
// @Success 200 {object} map[string]interface{} "List of event logs with pagination"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to fetch event logs"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/event-logs [get]
|
||||
func ListEventLogs(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse query parameters
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
|
||||
// Validate pagination
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
// Build query
|
||||
query := database.EventLogQuery{
|
||||
Topic: c.Query("topic"),
|
||||
Name: c.Query("name"),
|
||||
Source: c.Query("source"),
|
||||
Workspace: c.Query("workspace"),
|
||||
RunID: c.Query("run_id"),
|
||||
WorkflowName: c.Query("workflow_name"),
|
||||
Offset: offset,
|
||||
Limit: limit,
|
||||
}
|
||||
|
||||
// Handle processed filter (optional bool)
|
||||
if processedStr := c.Query("processed"); processedStr != "" {
|
||||
processed := processedStr == "true"
|
||||
query.Processed = &processed
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Fetch event logs
|
||||
result, err := database.ListEventLogs(ctx, query)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/functions"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/template"
|
||||
)
|
||||
|
||||
// FunctionEval executes a utility function script
|
||||
// @Summary Execute utility function
|
||||
// @Description Execute a utility function script with template rendering and JavaScript execution
|
||||
// @Tags Functions
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body FunctionEvalRequest true "Function evaluation request"
|
||||
// @Success 200 {object} map[string]interface{} "Evaluation result"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/functions/eval [post]
|
||||
func FunctionEval(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req FunctionEvalRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
if req.Script == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "script is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Build context with target and params
|
||||
ctx := make(map[string]interface{})
|
||||
if req.Target != "" {
|
||||
ctx["target"] = req.Target
|
||||
}
|
||||
for k, v := range req.Params {
|
||||
ctx[k] = v
|
||||
}
|
||||
|
||||
// Render template variables ({{target}}, etc.)
|
||||
templateEngine := template.NewEngine()
|
||||
renderedScript, err := templateEngine.Render(req.Script, ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Template rendering failed: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Execute as JavaScript using Otto runtime
|
||||
registry := functions.NewRegistry()
|
||||
result, err := registry.Execute(renderedScript, ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Execution failed: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"result": result,
|
||||
"rendered_script": renderedScript,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// FunctionList returns all available utility functions
|
||||
// @Summary List utility functions
|
||||
// @Description Get a flat list of all available utility functions with metadata
|
||||
// @Tags Functions
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]interface{} "List of functions with total count"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/functions/list [get]
|
||||
func FunctionList(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Get function registry from constants (single source of truth)
|
||||
registry := functions.FunctionRegistry()
|
||||
|
||||
// Convert to flat array format
|
||||
var result []FunctionListResponse
|
||||
for category, funcs := range registry {
|
||||
for _, fn := range funcs {
|
||||
result = append(result, FunctionListResponse{
|
||||
Name: fn.Signature,
|
||||
Description: fn.Description,
|
||||
ReturnType: fn.ReturnType,
|
||||
Example: fn.Example,
|
||||
Tags: []string{category},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"functions": result,
|
||||
"total": len(result),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,617 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/core"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func setupTestWorkflowDir(t *testing.T) (*config.Config, string) {
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
// Create modules directory
|
||||
modulesDir := filepath.Join(tmpDir, "modules")
|
||||
err := os.MkdirAll(modulesDir, 0755)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create a test flow
|
||||
flowContent := `kind: flow
|
||||
name: test-flow
|
||||
description: Test flow for API testing
|
||||
|
||||
params:
|
||||
- name: target
|
||||
required: true
|
||||
|
||||
modules:
|
||||
- name: test-module
|
||||
path: modules/test-module.yaml
|
||||
`
|
||||
err = os.WriteFile(filepath.Join(tmpDir, "test-flow.yaml"), []byte(flowContent), 0644)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create a test module
|
||||
moduleContent := `kind: module
|
||||
name: test-module
|
||||
description: Test module for API testing
|
||||
|
||||
params:
|
||||
- name: target
|
||||
required: true
|
||||
- name: threads
|
||||
default: "10"
|
||||
|
||||
steps:
|
||||
- name: echo-test
|
||||
type: bash
|
||||
command: echo "Hello {{target}}"
|
||||
`
|
||||
err = os.WriteFile(filepath.Join(modulesDir, "test-module.yaml"), []byte(moduleContent), 0644)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Create another module
|
||||
module2Content := `kind: module
|
||||
name: scan-module
|
||||
description: Scan module for testing
|
||||
|
||||
trigger:
|
||||
- name: manual
|
||||
on: manual
|
||||
enabled: true
|
||||
- name: cron
|
||||
on: cron
|
||||
schedule: "0 0 * * *"
|
||||
enabled: false
|
||||
|
||||
params:
|
||||
- name: target
|
||||
required: true
|
||||
|
||||
steps:
|
||||
- name: scan
|
||||
type: bash
|
||||
command: echo "Scanning {{target}}"
|
||||
- name: report
|
||||
type: bash
|
||||
command: echo "Reporting"
|
||||
`
|
||||
err = os.WriteFile(filepath.Join(modulesDir, "scan-module.yaml"), []byte(module2Content), 0644)
|
||||
require.NoError(t, err)
|
||||
|
||||
cfg := &config.Config{
|
||||
WorkflowsPath: tmpDir,
|
||||
}
|
||||
|
||||
return cfg, tmpDir
|
||||
}
|
||||
|
||||
func TestHealthCheck(t *testing.T) {
|
||||
app := fiber.New()
|
||||
app.Get("/health", HealthCheck)
|
||||
|
||||
req := httptest.NewRequest("GET", "/health", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "ok", result["status"])
|
||||
}
|
||||
|
||||
func TestReadinessCheck(t *testing.T) {
|
||||
app := fiber.New()
|
||||
app.Get("/ready", ReadinessCheck)
|
||||
|
||||
req := httptest.NewRequest("GET", "/ready", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "ready", result["status"])
|
||||
}
|
||||
|
||||
func TestListWorkflows(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows", ListWorkflows(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
data, ok := result["data"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Len(t, data, 3) // 1 flow + 2 modules
|
||||
}
|
||||
|
||||
func TestListWorkflowsVerbose(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows", ListWorkflowsVerbose(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
data, ok := result["data"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Len(t, data, 3)
|
||||
|
||||
// Check that verbose data includes params
|
||||
for _, item := range data {
|
||||
wf, ok := item.(map[string]interface{})
|
||||
require.True(t, ok)
|
||||
|
||||
assert.Contains(t, wf, "name")
|
||||
assert.Contains(t, wf, "kind")
|
||||
assert.Contains(t, wf, "description")
|
||||
assert.Contains(t, wf, "params")
|
||||
assert.Contains(t, wf, "required_params")
|
||||
}
|
||||
|
||||
// Check count
|
||||
count, ok := result["count"].(float64)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, float64(3), count)
|
||||
}
|
||||
|
||||
func TestGetWorkflow(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows/:name", GetWorkflow(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows/test-module", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "test-module", result["name"])
|
||||
assert.Equal(t, "module", result["kind"])
|
||||
assert.Equal(t, "Test module for API testing", result["description"])
|
||||
assert.Contains(t, result, "params")
|
||||
}
|
||||
|
||||
func TestGetWorkflow_NotFound(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows/:name", GetWorkflow(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows/nonexistent", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusNotFound, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, true, result["error"])
|
||||
assert.Contains(t, result["message"], "not found")
|
||||
}
|
||||
|
||||
func TestListWorkspaceNames(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
dbPath := filepath.Join(tmpDir, "test.sqlite")
|
||||
cfg := &config.Config{
|
||||
BaseFolder: tmpDir,
|
||||
Database: config.DatabaseConfig{
|
||||
DBEngine: "sqlite",
|
||||
DBPath: dbPath,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := database.Connect(cfg)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() {
|
||||
_ = database.Close()
|
||||
database.SetDB(nil)
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
require.NoError(t, database.Migrate(ctx))
|
||||
|
||||
now := time.Now()
|
||||
ws1 := &database.Workspace{Name: "b.example", DataSource: "local", Tags: []string{}, CreatedAt: now, UpdatedAt: now}
|
||||
ws2 := &database.Workspace{Name: "a.example", DataSource: "local", Tags: []string{}, CreatedAt: now, UpdatedAt: now}
|
||||
_, err = database.GetDB().NewInsert().Model(ws1).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
_, err = database.GetDB().NewInsert().Model(ws2).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workspace-names", ListWorkspaceNames(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workspace-names", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var names []string
|
||||
err = json.Unmarshal(body, &names)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"a.example", "b.example"}, names)
|
||||
}
|
||||
|
||||
func TestListWorkspaces_FilesystemIncludesWorkspaceFolders(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
dbPath := filepath.Join(tmpDir, "test.sqlite")
|
||||
workspacesDir := filepath.Join(tmpDir, "workspaces")
|
||||
require.NoError(t, os.MkdirAll(filepath.Join(workspacesDir, "fs-only"), 0755))
|
||||
require.NoError(t, os.MkdirAll(filepath.Join(workspacesDir, "shared"), 0755))
|
||||
|
||||
cfg := &config.Config{
|
||||
BaseFolder: tmpDir,
|
||||
WorkspacesPath: workspacesDir,
|
||||
Database: config.DatabaseConfig{
|
||||
DBEngine: "sqlite",
|
||||
DBPath: dbPath,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := database.Connect(cfg)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() {
|
||||
_ = database.Close()
|
||||
database.SetDB(nil)
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
require.NoError(t, database.Migrate(ctx))
|
||||
|
||||
now := time.Now()
|
||||
asset1 := &database.Asset{Workspace: "db-only", AssetValue: "http://example.com", CreatedAt: now, UpdatedAt: now}
|
||||
asset2 := &database.Asset{Workspace: "shared", AssetValue: "http://shared.example.com", CreatedAt: now, UpdatedAt: now}
|
||||
_, err = database.GetDB().NewInsert().Model(asset1).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
_, err = database.GetDB().NewInsert().Model(asset2).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workspaces", ListWorkspaces(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workspaces?filesystem=true&limit=100", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
data, ok := result["data"].([]interface{})
|
||||
require.True(t, ok)
|
||||
|
||||
foundTotalAssets := map[string]float64{}
|
||||
foundDataSource := map[string]string{}
|
||||
foundTags := map[string][]interface{}{}
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
require.True(t, ok)
|
||||
name, _ := m["name"].(string)
|
||||
totalAssets, _ := m["total_assets"].(float64)
|
||||
dataSource, _ := m["data_source"].(string)
|
||||
tags, _ := m["tags"].([]interface{})
|
||||
foundTotalAssets[name] = totalAssets
|
||||
foundDataSource[name] = dataSource
|
||||
foundTags[name] = tags
|
||||
}
|
||||
|
||||
assert.Contains(t, foundTotalAssets, "db-only")
|
||||
assert.Contains(t, foundTotalAssets, "fs-only")
|
||||
assert.Contains(t, foundTotalAssets, "shared")
|
||||
assert.Equal(t, float64(1), foundTotalAssets["db-only"])
|
||||
assert.Equal(t, float64(0), foundTotalAssets["fs-only"])
|
||||
assert.Equal(t, float64(1), foundTotalAssets["shared"])
|
||||
|
||||
assert.Equal(t, "filesystem", foundDataSource["db-only"])
|
||||
assert.Equal(t, "filesystem", foundDataSource["fs-only"])
|
||||
assert.Equal(t, "filesystem", foundDataSource["shared"])
|
||||
|
||||
assert.Contains(t, foundTags["db-only"], "filesystem")
|
||||
assert.Contains(t, foundTags["fs-only"], "filesystem")
|
||||
assert.Contains(t, foundTags["shared"], "filesystem")
|
||||
|
||||
assert.NotContains(t, foundTags["db-only"], "filesystem-only")
|
||||
assert.Contains(t, foundTags["fs-only"], "filesystem-only")
|
||||
assert.NotContains(t, foundTags["shared"], "filesystem-only")
|
||||
|
||||
pagination, ok := result["pagination"].(map[string]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, float64(3), pagination["total"])
|
||||
}
|
||||
|
||||
func TestListArtifactsVerifyExist(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
dbPath := filepath.Join(tmpDir, "test.sqlite")
|
||||
cfg := &config.Config{
|
||||
BaseFolder: tmpDir,
|
||||
Database: config.DatabaseConfig{
|
||||
DBEngine: "sqlite",
|
||||
DBPath: dbPath,
|
||||
},
|
||||
}
|
||||
|
||||
_, err := database.Connect(cfg)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() {
|
||||
_ = database.Close()
|
||||
database.SetDB(nil)
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
require.NoError(t, database.Migrate(ctx))
|
||||
|
||||
filePath := filepath.Join(tmpDir, "out.txt")
|
||||
require.NoError(t, os.WriteFile(filePath, []byte("ok"), 0644))
|
||||
|
||||
folderPath := filepath.Join(tmpDir, "outdir")
|
||||
require.NoError(t, os.MkdirAll(folderPath, 0755))
|
||||
|
||||
now := time.Now()
|
||||
art1 := &database.Artifact{
|
||||
ID: "a1",
|
||||
RunID: "r1",
|
||||
Workspace: "w1",
|
||||
Name: "file",
|
||||
ArtifactPath: filePath,
|
||||
ArtifactType: database.ArtifactTypeOutput,
|
||||
ContentType: database.ContentTypeText,
|
||||
SizeBytes: 2,
|
||||
LineCount: 1,
|
||||
CreatedAt: now,
|
||||
}
|
||||
art2 := &database.Artifact{
|
||||
ID: "a2",
|
||||
RunID: "r1",
|
||||
Workspace: "w1",
|
||||
Name: "folder",
|
||||
ArtifactPath: folderPath,
|
||||
ArtifactType: database.ArtifactTypeOutput,
|
||||
ContentType: database.ContentTypeFolder,
|
||||
CreatedAt: now,
|
||||
}
|
||||
art3 := &database.Artifact{
|
||||
ID: "a3",
|
||||
RunID: "r1",
|
||||
Workspace: "w1",
|
||||
Name: "missing",
|
||||
ArtifactPath: filepath.Join(tmpDir, "missing.txt"),
|
||||
ArtifactType: database.ArtifactTypeOutput,
|
||||
ContentType: database.ContentTypeText,
|
||||
CreatedAt: now,
|
||||
}
|
||||
|
||||
_, err = database.GetDB().NewInsert().Model(art1).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
_, err = database.GetDB().NewInsert().Model(art2).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
_, err = database.GetDB().NewInsert().Model(art3).Exec(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/artifacts", ListArtifacts(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/artifacts?verify_exist=true&limit=100", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
data, ok := result["data"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, 2, len(data))
|
||||
|
||||
byID := map[string]map[string]interface{}{}
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
require.True(t, ok)
|
||||
id, _ := m["id"].(string)
|
||||
if id != "" {
|
||||
byID[id] = m
|
||||
}
|
||||
}
|
||||
|
||||
assert.Equal(t, true, byID["a1"]["path_exists"])
|
||||
assert.Equal(t, false, byID["a1"]["path_is_dir"])
|
||||
assert.Equal(t, true, byID["a2"]["path_exists"])
|
||||
assert.Equal(t, true, byID["a2"]["path_is_dir"])
|
||||
_, hasMissing := byID["a3"]
|
||||
assert.False(t, hasMissing)
|
||||
}
|
||||
|
||||
func TestGetWorkflowVerbose(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows/:name", GetWorkflowVerbose(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows/scan-module?json=true", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "scan-module", result["name"])
|
||||
assert.Equal(t, "module", result["kind"])
|
||||
|
||||
// Check params
|
||||
params, ok := result["params"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Len(t, params, 1)
|
||||
|
||||
// Check steps
|
||||
steps, ok := result["steps"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Len(t, steps, 2)
|
||||
|
||||
// Check triggers
|
||||
triggers, ok := result["triggers"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Len(t, triggers, 2)
|
||||
}
|
||||
|
||||
func TestGetWorkflowVerbose_Flow(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows/:name", GetWorkflowVerbose(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows/test-flow?json=true", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "test-flow", result["name"])
|
||||
assert.Equal(t, "flow", result["kind"])
|
||||
|
||||
// Check modules
|
||||
modules, ok := result["modules"].([]interface{})
|
||||
require.True(t, ok)
|
||||
assert.Len(t, modules, 1)
|
||||
}
|
||||
|
||||
func TestValidateWorkflow(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows/:name/validate", ValidateWorkflow(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows/test-module/validate", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, true, result["valid"])
|
||||
}
|
||||
|
||||
func TestValidateWorkflow_NotFound(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/workflows/:name/validate", ValidateWorkflow(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/workflows/nonexistent/validate", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusNotFound, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestReloadWorkflows(t *testing.T) {
|
||||
cfg, _ := setupTestWorkflowDir(t)
|
||||
|
||||
app := fiber.New()
|
||||
app.Post("/workflows/reload", ReloadWorkflows(cfg))
|
||||
|
||||
req := httptest.NewRequest("POST", "/workflows/reload", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Contains(t, result["message"], "reloaded")
|
||||
}
|
||||
|
||||
func TestGetSettings(t *testing.T) {
|
||||
cfg := &config.Config{
|
||||
BaseFolder: "/test/base",
|
||||
Server: config.ServerConfig{
|
||||
Host: "localhost",
|
||||
Port: 8811,
|
||||
},
|
||||
}
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/settings", GetSettings(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/settings", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "/test/base", result["base_folder"])
|
||||
assert.Equal(t, core.VERSION, result["version"])
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
)
|
||||
|
||||
// ServerInfoData holds cached server info set once at startup
|
||||
type ServerInfoData struct {
|
||||
License string
|
||||
Version string
|
||||
Binary string
|
||||
Repo string
|
||||
Author string
|
||||
Docs string
|
||||
}
|
||||
|
||||
// cachedServerInfo holds server info set once at startup
|
||||
var cachedServerInfo *ServerInfoData
|
||||
|
||||
// SetServerInfo sets the cached server info (called once at server startup)
|
||||
func SetServerInfo(info *ServerInfoData) {
|
||||
cachedServerInfo = info
|
||||
}
|
||||
|
||||
// HealthCheck handles health check requests
|
||||
// @Summary Health check
|
||||
// @Description Check if the server is running
|
||||
// @Tags Health
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]string "status: ok"
|
||||
// @Router /health [get]
|
||||
func HealthCheck(c *fiber.Ctx) error {
|
||||
return c.JSON(fiber.Map{
|
||||
"status": "ok",
|
||||
})
|
||||
}
|
||||
|
||||
// ReadinessCheck handles readiness check requests
|
||||
// @Summary Readiness check
|
||||
// @Description Check if the server is ready to accept requests
|
||||
// @Tags Health
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]string "status: ready"
|
||||
// @Router /health/ready [get]
|
||||
func ReadinessCheck(c *fiber.Ctx) error {
|
||||
// Check database connection, etc.
|
||||
return c.JSON(fiber.Map{
|
||||
"status": "ready",
|
||||
})
|
||||
}
|
||||
|
||||
// Root handles the root endpoint showing version info
|
||||
// @Summary Server info
|
||||
// @Description Get server version and info
|
||||
// @Tags Info
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]string "Server information"
|
||||
// @Router / [get]
|
||||
func Root(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
return c.JSON(fiber.Map{
|
||||
"message": fmt.Sprintf("Oh dear me, how delightful to notice you're taking a look at this! I'm ever so pleased to let you know that %s is ticking along quite nicely, thank you.", cachedServerInfo.Binary),
|
||||
"version": cachedServerInfo.Version,
|
||||
"repo": cachedServerInfo.Repo,
|
||||
"author": cachedServerInfo.Author,
|
||||
"docs": cachedServerInfo.Docs,
|
||||
"license": cachedServerInfo.License,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ServerInfo handles the /server-info endpoint
|
||||
// @Summary Server info JSON
|
||||
// @Description Get server version and info in JSON
|
||||
// @Tags Info
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]string "Server information"
|
||||
// @Router /server-info [get]
|
||||
func ServerInfo(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
return c.JSON(fiber.Map{
|
||||
"message": fmt.Sprintf("Oh dear me, how delightful to notice you're taking a look at this! I'm ever so pleased to let you know that %s is ticking along quite nicely, thank you.", cachedServerInfo.Binary),
|
||||
"version": cachedServerInfo.Version,
|
||||
"repo": cachedServerInfo.Repo,
|
||||
"author": cachedServerInfo.Author,
|
||||
"docs": cachedServerInfo.Docs,
|
||||
"license": cachedServerInfo.License,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,372 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/installer"
|
||||
"github.com/j3ssie/osmedeus/v5/public"
|
||||
)
|
||||
|
||||
// GetRegistryInfo returns binary registry with installation status
|
||||
// @Summary Get registry info
|
||||
// @Description Get binary registry with mode support (direct-fetch or nix-build)
|
||||
// @Tags Install
|
||||
// @Produce json
|
||||
// @Param registry_mode query string false "Registry mode: direct-fetch or nix-build" default(direct-fetch)
|
||||
// @Success 200 {object} map[string]interface{} "Registry data"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to load registry"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/registry-info [get]
|
||||
func GetRegistryInfo(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
registryMode := c.Query("registry_mode", "direct-fetch")
|
||||
|
||||
switch registryMode {
|
||||
case "nix-build":
|
||||
return getNixBuildRegistry(c)
|
||||
case "direct-fetch":
|
||||
fallthrough
|
||||
default:
|
||||
return getDirectFetchRegistry(c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// getDirectFetchRegistry returns the direct-fetch registry (existing behavior)
|
||||
func getDirectFetchRegistry(c *fiber.Ctx) error {
|
||||
registry, err := installer.LoadRegistry("", nil)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to load registry: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Build response with installation status for each binary
|
||||
binariesWithStatus := make(map[string]BinaryStatusEntry)
|
||||
for name, entry := range registry {
|
||||
path, _ := exec.LookPath(name)
|
||||
binariesWithStatus[name] = BinaryStatusEntry{
|
||||
Desc: entry.Desc,
|
||||
RepoLink: entry.RepoLink,
|
||||
Version: entry.Version,
|
||||
Tags: entry.Tags,
|
||||
ValidateCommand: entry.ValidateCommand,
|
||||
Linux: entry.Linux,
|
||||
Darwin: entry.Darwin,
|
||||
Windows: entry.Windows,
|
||||
CommandLinux: entry.CommandLinux,
|
||||
CommandDarwin: entry.CommandDarwin,
|
||||
CommandDual: entry.CommandDual,
|
||||
MultiCommandsLinux: entry.MultiCommandsLinux,
|
||||
MultiCommandsDarwin: entry.MultiCommandsDarwin,
|
||||
Installed: installer.IsBinaryInstalled(name, &entry),
|
||||
Path: path,
|
||||
}
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"registry_mode": "direct-fetch",
|
||||
"registry_url": installer.DefaultRegistryURL,
|
||||
"binaries": binariesWithStatus,
|
||||
})
|
||||
}
|
||||
|
||||
// getNixBuildRegistry returns Nix flake binaries with registry metadata
|
||||
func getNixBuildRegistry(c *fiber.Ctx) error {
|
||||
// Parse flake.nix
|
||||
flakeContent, err := public.GetFlakeNix()
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to read flake.nix: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
categories, err := installer.ParseFlakeNixBinariesFromString(string(flakeContent))
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to parse flake.nix: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Load registry for metadata (desc, tags)
|
||||
registry, _ := installer.LoadRegistry("", nil)
|
||||
|
||||
// Build response with categories and tool metadata
|
||||
categoriesData := make([]map[string]interface{}, 0)
|
||||
for _, cat := range categories {
|
||||
toolsData := make([]map[string]interface{}, 0)
|
||||
for _, tool := range cat.Tools {
|
||||
// Get entry for validation command check
|
||||
var entryPtr *installer.BinaryEntry
|
||||
if registry != nil {
|
||||
if entry, ok := registry[tool]; ok {
|
||||
entryPtr = &entry
|
||||
}
|
||||
}
|
||||
|
||||
toolData := map[string]interface{}{
|
||||
"name": tool,
|
||||
"installed": installer.IsBinaryInstalled(tool, entryPtr),
|
||||
}
|
||||
if entryPtr != nil {
|
||||
toolData["desc"] = entryPtr.Desc
|
||||
toolData["tags"] = entryPtr.Tags
|
||||
toolData["version"] = entryPtr.Version
|
||||
toolData["repo_link"] = entryPtr.RepoLink
|
||||
if entryPtr.ValidateCommand != "" {
|
||||
toolData["valide-command"] = entryPtr.ValidateCommand
|
||||
}
|
||||
}
|
||||
if path, err := exec.LookPath(tool); err == nil {
|
||||
toolData["path"] = path
|
||||
}
|
||||
toolsData = append(toolsData, toolData)
|
||||
}
|
||||
catData := map[string]interface{}{
|
||||
"name": cat.Name,
|
||||
"tools": toolsData,
|
||||
}
|
||||
categoriesData = append(categoriesData, catData)
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"registry_mode": "nix-build",
|
||||
"nix_installed": installer.IsNixInstalled(),
|
||||
"categories": categoriesData,
|
||||
})
|
||||
}
|
||||
|
||||
// RegistryInstall handles binary or workflow installation via API
|
||||
// @Summary Install binaries or workflows
|
||||
// @Description Install binaries from registry or workflows from git/zip URL. Supports direct-fetch and nix-build modes.
|
||||
// @Tags Install
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body InstallRequest true "Installation configuration"
|
||||
// @Success 200 {object} map[string]interface{} "Installation result"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 500 {object} map[string]interface{} "Installation failed"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/registry-install [post]
|
||||
func RegistryInstall(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req InstallRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
// Default registry mode
|
||||
if req.RegistryMode == "" {
|
||||
req.RegistryMode = "direct-fetch"
|
||||
}
|
||||
|
||||
// Validate type
|
||||
if req.Type != "binary" && req.Type != "workflow" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "type must be 'binary' or 'workflow'",
|
||||
})
|
||||
}
|
||||
|
||||
inst := installer.NewInstaller(cfg.BaseFolder, cfg.WorkflowsPath, cfg.BinariesPath, nil)
|
||||
|
||||
switch req.Type {
|
||||
case "binary":
|
||||
return installBinaries(c, cfg, inst, req)
|
||||
case "workflow":
|
||||
return installWorkflow(c, inst, req)
|
||||
default:
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid installation type",
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// installBinaries handles binary installation with mode support
|
||||
func installBinaries(c *fiber.Ctx, cfg *config.Config, inst *installer.Installer, req InstallRequest) error {
|
||||
switch req.RegistryMode {
|
||||
case "nix-build":
|
||||
return installBinariesViaNix(c, cfg, inst, req)
|
||||
case "direct-fetch":
|
||||
fallthrough
|
||||
default:
|
||||
return installBinariesDirectFetch(c, inst, req)
|
||||
}
|
||||
}
|
||||
|
||||
// installBinariesDirectFetch handles binary installation via direct download
|
||||
func installBinariesDirectFetch(c *fiber.Ctx, inst *installer.Installer, req InstallRequest) error {
|
||||
registryURL := req.RegistryURL
|
||||
if registryURL == "" {
|
||||
registryURL = installer.DefaultRegistryURL
|
||||
}
|
||||
|
||||
// Load registry first
|
||||
registry, err := installer.LoadRegistry(registryURL, nil)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to load registry: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
var installed []string
|
||||
var failed []map[string]string
|
||||
|
||||
if req.InstallAll {
|
||||
// Install all binaries from registry
|
||||
for name := range registry {
|
||||
if err := installer.InstallBinary(name, registry, inst.BinariesFolder, nil); err != nil {
|
||||
failed = append(failed, map[string]string{
|
||||
"name": name,
|
||||
"error": err.Error(),
|
||||
})
|
||||
} else {
|
||||
installed = append(installed, name)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Install specified binaries
|
||||
if len(req.Names) == 0 {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "names array is required for binary installation (or set install_all=true)",
|
||||
})
|
||||
}
|
||||
|
||||
for _, name := range req.Names {
|
||||
if err := installer.InstallBinary(name, registry, inst.BinariesFolder, nil); err != nil {
|
||||
failed = append(failed, map[string]string{
|
||||
"name": name,
|
||||
"error": err.Error(),
|
||||
})
|
||||
} else {
|
||||
installed = append(installed, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
response := fiber.Map{
|
||||
"message": "Binary installation completed",
|
||||
"registry_mode": "direct-fetch",
|
||||
"installed": installed,
|
||||
"installed_count": len(installed),
|
||||
"binaries_folder": inst.BinariesFolder,
|
||||
}
|
||||
|
||||
if len(failed) > 0 {
|
||||
response["failed"] = failed
|
||||
response["failed_count"] = len(failed)
|
||||
}
|
||||
|
||||
return c.JSON(response)
|
||||
}
|
||||
|
||||
// installBinariesViaNix handles binary installation via Nix
|
||||
func installBinariesViaNix(c *fiber.Ctx, cfg *config.Config, _ *installer.Installer, req InstallRequest) error {
|
||||
// Check if Nix is installed
|
||||
if !installer.IsNixInstalled() {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Nix is not installed. Install Nix first or use registry_mode=direct-fetch",
|
||||
})
|
||||
}
|
||||
|
||||
// Get binaries folder
|
||||
binariesFolder := cfg.BinariesPath
|
||||
if binariesFolder == "" {
|
||||
binariesFolder = filepath.Join(cfg.BaseFolder, "binaries")
|
||||
}
|
||||
|
||||
var names []string
|
||||
if req.InstallAll {
|
||||
// Get all binaries from flake
|
||||
flakeContent, err := public.GetFlakeNix()
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to read flake.nix: " + err.Error(),
|
||||
})
|
||||
}
|
||||
categories, err := installer.ParseFlakeNixBinariesFromString(string(flakeContent))
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to parse flake.nix: " + err.Error(),
|
||||
})
|
||||
}
|
||||
names = installer.GetAllFlakeBinaries(categories)
|
||||
} else {
|
||||
if len(req.Names) == 0 {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "names array is required for binary installation (or set install_all=true)",
|
||||
})
|
||||
}
|
||||
names = req.Names
|
||||
}
|
||||
|
||||
// Install each binary via Nix
|
||||
var installed []string
|
||||
var failed []map[string]string
|
||||
|
||||
for _, name := range names {
|
||||
if err := installer.InstallBinaryViaNix(name, "", binariesFolder); err != nil {
|
||||
failed = append(failed, map[string]string{
|
||||
"name": name,
|
||||
"error": err.Error(),
|
||||
})
|
||||
} else {
|
||||
installed = append(installed, name)
|
||||
}
|
||||
}
|
||||
|
||||
response := fiber.Map{
|
||||
"message": "Nix binary installation completed",
|
||||
"registry_mode": "nix-build",
|
||||
"installed": installed,
|
||||
"installed_count": len(installed),
|
||||
"binaries_folder": binariesFolder,
|
||||
}
|
||||
|
||||
if len(failed) > 0 {
|
||||
response["failed"] = failed
|
||||
response["failed_count"] = len(failed)
|
||||
}
|
||||
|
||||
return c.JSON(response)
|
||||
}
|
||||
|
||||
// installWorkflow handles workflow installation
|
||||
func installWorkflow(c *fiber.Ctx, inst *installer.Installer, req InstallRequest) error {
|
||||
if req.Source == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "source is required for workflow installation (git URL, zip URL, or local path)",
|
||||
})
|
||||
}
|
||||
|
||||
if err := inst.InstallWorkflow(req.Source); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to install workflow: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Workflow installed successfully",
|
||||
"source": req.Source,
|
||||
"workflow_folder": inst.WorkflowFolder,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
// JobStatus represents the aggregated status of a job
|
||||
type JobStatus struct {
|
||||
JobID string `json:"job_id"`
|
||||
Status string `json:"status"` // pending, running, completed, failed, partial
|
||||
Runs []*database.Run `json:"runs"`
|
||||
Progress JobProgress `json:"progress"`
|
||||
}
|
||||
|
||||
// JobProgress represents progress statistics for a job
|
||||
type JobProgress struct {
|
||||
Total int `json:"total"`
|
||||
Pending int `json:"pending"`
|
||||
Running int `json:"running"`
|
||||
Completed int `json:"completed"`
|
||||
Failed int `json:"failed"`
|
||||
}
|
||||
|
||||
// GetJobStatus handles getting the status of a job (group of runs)
|
||||
// @Summary Get job status
|
||||
// @Description Get the aggregated status of a job and its runs
|
||||
// @Tags Jobs
|
||||
// @Produce json
|
||||
// @Param id path string true "Job ID"
|
||||
// @Success 200 {object} map[string]interface{} "Job status"
|
||||
// @Failure 404 {object} map[string]interface{} "Job not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/jobs/{id} [get]
|
||||
func GetJobStatus(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
jobID := c.Params("id")
|
||||
if jobID == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Job ID is required",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
runs, err := database.GetRunsByJobID(ctx, jobID)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
if len(runs) == 0 {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Job not found",
|
||||
})
|
||||
}
|
||||
|
||||
// Calculate progress
|
||||
progress := JobProgress{Total: len(runs)}
|
||||
for _, run := range runs {
|
||||
switch run.Status {
|
||||
case "pending":
|
||||
progress.Pending++
|
||||
case "running":
|
||||
progress.Running++
|
||||
case "completed":
|
||||
progress.Completed++
|
||||
case "failed":
|
||||
progress.Failed++
|
||||
}
|
||||
}
|
||||
|
||||
// Determine aggregate status
|
||||
status := aggregateStatus(progress)
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"job_id": jobID,
|
||||
"status": status,
|
||||
"runs": runs,
|
||||
"progress": progress,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// aggregateStatus determines the overall job status based on run statuses
|
||||
func aggregateStatus(progress JobProgress) string {
|
||||
if progress.Total == 0 {
|
||||
return "pending"
|
||||
}
|
||||
if progress.Running > 0 {
|
||||
return "running"
|
||||
}
|
||||
if progress.Pending > 0 && progress.Completed == 0 && progress.Failed == 0 {
|
||||
return "pending"
|
||||
}
|
||||
if progress.Completed == progress.Total {
|
||||
return "completed"
|
||||
}
|
||||
if progress.Failed == progress.Total {
|
||||
return "failed"
|
||||
}
|
||||
if progress.Failed > 0 || progress.Completed > 0 {
|
||||
return "partial" // some completed, some failed
|
||||
}
|
||||
return "pending"
|
||||
}
|
||||
@@ -0,0 +1,507 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/core"
|
||||
)
|
||||
|
||||
// LLMChatRequest represents a direct LLM chat completion request
|
||||
type LLMChatRequest struct {
|
||||
Messages []core.LLMMessage `json:"messages"`
|
||||
Model string `json:"model,omitempty"`
|
||||
MaxTokens int `json:"max_tokens,omitempty"`
|
||||
Temperature *float64 `json:"temperature,omitempty"`
|
||||
TopP *float64 `json:"top_p,omitempty"`
|
||||
TopK *int `json:"top_k,omitempty"`
|
||||
N int `json:"n,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
Tools []core.LLMTool `json:"tools,omitempty"`
|
||||
ToolChoice interface{} `json:"tool_choice,omitempty"`
|
||||
ResponseFormat *core.LLMResponseFormat `json:"response_format,omitempty"`
|
||||
}
|
||||
|
||||
// LLMChatResponse represents the chat completion response
|
||||
type LLMChatResponse struct {
|
||||
ID string `json:"id"`
|
||||
Model string `json:"model"`
|
||||
Content interface{} `json:"content"`
|
||||
FinishReason string `json:"finish_reason"`
|
||||
ToolCalls []core.LLMToolCall `json:"tool_calls,omitempty"`
|
||||
Usage map[string]int `json:"usage"`
|
||||
}
|
||||
|
||||
// LLMEmbeddingRequest represents an embedding request
|
||||
type LLMEmbeddingRequest struct {
|
||||
Input []string `json:"input"`
|
||||
Model string `json:"model,omitempty"`
|
||||
}
|
||||
|
||||
// LLMEmbeddingResponse represents the embedding response
|
||||
type LLMEmbeddingResponse struct {
|
||||
Model string `json:"model"`
|
||||
Embeddings [][]float64 `json:"embeddings"`
|
||||
Usage map[string]int `json:"usage"`
|
||||
}
|
||||
|
||||
// Internal types for API communication
|
||||
type llmChatAPIRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []llmChatMessage `json:"messages"`
|
||||
MaxTokens int `json:"max_tokens,omitempty"`
|
||||
Temperature float64 `json:"temperature,omitempty"`
|
||||
TopP float64 `json:"top_p,omitempty"`
|
||||
TopK int `json:"top_k,omitempty"`
|
||||
N int `json:"n,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
Tools []core.LLMTool `json:"tools,omitempty"`
|
||||
ToolChoice interface{} `json:"tool_choice,omitempty"`
|
||||
ResponseFormat *core.LLMResponseFormat `json:"response_format,omitempty"`
|
||||
}
|
||||
|
||||
type llmChatMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content interface{} `json:"content"`
|
||||
Name string `json:"name,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
ToolCalls []core.LLMToolCall `json:"tool_calls,omitempty"`
|
||||
}
|
||||
|
||||
type llmChatAPIResponse struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
Created int64 `json:"created"`
|
||||
Model string `json:"model"`
|
||||
Choices []llmAPIChoice `json:"choices"`
|
||||
Usage llmAPIUsage `json:"usage"`
|
||||
Error *llmAPIError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type llmAPIChoice struct {
|
||||
Index int `json:"index"`
|
||||
Message llmChatMessage `json:"message"`
|
||||
FinishReason string `json:"finish_reason"`
|
||||
}
|
||||
|
||||
type llmAPIUsage struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
CompletionTokens int `json:"completion_tokens"`
|
||||
TotalTokens int `json:"total_tokens"`
|
||||
}
|
||||
|
||||
type llmAPIError struct {
|
||||
Message string `json:"message"`
|
||||
Type string `json:"type"`
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
type llmEmbeddingAPIRequest struct {
|
||||
Model string `json:"model"`
|
||||
Input []string `json:"input"`
|
||||
EncodingFormat string `json:"encoding_format,omitempty"`
|
||||
}
|
||||
|
||||
type llmEmbeddingAPIResponse struct {
|
||||
Object string `json:"object"`
|
||||
Data []llmEmbeddingData `json:"data"`
|
||||
Model string `json:"model"`
|
||||
Usage llmEmbeddingAPIUsage `json:"usage"`
|
||||
Error *llmAPIError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type llmEmbeddingData struct {
|
||||
Object string `json:"object"`
|
||||
Embedding []float64 `json:"embedding"`
|
||||
Index int `json:"index"`
|
||||
}
|
||||
|
||||
type llmEmbeddingAPIUsage struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
TotalTokens int `json:"total_tokens"`
|
||||
}
|
||||
|
||||
// LLMChat handles direct LLM chat completion requests
|
||||
// @Summary LLM Chat Completion
|
||||
// @Description Send a chat completion request to the configured LLM provider (OpenAI-compatible)
|
||||
// @Tags LLM
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body LLMChatRequest true "Chat request"
|
||||
// @Success 200 {object} LLMChatResponse "Chat response"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 500 {object} map[string]interface{} "LLM error"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/llm/v1/chat/completions [post]
|
||||
func LLMChat(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req LLMChatRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Validate messages
|
||||
if len(req.Messages) == 0 {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Messages field is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate LLM configuration
|
||||
if cfg.LLM.GetProviderCount() == 0 {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "No LLM providers configured",
|
||||
})
|
||||
}
|
||||
|
||||
// Get provider
|
||||
provider := cfg.LLM.GetCurrentProvider()
|
||||
if provider == nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "No LLM provider available",
|
||||
})
|
||||
}
|
||||
|
||||
// Build API request
|
||||
apiReq := buildLLMChatAPIRequest(&req, cfg, provider)
|
||||
|
||||
// Execute request with retry
|
||||
ctx, cancel := context.WithTimeout(c.Context(), 120*time.Second)
|
||||
defer cancel()
|
||||
|
||||
response, err := executeLLMChatRequest(ctx, cfg, provider, apiReq)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "LLM request failed: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
if response.Error != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": fmt.Sprintf("LLM API error: %s (%s)", response.Error.Message, response.Error.Type),
|
||||
})
|
||||
}
|
||||
|
||||
// Build response
|
||||
result := &LLMChatResponse{
|
||||
ID: response.ID,
|
||||
Model: response.Model,
|
||||
Usage: map[string]int{
|
||||
"prompt_tokens": response.Usage.PromptTokens,
|
||||
"completion_tokens": response.Usage.CompletionTokens,
|
||||
"total_tokens": response.Usage.TotalTokens,
|
||||
},
|
||||
}
|
||||
|
||||
if len(response.Choices) > 0 {
|
||||
choice := response.Choices[0]
|
||||
result.Content = choice.Message.Content
|
||||
result.FinishReason = choice.FinishReason
|
||||
if len(choice.Message.ToolCalls) > 0 {
|
||||
result.ToolCalls = choice.Message.ToolCalls
|
||||
}
|
||||
}
|
||||
|
||||
return c.JSON(result)
|
||||
}
|
||||
}
|
||||
|
||||
// LLMEmbedding handles embedding generation requests
|
||||
// @Summary Generate Embeddings
|
||||
// @Description Generate embeddings for input text using the configured LLM provider
|
||||
// @Tags LLM
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body LLMEmbeddingRequest true "Embedding request"
|
||||
// @Success 200 {object} LLMEmbeddingResponse "Embedding response"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 500 {object} map[string]interface{} "LLM error"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/llm/v1/embeddings [post]
|
||||
func LLMEmbedding(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req LLMEmbeddingRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Validate input
|
||||
if len(req.Input) == 0 {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Input field is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate LLM configuration
|
||||
if cfg.LLM.GetProviderCount() == 0 {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "No LLM providers configured",
|
||||
})
|
||||
}
|
||||
|
||||
// Get provider
|
||||
provider := cfg.LLM.GetCurrentProvider()
|
||||
if provider == nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "No LLM provider available",
|
||||
})
|
||||
}
|
||||
|
||||
// Build API request
|
||||
apiReq := &llmEmbeddingAPIRequest{
|
||||
Model: req.Model,
|
||||
Input: req.Input,
|
||||
}
|
||||
if apiReq.Model == "" {
|
||||
apiReq.Model = provider.Model
|
||||
}
|
||||
|
||||
// Execute request
|
||||
ctx, cancel := context.WithTimeout(c.Context(), 120*time.Second)
|
||||
defer cancel()
|
||||
|
||||
response, err := executeLLMEmbeddingRequest(ctx, cfg, provider, apiReq)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Embedding request failed: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
if response.Error != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": fmt.Sprintf("Embedding API error: %s (%s)", response.Error.Message, response.Error.Type),
|
||||
})
|
||||
}
|
||||
|
||||
// Build response
|
||||
embeddings := make([][]float64, len(response.Data))
|
||||
for i, d := range response.Data {
|
||||
embeddings[i] = d.Embedding
|
||||
}
|
||||
|
||||
result := &LLMEmbeddingResponse{
|
||||
Model: response.Model,
|
||||
Embeddings: embeddings,
|
||||
Usage: map[string]int{
|
||||
"prompt_tokens": response.Usage.PromptTokens,
|
||||
"total_tokens": response.Usage.TotalTokens,
|
||||
},
|
||||
}
|
||||
|
||||
return c.JSON(result)
|
||||
}
|
||||
}
|
||||
|
||||
// buildLLMChatAPIRequest builds the API request from handler request
|
||||
func buildLLMChatAPIRequest(req *LLMChatRequest, cfg *config.Config, provider *config.LLMProvider) *llmChatAPIRequest {
|
||||
apiReq := &llmChatAPIRequest{
|
||||
Model: req.Model,
|
||||
MaxTokens: req.MaxTokens,
|
||||
N: req.N,
|
||||
Stream: req.Stream,
|
||||
Tools: req.Tools,
|
||||
ToolChoice: req.ToolChoice,
|
||||
ResponseFormat: req.ResponseFormat,
|
||||
}
|
||||
|
||||
// Use defaults from config if not specified
|
||||
if apiReq.Model == "" {
|
||||
apiReq.Model = provider.Model
|
||||
}
|
||||
if apiReq.MaxTokens == 0 {
|
||||
apiReq.MaxTokens = cfg.LLM.MaxTokens
|
||||
}
|
||||
|
||||
// Temperature - use request value, config default, or 0.7
|
||||
if req.Temperature != nil {
|
||||
apiReq.Temperature = *req.Temperature
|
||||
} else if cfg.LLM.Temperature > 0 {
|
||||
apiReq.Temperature = cfg.LLM.Temperature
|
||||
} else {
|
||||
apiReq.Temperature = 0.7
|
||||
}
|
||||
|
||||
// TopP
|
||||
if req.TopP != nil {
|
||||
apiReq.TopP = *req.TopP
|
||||
} else {
|
||||
apiReq.TopP = cfg.LLM.TopP
|
||||
}
|
||||
|
||||
// TopK
|
||||
if req.TopK != nil {
|
||||
apiReq.TopK = *req.TopK
|
||||
} else {
|
||||
apiReq.TopK = cfg.LLM.TopK
|
||||
}
|
||||
|
||||
// Convert messages
|
||||
messages := make([]llmChatMessage, 0, len(req.Messages)+1)
|
||||
|
||||
// Auto-prepend system prompt if configured and no system message exists
|
||||
if cfg.LLM.SystemPrompt != "" {
|
||||
hasSystem := false
|
||||
for _, msg := range req.Messages {
|
||||
if msg.Role == core.LLMRoleSystem {
|
||||
hasSystem = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !hasSystem {
|
||||
messages = append(messages, llmChatMessage{
|
||||
Role: string(core.LLMRoleSystem),
|
||||
Content: cfg.LLM.SystemPrompt,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
for _, msg := range req.Messages {
|
||||
messages = append(messages, llmChatMessage{
|
||||
Role: string(msg.Role),
|
||||
Content: msg.Content,
|
||||
Name: msg.Name,
|
||||
ToolCallID: msg.ToolCallID,
|
||||
ToolCalls: msg.ToolCalls,
|
||||
})
|
||||
}
|
||||
|
||||
apiReq.Messages = messages
|
||||
|
||||
return apiReq
|
||||
}
|
||||
|
||||
// executeLLMChatRequest executes an HTTP request to the LLM provider
|
||||
func executeLLMChatRequest(ctx context.Context, cfg *config.Config, provider *config.LLMProvider, apiReq *llmChatAPIRequest) (*llmChatAPIResponse, error) {
|
||||
body, err := json.Marshal(apiReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", provider.BaseURL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if provider.AuthToken != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+provider.AuthToken)
|
||||
}
|
||||
|
||||
// Parse custom headers
|
||||
if cfg.LLM.CustomHeaders != "" {
|
||||
for _, h := range strings.Split(cfg.LLM.CustomHeaders, ",") {
|
||||
if parts := strings.SplitN(strings.TrimSpace(h), ":", 2); len(parts) == 2 {
|
||||
req.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 120 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
var response llmChatAPIResponse
|
||||
if err := json.Unmarshal(respBody, &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse response: %w (body: %s)", err, string(respBody))
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
if response.Error != nil {
|
||||
return &response, fmt.Errorf("HTTP %d: %s", resp.StatusCode, response.Error.Message)
|
||||
}
|
||||
return &response, fmt.Errorf("HTTP %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
return &response, nil
|
||||
}
|
||||
|
||||
// executeLLMEmbeddingRequest executes an embedding request to the LLM provider
|
||||
func executeLLMEmbeddingRequest(ctx context.Context, cfg *config.Config, provider *config.LLMProvider, apiReq *llmEmbeddingAPIRequest) (*llmEmbeddingAPIResponse, error) {
|
||||
body, err := json.Marshal(apiReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
// Determine embedding endpoint
|
||||
embeddingURL := provider.BaseURL
|
||||
if strings.HasSuffix(embeddingURL, "/chat/completions") {
|
||||
embeddingURL = strings.Replace(embeddingURL, "/chat/completions", "/embeddings", 1)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", embeddingURL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if provider.AuthToken != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+provider.AuthToken)
|
||||
}
|
||||
|
||||
// Parse custom headers
|
||||
if cfg.LLM.CustomHeaders != "" {
|
||||
for _, h := range strings.Split(cfg.LLM.CustomHeaders, ",") {
|
||||
if parts := strings.SplitN(strings.TrimSpace(h), ":", 2); len(parts) == 2 {
|
||||
req.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 120 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
var response llmEmbeddingAPIResponse
|
||||
if err := json.Unmarshal(respBody, &response); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse response: %w (body: %s)", err, string(respBody))
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
if response.Error != nil {
|
||||
return &response, fmt.Errorf("HTTP %d: %s", resp.StatusCode, response.Error.Message)
|
||||
}
|
||||
return &response, fmt.Errorf("HTTP %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
return &response, nil
|
||||
}
|
||||
@@ -0,0 +1,548 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/google/uuid"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/core"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/executor"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/parser"
|
||||
)
|
||||
|
||||
// generateEmptyTarget creates a placeholder target name for empty_target mode
|
||||
func generateEmptyTarget() string {
|
||||
const chars = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||
random := make([]byte, 6)
|
||||
for i := range random {
|
||||
random[i] = chars[rand.Intn(len(chars))]
|
||||
}
|
||||
return fmt.Sprintf("empty-%s-%d", string(random), time.Now().Unix())
|
||||
}
|
||||
|
||||
// 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) {
|
||||
now := time.Now()
|
||||
runID := uuid.New().String()
|
||||
|
||||
paramsInterface := make(map[string]interface{})
|
||||
for k, v := range params {
|
||||
paramsInterface[k] = v
|
||||
}
|
||||
|
||||
run := &database.Run{
|
||||
ID: uuid.New().String(),
|
||||
RunID: runID,
|
||||
WorkflowName: workflow.Name,
|
||||
WorkflowKind: string(workflow.Kind),
|
||||
Target: target,
|
||||
Params: paramsInterface,
|
||||
Status: "running",
|
||||
TriggerType: triggerType,
|
||||
JobID: jobID,
|
||||
StartedAt: &now,
|
||||
TotalSteps: len(workflow.Steps),
|
||||
}
|
||||
|
||||
if err := database.CreateRun(ctx, run); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return run, nil
|
||||
}
|
||||
|
||||
// collectTargetsFromRequest collects targets from Target, Targets, and TargetFile fields
|
||||
func collectTargetsFromRequest(req *CreateRunRequest) ([]string, error) {
|
||||
var allTargets []string
|
||||
|
||||
// 1. Add single target if provided
|
||||
if req.Target != "" {
|
||||
allTargets = append(allTargets, req.Target)
|
||||
}
|
||||
|
||||
// 2. Add targets array if provided
|
||||
allTargets = append(allTargets, req.Targets...)
|
||||
|
||||
// 3. Read targets from file if provided
|
||||
if req.TargetFile != "" {
|
||||
fileTargets, err := readTargetsFromFile(req.TargetFile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read target file: %w", err)
|
||||
}
|
||||
allTargets = append(allTargets, fileTargets...)
|
||||
}
|
||||
|
||||
// Deduplicate and filter empty entries
|
||||
return deduplicateTargets(allTargets), nil
|
||||
}
|
||||
|
||||
// executeRunsConcurrently runs workflows for multiple targets with concurrency control
|
||||
func executeRunsConcurrently(
|
||||
workflow *core.Workflow,
|
||||
targets []string,
|
||||
baseParams map[string]string,
|
||||
cfg *config.Config,
|
||||
maxConcurrency int,
|
||||
isFlow bool,
|
||||
jobID string,
|
||||
) {
|
||||
if maxConcurrency <= 0 {
|
||||
maxConcurrency = 1
|
||||
}
|
||||
|
||||
sem := make(chan struct{}, maxConcurrency) // Semaphore
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for _, target := range targets {
|
||||
wg.Add(1)
|
||||
go func(t string) {
|
||||
defer wg.Done()
|
||||
|
||||
// Acquire semaphore
|
||||
sem <- struct{}{}
|
||||
defer func() { <-sem }()
|
||||
|
||||
// Clone params and set target
|
||||
targetParams := make(map[string]string)
|
||||
for k, v := range baseParams {
|
||||
targetParams[k] = v
|
||||
}
|
||||
targetParams["target"] = t
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create run record in database
|
||||
run, err := createRunRecord(ctx, cfg, workflow, t, targetParams, "api", jobID)
|
||||
var runID string
|
||||
if err == nil && run != nil {
|
||||
runID = run.RunID
|
||||
}
|
||||
|
||||
// Execute workflow
|
||||
exec := executor.NewExecutor()
|
||||
exec.SetServerMode(true) // Enable file logging for server mode
|
||||
|
||||
// Set up database progress tracking
|
||||
if runID != "" {
|
||||
exec.SetDBRunID(runID)
|
||||
exec.SetOnStepCompleted(func(stepCtx context.Context, dbRunID string) {
|
||||
_ = database.IncrementRunCompletedSteps(stepCtx, dbRunID)
|
||||
})
|
||||
}
|
||||
|
||||
var execErr error
|
||||
if isFlow && workflow.IsFlow() {
|
||||
_, execErr = exec.ExecuteFlow(ctx, workflow, targetParams, cfg)
|
||||
} else {
|
||||
_, execErr = exec.ExecuteModule(ctx, workflow, targetParams, cfg)
|
||||
}
|
||||
|
||||
// Update run status in database
|
||||
if runID != "" {
|
||||
if execErr != nil {
|
||||
_ = database.UpdateRunStatus(ctx, runID, "failed", execErr.Error())
|
||||
} else {
|
||||
_ = database.UpdateRunStatus(ctx, runID, "completed", "")
|
||||
}
|
||||
}
|
||||
}(target)
|
||||
}
|
||||
|
||||
// Wait for all runs to complete (optional - could also return immediately)
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// CreateRun handles run creation
|
||||
// @Summary Create a new run
|
||||
// @Description Execute a workflow against one or more targets. Supports multiple targets via array or file, concurrency control, priority levels, custom timeouts, runner configuration (host/docker/ssh), and scheduling via cron expressions.
|
||||
// @Tags Runs
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param run body CreateRunRequest true "Run configuration with optional priority, timeout, runner config, and scheduling"
|
||||
// @Success 202 {object} map[string]interface{} "Run started"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 404 {object} map[string]interface{} "Workflow not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/runs [post]
|
||||
func CreateRun(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req CreateRunRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
// Handle empty target mode
|
||||
if req.EmptyTarget {
|
||||
req.Target = generateEmptyTarget()
|
||||
}
|
||||
|
||||
// Collect all targets from Target, Targets, and TargetFile
|
||||
targets, err := collectTargetsFromRequest(&req)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
if len(targets) == 0 {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "At least one target is required (target, targets, target_file, or empty_target)",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate heuristics_check if provided
|
||||
if req.HeuristicsCheck != "" {
|
||||
validHeuristics := map[string]bool{"none": true, "basic": true, "advanced": true}
|
||||
if !validHeuristics[req.HeuristicsCheck] {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid heuristics_check value. Must be: none, basic, or advanced",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Determine workflow name and kind
|
||||
workflowName := req.Flow
|
||||
isFlow := true
|
||||
if workflowName == "" {
|
||||
workflowName = req.Module
|
||||
isFlow = false
|
||||
}
|
||||
|
||||
if workflowName == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Either 'flow' or 'module' is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Load workflow
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflow, err := loader.LoadWorkflow(workflowName)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workflow not found",
|
||||
})
|
||||
}
|
||||
|
||||
// Initialize params
|
||||
params := req.Params
|
||||
if params == nil {
|
||||
params = make(map[string]string)
|
||||
}
|
||||
|
||||
// Set default priority if not specified
|
||||
priority := req.Priority
|
||||
if priority == "" {
|
||||
priority = "medium"
|
||||
}
|
||||
|
||||
// Add runner configuration to params if specified
|
||||
if req.RunnerType != "" {
|
||||
params["runner_type"] = req.RunnerType
|
||||
}
|
||||
if req.DockerImage != "" {
|
||||
params["docker_image"] = req.DockerImage
|
||||
}
|
||||
if req.SSHHost != "" {
|
||||
params["ssh_host"] = req.SSHHost
|
||||
}
|
||||
if req.Timeout > 0 {
|
||||
params["timeout"] = fmt.Sprintf("%d", req.Timeout)
|
||||
}
|
||||
|
||||
// Add execution options to params
|
||||
if req.ThreadsHold > 0 {
|
||||
params["threads_hold"] = fmt.Sprintf("%d", req.ThreadsHold)
|
||||
}
|
||||
if req.HeuristicsCheck != "" {
|
||||
params["heuristics_check"] = req.HeuristicsCheck
|
||||
}
|
||||
if req.Repeat {
|
||||
params["repeat"] = "true"
|
||||
}
|
||||
if req.RepeatWaitTime != "" {
|
||||
params["repeat_wait_time"] = req.RepeatWaitTime
|
||||
}
|
||||
|
||||
// Set concurrency default
|
||||
concurrency := req.Concurrency
|
||||
if concurrency <= 0 {
|
||||
concurrency = 1
|
||||
}
|
||||
|
||||
cfgCopy := config.Get()
|
||||
|
||||
// Generate a job ID for grouping runs from this request
|
||||
jobID := uuid.New().String()[:8]
|
||||
|
||||
// Create run record(s) and execute
|
||||
var runIDs []string
|
||||
if len(targets) == 1 {
|
||||
// Single target - existing behavior
|
||||
params["target"] = targets[0]
|
||||
|
||||
// Create run record in database
|
||||
ctx := context.Background()
|
||||
run, _ := createRunRecord(ctx, cfgCopy, workflow, targets[0], params, "api", jobID)
|
||||
if run != nil {
|
||||
runIDs = append(runIDs, run.RunID)
|
||||
}
|
||||
|
||||
exec := executor.NewExecutor()
|
||||
exec.SetServerMode(true) // Enable file logging for server mode
|
||||
|
||||
// Set up database progress tracking
|
||||
if run != nil {
|
||||
exec.SetDBRunID(run.RunID)
|
||||
exec.SetOnStepCompleted(func(stepCtx context.Context, dbRunID string) {
|
||||
_ = database.IncrementRunCompletedSteps(stepCtx, dbRunID)
|
||||
})
|
||||
}
|
||||
|
||||
go func(runID string) {
|
||||
ctx := context.Background()
|
||||
var execErr error
|
||||
if isFlow && workflow.IsFlow() {
|
||||
_, execErr = exec.ExecuteFlow(ctx, workflow, params, cfgCopy)
|
||||
} else {
|
||||
_, execErr = exec.ExecuteModule(ctx, workflow, params, cfgCopy)
|
||||
}
|
||||
// Update run status in database
|
||||
if runID != "" {
|
||||
if execErr != nil {
|
||||
_ = database.UpdateRunStatus(ctx, runID, "failed", execErr.Error())
|
||||
} else {
|
||||
_ = database.UpdateRunStatus(ctx, runID, "completed", "")
|
||||
}
|
||||
}
|
||||
}(func() string {
|
||||
if run != nil {
|
||||
return run.RunID
|
||||
}
|
||||
return ""
|
||||
}())
|
||||
} else {
|
||||
// Multiple targets - concurrent execution
|
||||
go executeRunsConcurrently(workflow, targets, params, cfgCopy, concurrency, isFlow, jobID)
|
||||
}
|
||||
|
||||
// Build response
|
||||
response := fiber.Map{
|
||||
"message": "Run started",
|
||||
"workflow": workflow.Name,
|
||||
"kind": workflow.Kind,
|
||||
"target_count": len(targets),
|
||||
"priority": priority,
|
||||
"job_id": jobID,
|
||||
"status": "queued",
|
||||
"poll_url": fmt.Sprintf("/osm/api/jobs/%s", jobID),
|
||||
}
|
||||
|
||||
// For single target, include target field and run_id for backward compatibility
|
||||
if len(targets) == 1 {
|
||||
response["target"] = targets[0]
|
||||
if len(runIDs) > 0 {
|
||||
response["run_id"] = runIDs[0]
|
||||
}
|
||||
} else {
|
||||
response["targets"] = targets
|
||||
response["concurrency"] = concurrency
|
||||
}
|
||||
|
||||
// Add optional fields to response
|
||||
if req.RunnerType != "" {
|
||||
response["runner_type"] = req.RunnerType
|
||||
}
|
||||
if req.Timeout > 0 {
|
||||
response["timeout"] = req.Timeout
|
||||
}
|
||||
if req.Schedule != "" {
|
||||
response["schedule"] = req.Schedule
|
||||
response["schedule_enabled"] = req.ScheduleEnabled
|
||||
}
|
||||
if req.ThreadsHold > 0 {
|
||||
response["threads_hold"] = req.ThreadsHold
|
||||
}
|
||||
if req.EmptyTarget {
|
||||
response["empty_target"] = true
|
||||
}
|
||||
if req.HeuristicsCheck != "" {
|
||||
response["heuristics_check"] = req.HeuristicsCheck
|
||||
}
|
||||
if req.Repeat {
|
||||
response["repeat"] = true
|
||||
if req.RepeatWaitTime != "" {
|
||||
response["repeat_wait_time"] = req.RepeatWaitTime
|
||||
}
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusAccepted).JSON(response)
|
||||
}
|
||||
}
|
||||
|
||||
// ListRuns handles listing runs
|
||||
// @Summary List runs
|
||||
// @Description Get a paginated list of workflow runs with optional filters
|
||||
// @Tags Runs
|
||||
// @Produce json
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return" default(20)
|
||||
// @Param status query string false "Filter by status (pending, running, completed, failed, cancelled)"
|
||||
// @Param workflow query string false "Filter by workflow name"
|
||||
// @Param target query string false "Filter by target (partial match)"
|
||||
// @Success 200 {object} map[string]interface{} "List of runs"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/runs [get]
|
||||
func ListRuns(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
status := c.Query("status")
|
||||
workflow := c.Query("workflow")
|
||||
target := c.Query("target")
|
||||
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
result, err := database.ListRuns(ctx, offset, limit, status, workflow, target)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetRun handles getting a single run
|
||||
// @Summary Get run details
|
||||
// @Description Get details of a specific run by ID, including steps and artifacts
|
||||
// @Tags Runs
|
||||
// @Produce json
|
||||
// @Param id path string true "Run ID or RunID"
|
||||
// @Param include_steps query bool false "Include step results" default(false)
|
||||
// @Param include_artifacts query bool false "Include artifacts" default(false)
|
||||
// @Success 200 {object} map[string]interface{} "Run details"
|
||||
// @Failure 404 {object} map[string]interface{} "Run not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/runs/{id} [get]
|
||||
func GetRun(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
includeSteps := c.Query("include_steps") == "true"
|
||||
includeArtifacts := c.Query("include_artifacts") == "true"
|
||||
|
||||
ctx := context.Background()
|
||||
run, err := database.GetRunByID(ctx, id, includeSteps, includeArtifacts)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Run not found",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{"data": run})
|
||||
}
|
||||
}
|
||||
|
||||
// CancelRun handles cancelling a run
|
||||
// @Summary Cancel a run
|
||||
// @Description Cancel a running workflow execution
|
||||
// @Tags Runs
|
||||
// @Produce json
|
||||
// @Param id path string true "Run ID or RunID"
|
||||
// @Success 200 {object} map[string]interface{} "Run cancelled"
|
||||
// @Failure 404 {object} map[string]interface{} "Run not found"
|
||||
// @Failure 400 {object} map[string]interface{} "Run cannot be cancelled"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/runs/{id} [delete]
|
||||
func CancelRun(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
ctx := context.Background()
|
||||
|
||||
run, err := database.GetRunByID(ctx, id, false, false)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Run not found",
|
||||
})
|
||||
}
|
||||
|
||||
if run.Status != "pending" && run.Status != "running" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": fmt.Sprintf("Cannot cancel run with status '%s'", run.Status),
|
||||
})
|
||||
}
|
||||
|
||||
err = database.UpdateRunStatus(ctx, id, "cancelled", "Cancelled by user")
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to cancel run",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Run cancelled successfully",
|
||||
"id": run.ID,
|
||||
"run_id": run.RunID,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetRunSteps handles getting run steps
|
||||
func GetRunSteps(c *fiber.Ctx) error {
|
||||
// TODO: Implement with database
|
||||
return c.JSON(fiber.Map{
|
||||
"data": []interface{}{},
|
||||
})
|
||||
}
|
||||
|
||||
// GetRunArtifacts handles getting run artifacts
|
||||
func GetRunArtifacts(c *fiber.Ctx) error {
|
||||
// TODO: Implement with database
|
||||
return c.JSON(fiber.Map{
|
||||
"data": []interface{}{},
|
||||
})
|
||||
}
|
||||
|
||||
// DownloadArtifact handles artifact download
|
||||
func DownloadArtifact(c *fiber.Ctx) error {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Artifact not found",
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/executor"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/parser"
|
||||
)
|
||||
|
||||
// CreateSchedule handles creating a new schedule
|
||||
// @Summary Create a new schedule
|
||||
// @Description Create a scheduled workflow execution with cron expression
|
||||
// @Tags Schedules
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param schedule body CreateScheduleRequest true "Schedule configuration"
|
||||
// @Success 201 {object} map[string]interface{} "Schedule created"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules [post]
|
||||
func CreateSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req CreateScheduleRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate required fields
|
||||
if req.Name == "" || req.WorkflowName == "" || req.Schedule == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "name, workflow_name, and schedule are required",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create schedule record
|
||||
schedule, err := database.CreateSchedule(ctx, database.CreateScheduleInput{
|
||||
Name: req.Name,
|
||||
WorkflowName: req.WorkflowName,
|
||||
WorkflowKind: req.WorkflowKind,
|
||||
Target: req.Target,
|
||||
Schedule: req.Schedule,
|
||||
Enabled: req.Enabled,
|
||||
})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusCreated).JSON(fiber.Map{
|
||||
"message": "Schedule created",
|
||||
"data": schedule,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ListSchedules handles listing all schedules
|
||||
// @Summary List all schedules
|
||||
// @Description Get a paginated list of all scheduled workflows
|
||||
// @Tags Schedules
|
||||
// @Produce json
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return" default(20)
|
||||
// @Success 200 {object} map[string]interface{} "List of schedules"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules [get]
|
||||
func ListSchedules(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
result, err := database.ListSchedules(ctx, offset, limit)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetSchedule handles getting a single schedule
|
||||
// @Summary Get schedule details
|
||||
// @Description Get details of a specific schedule by ID
|
||||
// @Tags Schedules
|
||||
// @Produce json
|
||||
// @Param id path string true "Schedule ID"
|
||||
// @Success 200 {object} map[string]interface{} "Schedule details"
|
||||
// @Failure 404 {object} map[string]interface{} "Schedule not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules/{id} [get]
|
||||
func GetSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
|
||||
ctx := context.Background()
|
||||
schedule, err := database.GetScheduleByID(ctx, id)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Schedule not found",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{"data": schedule})
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateSchedule handles updating a schedule
|
||||
// @Summary Update a schedule
|
||||
// @Description Update an existing schedule
|
||||
// @Tags Schedules
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param id path string true "Schedule ID"
|
||||
// @Param schedule body UpdateScheduleRequest true "Schedule update data"
|
||||
// @Success 200 {object} map[string]interface{} "Schedule updated"
|
||||
// @Failure 404 {object} map[string]interface{} "Schedule not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules/{id} [put]
|
||||
func UpdateSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
|
||||
var req UpdateScheduleRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
schedule, err := database.UpdateSchedule(ctx, id, database.UpdateScheduleInput{
|
||||
Name: req.Name,
|
||||
Target: req.Target,
|
||||
Schedule: req.Schedule,
|
||||
Enabled: req.Enabled,
|
||||
})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Schedule updated",
|
||||
"data": schedule,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteSchedule handles deleting a schedule
|
||||
// @Summary Delete a schedule
|
||||
// @Description Delete a schedule by ID
|
||||
// @Tags Schedules
|
||||
// @Produce json
|
||||
// @Param id path string true "Schedule ID"
|
||||
// @Success 200 {object} map[string]interface{} "Schedule deleted"
|
||||
// @Failure 404 {object} map[string]interface{} "Schedule not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules/{id} [delete]
|
||||
func DeleteSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
|
||||
ctx := context.Background()
|
||||
if err := database.DeleteSchedule(ctx, id); err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{"message": "Schedule deleted"})
|
||||
}
|
||||
}
|
||||
|
||||
// EnableSchedule handles enabling a schedule
|
||||
// @Summary Enable a schedule
|
||||
// @Description Enable a disabled schedule
|
||||
// @Tags Schedules
|
||||
// @Produce json
|
||||
// @Param id path string true "Schedule ID"
|
||||
// @Success 200 {object} map[string]interface{} "Schedule enabled"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules/{id}/enable [post]
|
||||
func EnableSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
|
||||
ctx := context.Background()
|
||||
enabled := true
|
||||
_, err := database.UpdateSchedule(ctx, id, database.UpdateScheduleInput{Enabled: &enabled})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{"message": "Schedule enabled"})
|
||||
}
|
||||
}
|
||||
|
||||
// DisableSchedule handles disabling a schedule
|
||||
// @Summary Disable a schedule
|
||||
// @Description Disable an enabled schedule
|
||||
// @Tags Schedules
|
||||
// @Produce json
|
||||
// @Param id path string true "Schedule ID"
|
||||
// @Success 200 {object} map[string]interface{} "Schedule disabled"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules/{id}/disable [post]
|
||||
func DisableSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
|
||||
ctx := context.Background()
|
||||
enabled := false
|
||||
_, err := database.UpdateSchedule(ctx, id, database.UpdateScheduleInput{Enabled: &enabled})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{"message": "Schedule disabled"})
|
||||
}
|
||||
}
|
||||
|
||||
// TriggerSchedule handles manually triggering a scheduled workflow
|
||||
// @Summary Trigger a schedule
|
||||
// @Description Manually trigger a scheduled workflow execution
|
||||
// @Tags Schedules
|
||||
// @Produce json
|
||||
// @Param id path string true "Schedule ID"
|
||||
// @Success 202 {object} map[string]interface{} "Schedule triggered"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/schedules/{id}/trigger [post]
|
||||
func TriggerSchedule(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
id := c.Params("id")
|
||||
|
||||
ctx := context.Background()
|
||||
schedule, err := database.GetScheduleByID(ctx, id)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Schedule not found",
|
||||
})
|
||||
}
|
||||
|
||||
// Load and execute the workflow
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflow, err := loader.LoadWorkflow(schedule.WorkflowName)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to load workflow",
|
||||
})
|
||||
}
|
||||
|
||||
// Execute in background
|
||||
exec := executor.NewExecutor()
|
||||
cfgCopy := config.Get()
|
||||
go func() {
|
||||
bgCtx := context.Background()
|
||||
params := make(map[string]string)
|
||||
if schedule.InputConfig != nil {
|
||||
for k, v := range schedule.InputConfig {
|
||||
if s, ok := v.(string); ok {
|
||||
params[k] = s
|
||||
}
|
||||
}
|
||||
}
|
||||
if workflow.IsFlow() {
|
||||
_, _ = exec.ExecuteFlow(bgCtx, workflow, params, cfgCopy)
|
||||
} else {
|
||||
_, _ = exec.ExecuteModule(bgCtx, workflow, params, cfgCopy)
|
||||
}
|
||||
}()
|
||||
|
||||
// Update last run time
|
||||
_ = database.UpdateScheduleLastRun(ctx, id)
|
||||
|
||||
return c.Status(fiber.StatusAccepted).JSON(fiber.Map{
|
||||
"message": "Schedule triggered",
|
||||
"schedule": schedule.Name,
|
||||
"workflow": schedule.WorkflowName,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/core"
|
||||
)
|
||||
|
||||
// GetSettings returns basic settings
|
||||
func GetSettings(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
return c.JSON(fiber.Map{
|
||||
"base_folder": cfg.BaseFolder,
|
||||
"server": fiber.Map{
|
||||
"host": cfg.Server.Host,
|
||||
"port": cfg.Server.Port,
|
||||
},
|
||||
"version": core.VERSION,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// // UpdateSettings handles settings update
|
||||
// func UpdateSettings(c *fiber.Ctx) error {
|
||||
// return c.JSON(fiber.Map{"message": "Settings updated"})
|
||||
// }
|
||||
|
||||
// GetSettingsYAML returns the entire YAML configuration with sensitive fields redacted
|
||||
// @Summary Get YAML configuration
|
||||
// @Description Returns the entire configuration file with sensitive fields redacted
|
||||
// @Tags Settings
|
||||
// @Produce text/yaml
|
||||
// @Success 200 {string} string "YAML configuration content"
|
||||
// @Failure 500 {object} map[string]interface{} "Internal server error"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/settings/yaml [get]
|
||||
func GetSettingsYAML(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Read the config file
|
||||
settingsPath := filepath.Join(cfg.BaseFolder, "osm-settings.yaml")
|
||||
content, err := os.ReadFile(settingsPath)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": fmt.Sprintf("Failed to read config file: %v", err),
|
||||
})
|
||||
}
|
||||
|
||||
// Redact sensitive fields
|
||||
redactedContent := redactSensitiveFields(string(content))
|
||||
|
||||
// Return as YAML
|
||||
c.Set("Content-Type", "text/yaml")
|
||||
return c.SendString(redactedContent)
|
||||
}
|
||||
}
|
||||
|
||||
// redactSensitiveFields redacts values of sensitive fields in YAML content
|
||||
// Fields containing: _key, secret, password, username, _token (case-insensitive)
|
||||
func redactSensitiveFields(content string) string {
|
||||
// Pattern matches YAML key-value pairs where key contains sensitive patterns
|
||||
// Handles both quoted and unquoted values, and preserves comments
|
||||
sensitivePatterns := []string{
|
||||
`_key`,
|
||||
`secret`,
|
||||
`password`,
|
||||
`username`,
|
||||
`_token`,
|
||||
}
|
||||
|
||||
// Build regex pattern for sensitive field names
|
||||
patternStr := `(?i)^(\s*)([\w-]*(?:` + strings.Join(sensitivePatterns, "|") + `)[\w-]*):\s*(.+)$`
|
||||
re := regexp.MustCompile(patternStr)
|
||||
|
||||
lines := strings.Split(content, "\n")
|
||||
for i, line := range lines {
|
||||
// Skip comments and empty lines
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if line matches sensitive pattern
|
||||
if matches := re.FindStringSubmatch(line); matches != nil {
|
||||
indent := matches[1]
|
||||
key := matches[2]
|
||||
value := matches[3]
|
||||
|
||||
// Don't redact if value is already empty or a placeholder
|
||||
if value == `""` || value == "''" || value == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// Redact the value
|
||||
lines[i] = fmt.Sprintf("%s%s: \"[REDACTED]\"", indent, key)
|
||||
}
|
||||
}
|
||||
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
// // UpdateSettingsYAML replaces the entire YAML configuration
|
||||
// // @Summary Update YAML configuration
|
||||
// // @Description Replaces the entire configuration file with the provided YAML content
|
||||
// // @Tags Settings
|
||||
// // @Accept text/yaml
|
||||
// // @Produce json
|
||||
// // @Param config body string true "YAML configuration content"
|
||||
// // @Success 200 {object} map[string]interface{} "Configuration updated successfully"
|
||||
// // @Failure 400 {object} map[string]interface{} "Invalid YAML"
|
||||
// // @Failure 500 {object} map[string]interface{} "Internal server error"
|
||||
// // @Security BearerAuth
|
||||
// // @Router /osm/api/settings/yaml [put]
|
||||
// func UpdateSettingsYAML(cfg *config.Config) fiber.Handler {
|
||||
// return func(c *fiber.Ctx) error {
|
||||
// return c.Status(fiber.StatusMethodNotAllowed).JSON(fiber.Map{
|
||||
// "error": true,
|
||||
// "message": "Updating osm-settings.yaml via API is disabled",
|
||||
// })
|
||||
// }
|
||||
// }
|
||||
@@ -0,0 +1,246 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/snapshot"
|
||||
)
|
||||
|
||||
// SnapshotExportRequest represents the request body for snapshot export
|
||||
type SnapshotExportRequest struct {
|
||||
Workspace string `json:"workspace"`
|
||||
}
|
||||
|
||||
// SnapshotImportURLRequest represents the request body for import via URL
|
||||
type SnapshotImportURLRequest struct {
|
||||
URL string `json:"url"`
|
||||
}
|
||||
|
||||
// ListSnapshots handles listing available snapshots
|
||||
// @Summary List snapshots
|
||||
// @Description Get a list of available snapshot files in the snapshot directory
|
||||
// @Tags Snapshots
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]interface{} "List of snapshots"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to list snapshots"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/snapshots [get]
|
||||
func ListSnapshots(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
snapshots, err := snapshot.ListSnapshots(cfg.SnapshotPath)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to list snapshots: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": snapshots,
|
||||
"count": len(snapshots),
|
||||
"path": cfg.SnapshotPath,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// SnapshotExport handles exporting a workspace to a snapshot
|
||||
// @Summary Export workspace snapshot
|
||||
// @Description Export a workspace to a compressed zip archive and download it
|
||||
// @Tags Snapshots
|
||||
// @Accept json
|
||||
// @Produce application/zip
|
||||
// @Param body body SnapshotExportRequest true "Workspace to export"
|
||||
// @Success 200 {file} binary "Snapshot zip file"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 404 {object} map[string]interface{} "Workspace not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to create snapshot"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/snapshots/export [post]
|
||||
func SnapshotExport(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
var req SnapshotExportRequest
|
||||
if err := c.BodyParser(&req); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
if req.Workspace == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workspace name is required",
|
||||
})
|
||||
}
|
||||
|
||||
workspacePath := filepath.Join(cfg.WorkspacesPath, req.Workspace)
|
||||
if _, err := os.Stat(workspacePath); os.IsNotExist(err) {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workspace not found: " + req.Workspace,
|
||||
})
|
||||
}
|
||||
|
||||
// Ensure snapshot directory exists
|
||||
if err := os.MkdirAll(cfg.SnapshotPath, 0755); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to create snapshot directory",
|
||||
})
|
||||
}
|
||||
|
||||
// Generate output path
|
||||
zipFilename := fmt.Sprintf("%s_%d.zip", req.Workspace, time.Now().Unix())
|
||||
outputPath := filepath.Join(cfg.SnapshotPath, zipFilename)
|
||||
|
||||
// Export workspace
|
||||
result, err := snapshot.ExportWorkspace(workspacePath, outputPath)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to create snapshot: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Set headers for file download
|
||||
c.Set("Content-Disposition", fmt.Sprintf("attachment; filename=%s", zipFilename))
|
||||
c.Set("Content-Type", "application/zip")
|
||||
c.Set("X-Snapshot-Size", fmt.Sprintf("%d", result.FileSize))
|
||||
|
||||
return c.SendFile(result.OutputPath)
|
||||
}
|
||||
}
|
||||
|
||||
// SnapshotImport handles importing a workspace from a snapshot
|
||||
// @Summary Import workspace snapshot
|
||||
// @Description Import a workspace from an uploaded zip file or URL
|
||||
// @Tags Snapshots
|
||||
// @Accept multipart/form-data
|
||||
// @Produce json
|
||||
// @Param file formData file false "Snapshot zip file to import"
|
||||
// @Param url formData string false "URL of snapshot to download and import"
|
||||
// @Param force formData bool false "Overwrite existing workspace if present"
|
||||
// @Param skip_db formData bool false "Skip database import (files only)"
|
||||
// @Success 200 {object} map[string]interface{} "Import result"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to import snapshot"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/snapshots/import [post]
|
||||
func SnapshotImport(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
force := c.FormValue("force") == "true"
|
||||
skipDB := c.FormValue("skip_db") == "true"
|
||||
url := c.FormValue("url")
|
||||
|
||||
var source string
|
||||
|
||||
// Check for file upload
|
||||
file, err := c.FormFile("file")
|
||||
if err == nil && file != nil {
|
||||
// Save uploaded file temporarily
|
||||
tempDir := os.TempDir()
|
||||
tempPath := filepath.Join(tempDir, file.Filename)
|
||||
|
||||
if err := c.SaveFile(file, tempPath); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to save uploaded file: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
source = tempPath
|
||||
defer func() { _ = os.Remove(tempPath) }()
|
||||
} else if url != "" {
|
||||
source = url
|
||||
} else {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Either file or url is required",
|
||||
})
|
||||
}
|
||||
|
||||
var result *snapshot.ImportResult
|
||||
|
||||
if force {
|
||||
result, err = snapshot.ForceImportWorkspace(source, cfg.WorkspacesPath, skipDB, cfg)
|
||||
} else {
|
||||
result, err = snapshot.ImportWorkspace(source, cfg.WorkspacesPath, skipDB, cfg)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to import snapshot: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Workspace imported successfully",
|
||||
"workspace": result.WorkspaceName,
|
||||
"local_path": result.LocalPath,
|
||||
"data_source": result.DataSource,
|
||||
"files_count": result.FilesCount,
|
||||
"warning": "Imported workspace database state may be unstable. Only import from trusted sources.",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteSnapshot handles deleting a snapshot file
|
||||
// @Summary Delete snapshot
|
||||
// @Description Delete a snapshot file by name
|
||||
// @Tags Snapshots
|
||||
// @Produce json
|
||||
// @Param name path string true "Snapshot filename"
|
||||
// @Success 200 {object} map[string]interface{} "Snapshot deleted"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Failure 404 {object} map[string]interface{} "Snapshot not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to delete snapshot"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/snapshots/{name} [delete]
|
||||
func DeleteSnapshot(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
name := c.Params("name")
|
||||
if name == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Snapshot name is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Sanitize path to prevent directory traversal
|
||||
if filepath.Base(name) != name {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid snapshot name",
|
||||
})
|
||||
}
|
||||
|
||||
snapshotPath := filepath.Join(cfg.SnapshotPath, name)
|
||||
|
||||
// Check if file exists
|
||||
if _, err := os.Stat(snapshotPath); os.IsNotExist(err) {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Snapshot not found: " + name,
|
||||
})
|
||||
}
|
||||
|
||||
// Delete the file
|
||||
if err := os.Remove(snapshotPath); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to delete snapshot: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Snapshot deleted successfully",
|
||||
"name": name,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func setupSnapshotTestConfig(t *testing.T) (*config.Config, func()) {
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
snapshotDir := filepath.Join(tmpDir, "snapshot")
|
||||
workspacesDir := filepath.Join(tmpDir, "workspaces")
|
||||
|
||||
err := os.MkdirAll(snapshotDir, 0755)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = os.MkdirAll(workspacesDir, 0755)
|
||||
require.NoError(t, err)
|
||||
|
||||
cfg := &config.Config{
|
||||
SnapshotPath: snapshotDir,
|
||||
WorkspacesPath: workspacesDir,
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
}
|
||||
|
||||
return cfg, cleanup
|
||||
}
|
||||
|
||||
func TestListSnapshots_Empty(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/snapshots", ListSnapshots(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/snapshots", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, float64(0), result["count"])
|
||||
// data can be null or empty array when no snapshots exist
|
||||
assert.Equal(t, cfg.SnapshotPath, result["path"])
|
||||
}
|
||||
|
||||
func TestListSnapshots_WithSnapshots(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
// Create test snapshot files
|
||||
err := os.WriteFile(filepath.Join(cfg.SnapshotPath, "example.com_123.zip"), []byte("test"), 0644)
|
||||
require.NoError(t, err)
|
||||
err = os.WriteFile(filepath.Join(cfg.SnapshotPath, "test.com_456.zip"), []byte("test2"), 0644)
|
||||
require.NoError(t, err)
|
||||
|
||||
app := fiber.New()
|
||||
app.Get("/snapshots", ListSnapshots(cfg))
|
||||
|
||||
req := httptest.NewRequest("GET", "/snapshots", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, float64(2), result["count"])
|
||||
}
|
||||
|
||||
func TestSnapshotExport_MissingWorkspace(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Post("/snapshots/export", SnapshotExport(cfg))
|
||||
|
||||
// Empty request body
|
||||
req := httptest.NewRequest("POST", "/snapshots/export", bytes.NewReader([]byte("{}")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusBadRequest, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, result["error"].(bool))
|
||||
assert.Contains(t, result["message"], "Workspace name is required")
|
||||
}
|
||||
|
||||
func TestSnapshotExport_WorkspaceNotFound(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Post("/snapshots/export", SnapshotExport(cfg))
|
||||
|
||||
reqBody := `{"workspace": "nonexistent.com"}`
|
||||
req := httptest.NewRequest("POST", "/snapshots/export", bytes.NewReader([]byte(reqBody)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusNotFound, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, result["error"].(bool))
|
||||
assert.Contains(t, result["message"], "Workspace not found")
|
||||
}
|
||||
|
||||
func TestSnapshotExport_Success(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
// Create a test workspace
|
||||
workspacePath := filepath.Join(cfg.WorkspacesPath, "example.com")
|
||||
err := os.MkdirAll(workspacePath, 0755)
|
||||
require.NoError(t, err)
|
||||
err = os.WriteFile(filepath.Join(workspacePath, "output.txt"), []byte("scan results"), 0644)
|
||||
require.NoError(t, err)
|
||||
|
||||
app := fiber.New()
|
||||
app.Post("/snapshots/export", SnapshotExport(cfg))
|
||||
|
||||
reqBody := `{"workspace": "example.com"}`
|
||||
req := httptest.NewRequest("POST", "/snapshots/export", bytes.NewReader([]byte(reqBody)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
assert.Equal(t, "application/zip", resp.Header.Get("Content-Type"))
|
||||
assert.Contains(t, resp.Header.Get("Content-Disposition"), "attachment")
|
||||
assert.NotEmpty(t, resp.Header.Get("X-Snapshot-Size"))
|
||||
}
|
||||
|
||||
func TestSnapshotImport_NoSource(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Post("/snapshots/import", SnapshotImport(cfg))
|
||||
|
||||
// Create empty multipart form
|
||||
body := &bytes.Buffer{}
|
||||
writer := multipart.NewWriter(body)
|
||||
err := writer.Close()
|
||||
require.NoError(t, err)
|
||||
|
||||
req := httptest.NewRequest("POST", "/snapshots/import", body)
|
||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusBadRequest, resp.StatusCode)
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(respBody, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, result["error"].(bool))
|
||||
assert.Contains(t, result["message"], "Either file or url is required")
|
||||
}
|
||||
|
||||
func TestSnapshotImport_InvalidRequest(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Post("/snapshots/import", SnapshotImport(cfg))
|
||||
|
||||
// Send request without multipart form
|
||||
req := httptest.NewRequest("POST", "/snapshots/import", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusBadRequest, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestDeleteSnapshot_NotFound(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Delete("/snapshots/:name", DeleteSnapshot(cfg))
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/snapshots/nonexistent.zip", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusNotFound, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, result["error"].(bool))
|
||||
assert.Contains(t, result["message"], "Snapshot not found")
|
||||
}
|
||||
|
||||
func TestDeleteSnapshot_InvalidPath(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Delete("/snapshots/:name", DeleteSnapshot(cfg))
|
||||
|
||||
// Fiber preserves URL-encoded values in params, so %2F stays as %2F
|
||||
// The handler's filepath.Base check only catches actual path separators
|
||||
// URL-encoded path traversal is treated as a literal filename and returns 404
|
||||
req := httptest.NewRequest("DELETE", "/snapshots/..%2F..%2Fetc%2Fpasswd", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Returns 404 because the URL-encoded string is treated as a literal filename
|
||||
assert.Equal(t, fiber.StatusNotFound, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestDeleteSnapshot_PathTraversalWithSlash(t *testing.T) {
|
||||
_, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
// Use wildcard to accept path with slashes for testing
|
||||
app.Delete("/snapshots/*", func(c *fiber.Ctx) error {
|
||||
name := c.Params("*")
|
||||
if name == "" || filepath.Base(name) != name {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid snapshot name",
|
||||
})
|
||||
}
|
||||
return c.JSON(fiber.Map{"name": name})
|
||||
})
|
||||
|
||||
// Test with actual path traversal (slashes in path)
|
||||
req := httptest.NewRequest("DELETE", "/snapshots/../../../etc/passwd", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusBadRequest, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, result["error"].(bool))
|
||||
assert.Contains(t, result["message"], "Invalid snapshot name")
|
||||
}
|
||||
|
||||
func TestDeleteSnapshot_Success(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
// Create a test snapshot file
|
||||
snapshotPath := filepath.Join(cfg.SnapshotPath, "example.com_123.zip")
|
||||
err := os.WriteFile(snapshotPath, []byte("test content"), 0644)
|
||||
require.NoError(t, err)
|
||||
|
||||
app := fiber.New()
|
||||
app.Delete("/snapshots/:name", DeleteSnapshot(cfg))
|
||||
|
||||
req := httptest.NewRequest("DELETE", "/snapshots/example.com_123.zip", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, fiber.StatusOK, resp.StatusCode)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
|
||||
var result map[string]interface{}
|
||||
err = json.Unmarshal(body, &result)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Contains(t, result["message"], "Snapshot deleted successfully")
|
||||
assert.Equal(t, "example.com_123.zip", result["name"])
|
||||
|
||||
// Verify file was deleted
|
||||
_, err = os.Stat(snapshotPath)
|
||||
assert.True(t, os.IsNotExist(err))
|
||||
}
|
||||
|
||||
func TestDeleteSnapshot_MissingName(t *testing.T) {
|
||||
cfg, cleanup := setupSnapshotTestConfig(t)
|
||||
defer cleanup()
|
||||
|
||||
app := fiber.New()
|
||||
app.Delete("/snapshots/:name", DeleteSnapshot(cfg))
|
||||
|
||||
// Empty name parameter - Fiber will route to empty string
|
||||
req := httptest.NewRequest("DELETE", "/snapshots/", nil)
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Fiber returns 404 for unmatched routes
|
||||
assert.Equal(t, fiber.StatusNotFound, resp.StatusCode)
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
// GetSystemStats returns aggregated system statistics
|
||||
// @Summary Get system statistics
|
||||
// @Description Get aggregated counts for workflows, runs, workspaces, assets, vulnerabilities, and schedules
|
||||
// @Tags Stats
|
||||
// @Produce json
|
||||
// @Success 200 {object} database.SystemStats "System statistics"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to get stats"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/stats [get]
|
||||
func GetSystemStats(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
ctx := context.Background()
|
||||
|
||||
// Get system stats
|
||||
stats, err := database.GetSystemStats(ctx, cfg.WorkflowsPath)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(stats)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,315 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"compress/flate"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/parser"
|
||||
)
|
||||
|
||||
// UploadFile handles uploading a file containing a list of inputs
|
||||
// @Summary Upload input file
|
||||
// @Description Upload a file containing a list of inputs (targets, URLs, etc.) for later use in runs
|
||||
// @Tags Files
|
||||
// @Accept multipart/form-data
|
||||
// @Produce json
|
||||
// @Param file formData file true "Input file to upload"
|
||||
// @Success 200 {object} map[string]interface{} "File uploaded with path"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/upload-file [post]
|
||||
func UploadFile(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Get the file from the form
|
||||
file, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "No file provided",
|
||||
})
|
||||
}
|
||||
|
||||
// Create uploads directory if it doesn't exist
|
||||
uploadsDir := cfg.DataPath + "/uploads"
|
||||
if err := os.MkdirAll(uploadsDir, 0755); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to create uploads directory",
|
||||
})
|
||||
}
|
||||
|
||||
// Generate a unique filename to avoid conflicts
|
||||
uniqueFilename := fmt.Sprintf("%d_%s", time.Now().UnixNano(), file.Filename)
|
||||
destPath := uploadsDir + "/" + uniqueFilename
|
||||
|
||||
// Save the file
|
||||
if err := c.SaveFile(file, destPath); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to save file",
|
||||
})
|
||||
}
|
||||
|
||||
// Count lines in the file
|
||||
lineCount := 0
|
||||
content, err := os.ReadFile(destPath)
|
||||
if err == nil {
|
||||
for _, b := range content {
|
||||
if b == '\n' {
|
||||
lineCount++
|
||||
}
|
||||
}
|
||||
// Add 1 if file doesn't end with newline but has content
|
||||
if len(content) > 0 && content[len(content)-1] != '\n' {
|
||||
lineCount++
|
||||
}
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "File uploaded",
|
||||
"filename": uniqueFilename,
|
||||
"path": destPath,
|
||||
"size": file.Size,
|
||||
"lines": lineCount,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// UploadWorkflow handles uploading a workflow YAML file
|
||||
// @Summary Upload workflow file
|
||||
// @Description Upload a raw YAML workflow file and save it to the workflows directory
|
||||
// @Tags Workflows
|
||||
// @Accept multipart/form-data
|
||||
// @Produce json
|
||||
// @Param file formData file true "Workflow YAML file"
|
||||
// @Success 201 {object} map[string]interface{} "Workflow uploaded"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid request or YAML"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workflow-upload [post]
|
||||
func UploadWorkflow(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Get the file from the form
|
||||
file, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "No file provided",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate file extension
|
||||
filename := file.Filename
|
||||
if !strings.HasSuffix(filename, ".yaml") && !strings.HasSuffix(filename, ".yml") {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Only YAML files (.yaml or .yml) are allowed",
|
||||
})
|
||||
}
|
||||
|
||||
// Open the file to read its content
|
||||
src, err := file.Open()
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to read file",
|
||||
})
|
||||
}
|
||||
defer func() { _ = src.Close() }()
|
||||
|
||||
// Read the content
|
||||
content, err := io.ReadAll(src)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to read file content",
|
||||
})
|
||||
}
|
||||
|
||||
// Parse the workflow to validate and get its properties
|
||||
workflow, err := parser.ParseContent(content)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid workflow YAML: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Validate workflow using the parser
|
||||
p := parser.NewParser()
|
||||
if err := p.Validate(workflow); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workflow validation failed: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Determine target directory based on workflow kind
|
||||
var targetDir string
|
||||
if workflow.IsFlow() {
|
||||
targetDir = cfg.WorkflowsPath + "/flows"
|
||||
} else {
|
||||
targetDir = cfg.WorkflowsPath + "/modules"
|
||||
}
|
||||
|
||||
// Create target directory if it doesn't exist
|
||||
if err := os.MkdirAll(targetDir, 0755); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to create workflows directory",
|
||||
})
|
||||
}
|
||||
|
||||
// Save the workflow file
|
||||
destPath := targetDir + "/" + filename
|
||||
if err := os.WriteFile(destPath, content, 0644); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to save workflow file",
|
||||
})
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusCreated).JSON(fiber.Map{
|
||||
"message": "Workflow uploaded",
|
||||
"name": workflow.Name,
|
||||
"kind": workflow.Kind,
|
||||
"description": workflow.Description,
|
||||
"path": destPath,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// SnapshotDownload compresses a workspace and serves it for download
|
||||
// @Summary Download workspace snapshot
|
||||
// @Description Compress a workspace folder into a zip file and download it
|
||||
// @Tags Snapshots
|
||||
// @Produce application/zip
|
||||
// @Param workspace_name path string true "Workspace name"
|
||||
// @Success 200 {file} file "Zip file download"
|
||||
// @Failure 404 {object} map[string]interface{} "Workspace not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to create snapshot"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/snapshot-download/{workspace_name} [get]
|
||||
func SnapshotDownload(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
workspaceName := c.Params("workspace_name")
|
||||
if workspaceName == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workspace name is required",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate workspace exists
|
||||
workspacePath := cfg.WorkspacesPath + "/" + workspaceName
|
||||
if _, err := os.Stat(workspacePath); os.IsNotExist(err) {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workspace not found: " + workspaceName,
|
||||
})
|
||||
}
|
||||
|
||||
// Create snapshot directory if it doesn't exist
|
||||
if err := os.MkdirAll(cfg.SnapshotPath, 0755); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to create snapshot directory",
|
||||
})
|
||||
}
|
||||
|
||||
// Generate zip filename with timestamp
|
||||
zipFilename := fmt.Sprintf("%s_%d.zip", workspaceName, time.Now().Unix())
|
||||
zipPath := cfg.SnapshotPath + "/" + zipFilename
|
||||
|
||||
// Create the zip file
|
||||
if err := createZipArchive(workspacePath, zipPath); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to create snapshot: " + err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Set headers for file download
|
||||
c.Set("Content-Disposition", fmt.Sprintf("attachment; filename=%s", zipFilename))
|
||||
c.Set("Content-Type", "application/zip")
|
||||
|
||||
// Send the file
|
||||
return c.SendFile(zipPath)
|
||||
}
|
||||
}
|
||||
|
||||
// createZipArchive creates a zip archive of the source directory
|
||||
func createZipArchive(source, target string) error {
|
||||
zipFile, err := os.Create(target)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = zipFile.Close() }()
|
||||
|
||||
archive := zip.NewWriter(zipFile)
|
||||
defer func() { _ = archive.Close() }()
|
||||
|
||||
// Register highest compression level for best compression ratio
|
||||
archive.RegisterCompressor(zip.Deflate, func(out io.Writer) (io.WriteCloser, error) {
|
||||
return flate.NewWriter(out, flate.BestCompression)
|
||||
})
|
||||
|
||||
// Walk through the source directory
|
||||
return filepath.Walk(source, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get relative path
|
||||
relPath, err := filepath.Rel(source, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Skip the root directory itself
|
||||
if relPath == "." {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Create header
|
||||
header, err := zip.FileInfoHeader(info)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
header.Name = relPath
|
||||
|
||||
if info.IsDir() {
|
||||
header.Name += "/"
|
||||
} else {
|
||||
header.Method = zip.Deflate
|
||||
}
|
||||
|
||||
// Create writer for this file
|
||||
writer, err := archive.CreateHeader(header)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// If it's a directory, we're done
|
||||
if info.IsDir() {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Copy file contents
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
|
||||
_, err = io.Copy(writer, file)
|
||||
return err
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,279 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
// ListVulnerabilities handles listing vulnerabilities with pagination and filtering
|
||||
// @Summary List vulnerabilities
|
||||
// @Description Get a paginated list of vulnerabilities with optional workspace, severity, and confidence filtering
|
||||
// @Tags Vulnerabilities
|
||||
// @Produce json
|
||||
// @Param workspace query string false "Filter by workspace name"
|
||||
// @Param severity query string false "Filter by severity (critical, high, medium, low, info)"
|
||||
// @Param confidence query string false "Filter by confidence (certain, firm, tentative, manual review required)"
|
||||
// @Param asset_value query string false "Filter by asset value (partial match)"
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return" default(20)
|
||||
// @Success 200 {object} map[string]interface{} "List of vulnerabilities with pagination"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to fetch vulnerabilities"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/vulnerabilities [get]
|
||||
func ListVulnerabilities(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse query parameters
|
||||
workspace := c.Query("workspace")
|
||||
severity := c.Query("severity")
|
||||
confidence := c.Query("confidence")
|
||||
assetValue := c.Query("asset_value")
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
|
||||
// Validate pagination
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Get vulnerabilities from database
|
||||
result, err := database.ListVulnerabilities(ctx, database.VulnerabilityQuery{
|
||||
Workspace: workspace,
|
||||
Severity: severity,
|
||||
Confidence: confidence,
|
||||
AssetValue: assetValue,
|
||||
Offset: offset,
|
||||
Limit: limit,
|
||||
})
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetVulnerability handles getting a single vulnerability by ID
|
||||
// @Summary Get vulnerability by ID
|
||||
// @Description Get a single vulnerability by its ID
|
||||
// @Tags Vulnerabilities
|
||||
// @Produce json
|
||||
// @Param id path int true "Vulnerability ID"
|
||||
// @Success 200 {object} map[string]interface{} "Vulnerability details"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid ID"
|
||||
// @Failure 404 {object} map[string]interface{} "Vulnerability not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to fetch vulnerability"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/vulnerabilities/{id} [get]
|
||||
func GetVulnerability(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse ID
|
||||
idStr := c.Params("id")
|
||||
id, err := strconv.ParseInt(idStr, 10, 64)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid vulnerability ID",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Get vulnerability
|
||||
vuln, err := database.GetVulnerabilityByID(ctx, id)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Vulnerability not found",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": vuln,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// CreateVulnerabilityInput represents the input for creating a vulnerability
|
||||
type CreateVulnerabilityInput struct {
|
||||
Workspace string `json:"workspace"`
|
||||
VulnInfo string `json:"vuln_info"`
|
||||
VulnTitle string `json:"vuln_title"`
|
||||
VulnDesc string `json:"vuln_desc"`
|
||||
VulnPOC string `json:"vuln_poc"`
|
||||
Severity string `json:"severity"`
|
||||
AssetType string `json:"asset_type"`
|
||||
AssetValue string `json:"asset_value"`
|
||||
Tags []string `json:"tags"`
|
||||
DetailHTTPRequest string `json:"detail_http_request"`
|
||||
DetailHTTPResponse string `json:"detail_http_response"`
|
||||
RawVulnJSON string `json:"raw_vuln_json"`
|
||||
}
|
||||
|
||||
// CreateVulnerability handles creating a new vulnerability
|
||||
// @Summary Create vulnerability
|
||||
// @Description Create a new vulnerability record
|
||||
// @Tags Vulnerabilities
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param vulnerability body CreateVulnerabilityInput true "Vulnerability data"
|
||||
// @Success 201 {object} map[string]interface{} "Created vulnerability"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid input"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to create vulnerability"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/vulnerabilities [post]
|
||||
func CreateVulnerability(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse input
|
||||
var input CreateVulnerabilityInput
|
||||
if err := c.BodyParser(&input); err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid request body",
|
||||
})
|
||||
}
|
||||
|
||||
// Validate required fields
|
||||
if input.Workspace == "" {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workspace is required",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Create vulnerability
|
||||
vuln := &database.Vulnerability{
|
||||
Workspace: input.Workspace,
|
||||
VulnInfo: input.VulnInfo,
|
||||
VulnTitle: input.VulnTitle,
|
||||
VulnDesc: input.VulnDesc,
|
||||
VulnPOC: input.VulnPOC,
|
||||
Severity: input.Severity,
|
||||
AssetType: input.AssetType,
|
||||
AssetValue: input.AssetValue,
|
||||
Tags: input.Tags,
|
||||
DetailHTTPRequest: input.DetailHTTPRequest,
|
||||
DetailHTTPResponse: input.DetailHTTPResponse,
|
||||
RawVulnJSON: input.RawVulnJSON,
|
||||
CreatedAt: time.Now(),
|
||||
UpdatedAt: time.Now(),
|
||||
}
|
||||
|
||||
if err := database.CreateVulnerabilityRecord(ctx, vuln); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.Status(fiber.StatusCreated).JSON(fiber.Map{
|
||||
"data": vuln,
|
||||
"message": "Vulnerability created successfully",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteVulnerability handles deleting a vulnerability
|
||||
// @Summary Delete vulnerability
|
||||
// @Description Delete a vulnerability by ID
|
||||
// @Tags Vulnerabilities
|
||||
// @Produce json
|
||||
// @Param id path int true "Vulnerability ID"
|
||||
// @Success 200 {object} map[string]interface{} "Vulnerability deleted"
|
||||
// @Failure 400 {object} map[string]interface{} "Invalid ID"
|
||||
// @Failure 404 {object} map[string]interface{} "Vulnerability not found"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to delete vulnerability"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/vulnerabilities/{id} [delete]
|
||||
func DeleteVulnerability(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse ID
|
||||
idStr := c.Params("id")
|
||||
id, err := strconv.ParseInt(idStr, 10, 64)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid vulnerability ID",
|
||||
})
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Delete vulnerability
|
||||
if err := database.DeleteVulnerabilityByID(ctx, id); err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Vulnerability deleted successfully",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetVulnerabilitySummary returns severity summary for a workspace
|
||||
// @Summary Get vulnerability summary
|
||||
// @Description Get a summary of vulnerabilities grouped by severity
|
||||
// @Tags Vulnerabilities
|
||||
// @Produce json
|
||||
// @Param workspace query string false "Filter by workspace name"
|
||||
// @Success 200 {object} map[string]interface{} "Vulnerability summary by severity"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to get summary"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/vulnerabilities/summary [get]
|
||||
func GetVulnerabilitySummary(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
workspace := c.Query("workspace")
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Get summary
|
||||
summary, err := database.GetVulnerabilitySummary(ctx, workspace)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
// Calculate total
|
||||
total := 0
|
||||
for _, count := range summary {
|
||||
total += count
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": fiber.Map{
|
||||
"by_severity": summary,
|
||||
"total": total,
|
||||
"workspace": workspace,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,461 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/core"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/parser"
|
||||
)
|
||||
|
||||
// ListWorkflows handles listing workflows
|
||||
func ListWorkflows(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflows, err := loader.LoadAllWorkflows()
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to load workflows",
|
||||
})
|
||||
}
|
||||
|
||||
data := make([]fiber.Map, 0, len(workflows))
|
||||
for _, w := range workflows {
|
||||
data = append(data, fiber.Map{
|
||||
"name": w.Name,
|
||||
"kind": w.Kind,
|
||||
"description": w.Description,
|
||||
"file_path": w.FilePath,
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetWorkflow handles getting a single workflow
|
||||
func GetWorkflow(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
name := c.Params("name")
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflow, err := loader.LoadWorkflow(name)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workflow not found",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(workflowToMap(workflow))
|
||||
}
|
||||
}
|
||||
|
||||
// ListWorkflowsVerbose handles listing workflows with verbose details
|
||||
// @Summary List all workflows
|
||||
// @Description Get a list of all available workflows with details
|
||||
// @Tags Workflows
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]interface{} "List of workflows"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to load workflows"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workflows [get]
|
||||
func ListWorkflowsVerbose(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Check source parameter - "filesystem" for direct loading, default is "db"
|
||||
source := c.Query("source", "db")
|
||||
|
||||
// If source is filesystem, use the original approach
|
||||
if source == "filesystem" {
|
||||
return listWorkflowsFromFilesystem(c, cfg)
|
||||
}
|
||||
|
||||
// Parse query parameters for DB-based listing
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "50"))
|
||||
kind := c.Query("kind", "")
|
||||
search := c.Query("search", "")
|
||||
tagsStr := c.Query("tags", "")
|
||||
|
||||
var tags []string
|
||||
if tagsStr != "" {
|
||||
tags = strings.Split(tagsStr, ",")
|
||||
for i := range tags {
|
||||
tags[i] = strings.TrimSpace(tags[i])
|
||||
}
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Query database
|
||||
query := database.WorkflowQuery{
|
||||
Tags: tags,
|
||||
Kind: kind,
|
||||
Search: search,
|
||||
Offset: offset,
|
||||
Limit: limit,
|
||||
}
|
||||
|
||||
result, err := database.ListWorkflowsFromDB(ctx, query)
|
||||
if err != nil {
|
||||
// Fallback to filesystem if DB query fails
|
||||
return listWorkflowsFromFilesystem(c, cfg)
|
||||
}
|
||||
|
||||
// If no workflows in DB, suggest indexing
|
||||
if result.TotalCount == 0 {
|
||||
// Check if DB is empty vs no matches
|
||||
totalCount, _ := database.GetWorkflowCount(ctx)
|
||||
if totalCount == 0 {
|
||||
return c.JSON(fiber.Map{
|
||||
"data": []interface{}{},
|
||||
"count": 0,
|
||||
"message": "No workflows indexed. Run 'osmedeus db index workflow' or POST /osm/api/workflows/refresh to index workflows.",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// listWorkflowsFromFilesystem loads workflows directly from filesystem
|
||||
func listWorkflowsFromFilesystem(c *fiber.Ctx, cfg *config.Config) error {
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflows, err := loader.LoadAllWorkflows()
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to load workflows",
|
||||
})
|
||||
}
|
||||
|
||||
data := make([]fiber.Map, 0, len(workflows))
|
||||
for _, w := range workflows {
|
||||
// Build params detail
|
||||
params := make([]fiber.Map, 0, len(w.Params))
|
||||
requiredParams := []string{}
|
||||
for _, p := range w.Params {
|
||||
params = append(params, fiber.Map{
|
||||
"name": p.Name,
|
||||
"default": p.Default,
|
||||
"required": p.Required,
|
||||
"generator": p.Generator,
|
||||
})
|
||||
if p.Required {
|
||||
requiredParams = append(requiredParams, p.Name)
|
||||
}
|
||||
}
|
||||
|
||||
data = append(data, fiber.Map{
|
||||
"name": w.Name,
|
||||
"kind": w.Kind,
|
||||
"description": w.Description,
|
||||
"tags": w.Tags,
|
||||
"file_path": w.FilePath,
|
||||
"params": params,
|
||||
"required_params": requiredParams,
|
||||
"step_count": len(w.Steps),
|
||||
"module_count": len(w.Modules),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": data,
|
||||
"count": len(data),
|
||||
})
|
||||
}
|
||||
|
||||
// GetWorkflowVerbose handles getting a single workflow with verbose details
|
||||
// @Summary Get workflow details
|
||||
// @Description Get workflow content. Returns raw YAML by default. Use json=true to get JSON with parsed details.
|
||||
// @Tags Workflows
|
||||
// @Produce json,text/yaml
|
||||
// @Param name path string true "Workflow name"
|
||||
// @Param json query bool false "Return JSON with parsed details instead of raw YAML"
|
||||
// @Success 200 {object} map[string]interface{} "Workflow details (JSON) or raw YAML content"
|
||||
// @Failure 404 {object} map[string]interface{} "Workflow not found"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workflows/{name} [get]
|
||||
func GetWorkflowVerbose(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
name := c.Params("name")
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflow, err := loader.LoadWorkflow(name)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workflow not found",
|
||||
})
|
||||
}
|
||||
|
||||
// Check if JSON format is requested (default is YAML)
|
||||
if c.Query("json") == "true" {
|
||||
return returnWorkflowJSON(c, workflow)
|
||||
}
|
||||
|
||||
// Return YAML by default
|
||||
content, err := os.ReadFile(workflow.FilePath)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to read workflow file",
|
||||
})
|
||||
}
|
||||
c.Set("Content-Type", "text/yaml; charset=utf-8")
|
||||
return c.SendString(string(content))
|
||||
}
|
||||
}
|
||||
|
||||
// returnWorkflowJSON returns workflow details as JSON
|
||||
func returnWorkflowJSON(c *fiber.Ctx, workflow *core.Workflow) error {
|
||||
// Build params detail
|
||||
params := make([]fiber.Map, 0, len(workflow.Params))
|
||||
for _, p := range workflow.Params {
|
||||
params = append(params, fiber.Map{
|
||||
"name": p.Name,
|
||||
"default": p.Default,
|
||||
"required": p.Required,
|
||||
"generator": p.Generator,
|
||||
})
|
||||
}
|
||||
|
||||
// Build steps detail with full information
|
||||
steps := make([]fiber.Map, 0, len(workflow.Steps))
|
||||
for i, s := range workflow.Steps {
|
||||
stepMap := fiber.Map{
|
||||
"index": i,
|
||||
"name": s.Name,
|
||||
"type": s.Type,
|
||||
"command": s.Command,
|
||||
"timeout": s.Timeout,
|
||||
"pre_condition": s.PreCondition,
|
||||
"step_runner": s.StepRunner,
|
||||
}
|
||||
// Add exports if present
|
||||
if len(s.Exports) > 0 {
|
||||
stepMap["exports"] = s.Exports
|
||||
}
|
||||
// Add step runner config if present
|
||||
if s.StepRunnerConfig != nil && s.StepRunnerConfig.RunnerConfig != nil {
|
||||
stepMap["step_runner_config"] = fiber.Map{
|
||||
"image": s.StepRunnerConfig.Image,
|
||||
"host": s.StepRunnerConfig.Host,
|
||||
"user": s.StepRunnerConfig.User,
|
||||
"volumes": s.StepRunnerConfig.Volumes,
|
||||
}
|
||||
}
|
||||
// Add parallel steps if present
|
||||
if len(s.ParallelSteps) > 0 {
|
||||
substeps := make([]fiber.Map, 0, len(s.ParallelSteps))
|
||||
for j, ss := range s.ParallelSteps {
|
||||
substeps = append(substeps, fiber.Map{
|
||||
"index": j,
|
||||
"name": ss.Name,
|
||||
"type": ss.Type,
|
||||
"command": ss.Command,
|
||||
})
|
||||
}
|
||||
stepMap["parallel_steps"] = substeps
|
||||
}
|
||||
steps = append(steps, stepMap)
|
||||
}
|
||||
|
||||
// Build modules detail with full information
|
||||
modules := make([]fiber.Map, 0, len(workflow.Modules))
|
||||
for i, m := range workflow.Modules {
|
||||
modules = append(modules, fiber.Map{
|
||||
"index": i,
|
||||
"name": m.Name,
|
||||
"path": m.Path,
|
||||
"depends_on": m.DependsOn,
|
||||
"condition": m.Condition,
|
||||
})
|
||||
}
|
||||
|
||||
// Build triggers detail
|
||||
triggers := make([]fiber.Map, 0, len(workflow.Triggers))
|
||||
for _, t := range workflow.Triggers {
|
||||
triggers = append(triggers, fiber.Map{
|
||||
"name": t.Name,
|
||||
"on": t.On,
|
||||
"schedule": t.Schedule,
|
||||
"enabled": t.Enabled,
|
||||
})
|
||||
}
|
||||
|
||||
// Build dependencies info
|
||||
var dependencies fiber.Map
|
||||
if workflow.Dependencies != nil {
|
||||
dependencies = fiber.Map{
|
||||
"commands": workflow.Dependencies.Commands,
|
||||
"files": workflow.Dependencies.Files,
|
||||
}
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"name": workflow.Name,
|
||||
"kind": workflow.Kind,
|
||||
"description": workflow.Description,
|
||||
"file_path": workflow.FilePath,
|
||||
"params": params,
|
||||
"steps": steps,
|
||||
"modules": modules,
|
||||
"triggers": triggers,
|
||||
"dependencies": dependencies,
|
||||
})
|
||||
}
|
||||
|
||||
// ValidateWorkflow handles workflow validation
|
||||
func ValidateWorkflow(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
name := c.Params("name")
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
workflow, err := loader.LoadWorkflow(name)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusNotFound).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Workflow not found",
|
||||
})
|
||||
}
|
||||
|
||||
if err := parser.Validate(workflow); err != nil {
|
||||
return c.JSON(fiber.Map{
|
||||
"valid": false,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"valid": true,
|
||||
"message": "Workflow is valid",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ReloadWorkflows handles reloading workflows from disk
|
||||
func ReloadWorkflows(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
loader := parser.NewLoader(cfg.WorkflowsPath)
|
||||
if err := loader.ReloadWorkflows(); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Failed to reload workflows",
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Workflows reloaded",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// RefreshWorkflowIndex re-indexes workflows from filesystem to database
|
||||
// @Summary Refresh workflow index
|
||||
// @Description Re-index all workflows from filesystem to database
|
||||
// @Tags Workflows
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param force query bool false "Force re-index all workflows regardless of checksum"
|
||||
// @Success 200 {object} map[string]interface{} "Indexing result"
|
||||
// @Failure 500 {object} map[string]interface{} "Indexing failed"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workflows/refresh [post]
|
||||
func RefreshWorkflowIndex(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
force := c.Query("force") == "true"
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Index workflows
|
||||
result, err := database.IndexWorkflowsFromFilesystem(ctx, cfg.WorkflowsPath, force)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"message": "Workflows indexed successfully",
|
||||
"added": result.Added,
|
||||
"updated": result.Updated,
|
||||
"removed": result.Removed,
|
||||
"errors": result.Errors,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// GetAllWorkflowTags returns all unique tags from indexed workflows
|
||||
// @Summary Get all workflow tags
|
||||
// @Description Get all unique tags from indexed workflows in database
|
||||
// @Tags Workflows
|
||||
// @Produce json
|
||||
// @Success 200 {object} map[string]interface{} "List of tags"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to get tags"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workflows/tags [get]
|
||||
func GetAllWorkflowTags(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
ctx := context.Background()
|
||||
|
||||
tags, err := database.GetAllTags(ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"tags": tags,
|
||||
"count": len(tags),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// workflowToMap converts a workflow to a map for JSON response
|
||||
func workflowToMap(w *core.Workflow) fiber.Map {
|
||||
params := make([]fiber.Map, 0, len(w.Params))
|
||||
for _, p := range w.Params {
|
||||
params = append(params, fiber.Map{
|
||||
"name": p.Name,
|
||||
"default": p.Default,
|
||||
"required": p.Required,
|
||||
"generator": p.Generator,
|
||||
})
|
||||
}
|
||||
|
||||
triggers := make([]fiber.Map, 0, len(w.Triggers))
|
||||
for _, t := range w.Triggers {
|
||||
triggers = append(triggers, fiber.Map{
|
||||
"name": t.Name,
|
||||
"on": t.On,
|
||||
"schedule": t.Schedule,
|
||||
"enabled": t.Enabled,
|
||||
})
|
||||
}
|
||||
|
||||
return fiber.Map{
|
||||
"name": w.Name,
|
||||
"kind": w.Kind,
|
||||
"description": w.Description,
|
||||
"file_path": w.FilePath,
|
||||
"params": params,
|
||||
"triggers": triggers,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,287 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
)
|
||||
|
||||
func listWorkspaceDirs(workspacesDir string) ([]string, error) {
|
||||
if workspacesDir == "" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
entries, err := os.ReadDir(workspacesDir)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(entries))
|
||||
for _, entry := range entries {
|
||||
name := entry.Name()
|
||||
if !isValidWorkspaceName(name) {
|
||||
continue
|
||||
}
|
||||
if entry.IsDir() {
|
||||
names = append(names, name)
|
||||
continue
|
||||
}
|
||||
if entry.Type()&os.ModeSymlink != 0 {
|
||||
info, err := os.Stat(filepath.Join(workspacesDir, name))
|
||||
if err == nil && info.IsDir() {
|
||||
names = append(names, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
sort.Strings(names)
|
||||
return names, nil
|
||||
}
|
||||
|
||||
func resolveWorkspacesDirForListing(cfg *config.Config) string {
|
||||
configured := ""
|
||||
if cfg != nil {
|
||||
configured = cfg.GetWorkspacesDir()
|
||||
}
|
||||
if configured != "" {
|
||||
if _, err := os.Stat(configured); err == nil {
|
||||
return configured
|
||||
}
|
||||
}
|
||||
|
||||
if home, err := os.UserHomeDir(); err == nil && home != "" {
|
||||
candidate := filepath.Join(home, "workspaces-osmedeus")
|
||||
if _, err := os.Stat(candidate); err == nil {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := os.Stat("/workspaces-osmedeus"); err == nil {
|
||||
return "/workspaces-osmedeus"
|
||||
}
|
||||
|
||||
return configured
|
||||
}
|
||||
|
||||
type filesystemWorkspaceRecord struct {
|
||||
Name string `json:"name"`
|
||||
LocalPath string `json:"local_path,omitempty"`
|
||||
DataSource string `json:"data_source"`
|
||||
TotalAssets int `json:"total_assets"`
|
||||
Tags []string `json:"tags"`
|
||||
}
|
||||
|
||||
// ListWorkspaces handles listing all workspaces
|
||||
// @Summary List all workspaces
|
||||
// @Description Get a list of all run workspaces. By default returns full workspace records from database. Use filesystem=true to list workspaces derived from assets.
|
||||
// @Tags Workspaces
|
||||
// @Produce json
|
||||
// @Param filesystem query bool false "List workspaces from filesystem/assets instead of workspaces table" default(false)
|
||||
// @Param offset query int false "Number of records to skip" default(0)
|
||||
// @Param limit query int false "Maximum number of records to return (max 10000)" default(20)
|
||||
// @Success 200 {object} map[string]interface{} "List of workspaces"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to read workspaces"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workspaces [get]
|
||||
func ListWorkspaces(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Parse query parameters
|
||||
filesystem := c.Query("filesystem", "false") == "true"
|
||||
offset, _ := strconv.Atoi(c.Query("offset", "0"))
|
||||
limit, _ := strconv.Atoi(c.Query("limit", "20"))
|
||||
|
||||
// Validate pagination
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 10000 {
|
||||
limit = 10000
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Return workspaces based on mode
|
||||
if filesystem {
|
||||
assetWorkspaces, err := database.ListAllWorkspacesFromAssets(ctx)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
inAssetsDBByName := make(map[string]bool, len(assetWorkspaces))
|
||||
for _, ws := range assetWorkspaces {
|
||||
inAssetsDBByName[ws.Name] = true
|
||||
}
|
||||
|
||||
workspacesDir := resolveWorkspacesDirForListing(cfg)
|
||||
workspaceDirs, err := listWorkspaceDirs(workspacesDir)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
assetCountByName := make(map[string]int, len(assetWorkspaces))
|
||||
for _, ws := range assetWorkspaces {
|
||||
assetCountByName[ws.Name] = ws.AssetCount
|
||||
}
|
||||
|
||||
hasDirByName := make(map[string]bool, len(workspaceDirs))
|
||||
for _, name := range workspaceDirs {
|
||||
hasDirByName[name] = true
|
||||
if _, ok := assetCountByName[name]; !ok {
|
||||
assetCountByName[name] = 0
|
||||
}
|
||||
}
|
||||
|
||||
records := make([]filesystemWorkspaceRecord, 0, len(assetCountByName))
|
||||
for name, assetCount := range assetCountByName {
|
||||
tags := []string{"filesystem"}
|
||||
if hasDirByName[name] && !inAssetsDBByName[name] {
|
||||
tags = append(tags, "filesystem-only")
|
||||
}
|
||||
rec := filesystemWorkspaceRecord{
|
||||
Name: name,
|
||||
DataSource: "filesystem",
|
||||
TotalAssets: assetCount,
|
||||
Tags: tags,
|
||||
}
|
||||
if hasDirByName[name] {
|
||||
rec.LocalPath = filepath.Join(workspacesDir, name)
|
||||
}
|
||||
records = append(records, rec)
|
||||
}
|
||||
sort.Slice(records, func(i, j int) bool {
|
||||
return records[i].Name < records[j].Name
|
||||
})
|
||||
|
||||
totalCount := len(records)
|
||||
if offset > totalCount {
|
||||
offset = totalCount
|
||||
}
|
||||
end := offset + limit
|
||||
if end > totalCount {
|
||||
end = totalCount
|
||||
}
|
||||
page := records[offset:end]
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": page,
|
||||
"workspaces_dir": workspacesDir,
|
||||
"pagination": fiber.Map{
|
||||
"total": totalCount,
|
||||
"offset": offset,
|
||||
"limit": limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Default: Get full workspace records from workspaces table
|
||||
result, err := database.ListWorkspacesFullFromDB(ctx, offset, limit)
|
||||
if err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(fiber.Map{
|
||||
"data": result.Data,
|
||||
"pagination": fiber.Map{
|
||||
"total": result.TotalCount,
|
||||
"offset": result.Offset,
|
||||
"limit": result.Limit,
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ListWorkspaceNames handles listing workspace names
|
||||
// @Summary List workspace names
|
||||
// @Description Get a sorted list of workspace names from the database
|
||||
// @Tags Workspaces
|
||||
// @Produce json
|
||||
// @Success 200 {array} string "Workspace names"
|
||||
// @Failure 500 {object} map[string]interface{} "Failed to list workspace names"
|
||||
// @Security BearerAuth
|
||||
// @Router /osm/api/workspace-names [get]
|
||||
func ListWorkspaceNames(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
ctx := context.Background()
|
||||
|
||||
db := database.GetDB()
|
||||
if db == nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Database not connected",
|
||||
})
|
||||
}
|
||||
|
||||
var names []string
|
||||
if err := db.NewSelect().Model((*database.Workspace)(nil)).Column("name").Order("name ASC").Scan(ctx, &names); err != nil {
|
||||
return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
|
||||
return c.JSON(names)
|
||||
}
|
||||
}
|
||||
|
||||
// isValidWorkspaceName validates workspace name to prevent path traversal
|
||||
func isValidWorkspaceName(name string) bool {
|
||||
// Reject empty, ".", "..", or names containing path separators
|
||||
if name == "" || name == "." || name == ".." {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(name, "/") || strings.Contains(name, "\\") {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(name, "..") {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// isPathUnderWorkspace ensures the file path is within the workspace folder
|
||||
func isPathUnderWorkspace(filePath, workspacePath string) bool {
|
||||
if filePath == "" || workspacePath == "" {
|
||||
return false
|
||||
}
|
||||
absFile, err := filepath.Abs(filePath)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
absWorkspace, err := filepath.Abs(workspacePath)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
realFile, err := filepath.EvalSymlinks(absFile)
|
||||
if err == nil {
|
||||
absFile = realFile
|
||||
}
|
||||
realWorkspace, err := filepath.EvalSymlinks(absWorkspace)
|
||||
if err == nil {
|
||||
absWorkspace = realWorkspace
|
||||
}
|
||||
|
||||
return strings.HasPrefix(absFile, absWorkspace+string(filepath.Separator)) || absFile == absWorkspace
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
)
|
||||
|
||||
// Claims represents JWT claims
|
||||
type Claims struct {
|
||||
Username string `json:"username"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// JWTAuth creates JWT authentication middleware
|
||||
func JWTAuth(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Get token from Authorization header
|
||||
authHeader := c.Get("Authorization")
|
||||
if authHeader == "" {
|
||||
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Missing authorization header",
|
||||
})
|
||||
}
|
||||
|
||||
// Check Bearer prefix
|
||||
parts := strings.Split(authHeader, " ")
|
||||
if len(parts) != 2 || parts[0] != "Bearer" {
|
||||
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid authorization header format",
|
||||
})
|
||||
}
|
||||
|
||||
tokenString := parts[1]
|
||||
|
||||
// Parse and validate token
|
||||
claims := &Claims{}
|
||||
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) {
|
||||
return []byte(cfg.Server.JWT.SecretSigningKey), nil
|
||||
})
|
||||
|
||||
if err != nil || !token.Valid {
|
||||
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid or expired token",
|
||||
})
|
||||
}
|
||||
|
||||
// Store claims in context
|
||||
c.Locals("user", claims)
|
||||
|
||||
return c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// GenerateToken generates a JWT token
|
||||
func GenerateToken(username string, cfg *config.Config) (string, error) {
|
||||
expiration := time.Duration(cfg.Server.JWT.ExpirationMinutes) * time.Minute
|
||||
|
||||
claims := &Claims{
|
||||
Username: username,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(expiration)),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(cfg.Server.JWT.SecretSigningKey))
|
||||
}
|
||||
|
||||
// GetUser gets the current user from context
|
||||
func GetUser(c *fiber.Ctx) *Claims {
|
||||
claims, ok := c.Locals("user").(*Claims)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return claims
|
||||
}
|
||||
|
||||
// APIKeyAuth creates API key authentication middleware
|
||||
func APIKeyAuth(cfg *config.Config) fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
apiKey := c.Get("x-osm-api-key")
|
||||
|
||||
if !isValidAPIKey(apiKey, cfg.Server.AuthAPIKey) {
|
||||
return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": "Invalid or missing API key",
|
||||
})
|
||||
}
|
||||
|
||||
return c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// isValidAPIKey validates the provided API key against expected value
|
||||
func isValidAPIKey(provided, expected string) bool {
|
||||
// Reject empty or whitespace-only keys
|
||||
trimmed := strings.TrimSpace(provided)
|
||||
if trimmed == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
// Reject suspicious placeholder values
|
||||
lower := strings.ToLower(trimmed)
|
||||
if lower == "null" || lower == "undefined" || lower == "nil" {
|
||||
return false
|
||||
}
|
||||
|
||||
// Compare with expected (case-sensitive, exact match)
|
||||
return provided == expected
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAPIKeyAuth(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
headerKey string
|
||||
configKey string
|
||||
wantStatus int
|
||||
}{
|
||||
{"valid key", "secret-key-123", "secret-key-123", fiber.StatusOK},
|
||||
{"missing header", "", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"whitespace only", " ", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"null string", "null", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"NULL uppercase", "NULL", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"undefined string", "undefined", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"nil string", "nil", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"wrong key", "wrong-key", "secret-key-123", fiber.StatusUnauthorized},
|
||||
{"case mismatch", "Secret-Key-123", "secret-key-123", fiber.StatusUnauthorized},
|
||||
// Note: HTTP headers with leading/trailing whitespace are trimmed by the HTTP library
|
||||
{"leading whitespace trimmed by http", " secret-key-123", "secret-key-123", fiber.StatusOK},
|
||||
{"trailing whitespace trimmed by http", "secret-key-123 ", "secret-key-123", fiber.StatusOK},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cfg := &config.Config{
|
||||
Server: config.ServerConfig{
|
||||
EnabledAuthAPI: true,
|
||||
AuthAPIKey: tt.configKey,
|
||||
},
|
||||
}
|
||||
|
||||
app := fiber.New()
|
||||
app.Use(APIKeyAuth(cfg))
|
||||
app.Get("/test", func(c *fiber.Ctx) error {
|
||||
return c.SendString("ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
if tt.headerKey != "" {
|
||||
req.Header.Set("x-osm-api-key", tt.headerKey)
|
||||
}
|
||||
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.wantStatus, resp.StatusCode)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsValidAPIKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
provided string
|
||||
expected string
|
||||
want bool
|
||||
}{
|
||||
{"exact match", "my-key", "my-key", true},
|
||||
{"empty provided", "", "my-key", false},
|
||||
{"whitespace only", " ", "my-key", false},
|
||||
{"null lowercase", "null", "my-key", false},
|
||||
{"null uppercase", "NULL", "my-key", false},
|
||||
{"null mixed case", "Null", "my-key", false},
|
||||
{"undefined lowercase", "undefined", "my-key", false},
|
||||
{"undefined uppercase", "UNDEFINED", "my-key", false},
|
||||
{"nil lowercase", "nil", "my-key", false},
|
||||
{"nil uppercase", "NIL", "my-key", false},
|
||||
{"wrong key", "other", "my-key", false},
|
||||
{"case sensitive mismatch", "My-Key", "my-key", false},
|
||||
{"leading whitespace", " my-key", "my-key", false},
|
||||
{"trailing whitespace", "my-key ", "my-key", false},
|
||||
{"both have whitespace identical", " my-key ", " my-key ", true}, // exact match even with whitespace
|
||||
{"special characters", "my-key!@#$%", "my-key!@#$%", true},
|
||||
{"long key", "this-is-a-very-long-api-key-with-many-characters-1234567890", "this-is-a-very-long-api-key-with-many-characters-1234567890", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := isValidAPIKey(tt.provided, tt.expected)
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAPIKeyAuth_ResponseBody(t *testing.T) {
|
||||
cfg := &config.Config{
|
||||
Server: config.ServerConfig{
|
||||
EnabledAuthAPI: true,
|
||||
AuthAPIKey: "test-key",
|
||||
},
|
||||
}
|
||||
|
||||
app := fiber.New()
|
||||
app.Use(APIKeyAuth(cfg))
|
||||
app.Get("/test", func(c *fiber.Ctx) error {
|
||||
return c.SendString("ok")
|
||||
})
|
||||
|
||||
// Test that invalid key returns proper error response
|
||||
req := httptest.NewRequest("GET", "/test", nil)
|
||||
req.Header.Set("x-osm-api-key", "wrong-key")
|
||||
|
||||
resp, err := app.Test(req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, fiber.StatusUnauthorized, resp.StatusCode)
|
||||
assert.Equal(t, "application/json", resp.Header.Get("Content-Type"))
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/logger"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// DebugRequestBody logs request bodies for POST/PUT/PATCH requests in debug mode
|
||||
func DebugRequestBody() fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
log := logger.Get()
|
||||
method := c.Method()
|
||||
|
||||
// Only log bodies for methods that typically have request bodies
|
||||
if method == "POST" || method == "PUT" || method == "PATCH" {
|
||||
body := c.Body()
|
||||
if len(body) > 0 {
|
||||
// Try to pretty-print JSON
|
||||
var prettyJSON bytes.Buffer
|
||||
if err := json.Indent(&prettyJSON, body, "", " "); err == nil {
|
||||
log.Debug("Request body",
|
||||
zap.String("method", method),
|
||||
zap.String("path", c.Path()),
|
||||
zap.String("body", prettyJSON.String()),
|
||||
)
|
||||
} else {
|
||||
// Not JSON or invalid JSON, log as-is (truncated if too long)
|
||||
bodyStr := string(body)
|
||||
if len(bodyStr) > 2000 {
|
||||
bodyStr = bodyStr[:2000] + "... (truncated)"
|
||||
}
|
||||
log.Debug("Request body",
|
||||
zap.String("method", method),
|
||||
zap.String("path", c.Path()),
|
||||
zap.String("body", bodyStr),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Log query parameters for all requests
|
||||
if c.Request().URI().QueryString() != nil && len(c.Request().URI().QueryString()) > 0 {
|
||||
log.Debug("Request query params",
|
||||
zap.String("method", method),
|
||||
zap.String("path", c.Path()),
|
||||
zap.String("query", string(c.Request().URI().QueryString())),
|
||||
)
|
||||
}
|
||||
|
||||
return c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// DebugErrorHandler wraps responses to log detailed error information
|
||||
func DebugErrorHandler(c *fiber.Ctx, err error) error {
|
||||
log := logger.Get()
|
||||
|
||||
code := fiber.StatusInternalServerError
|
||||
if e, ok := err.(*fiber.Error); ok {
|
||||
code = e.Code
|
||||
}
|
||||
|
||||
// Log detailed error information
|
||||
log.Error("Request error",
|
||||
zap.String("method", c.Method()),
|
||||
zap.String("path", c.Path()),
|
||||
zap.Int("status", code),
|
||||
zap.Error(err),
|
||||
zap.String("ip", c.IP()),
|
||||
zap.String("user_agent", c.Get("User-Agent")),
|
||||
)
|
||||
|
||||
// Log request body for failed POST/PUT/PATCH requests
|
||||
if c.Method() == "POST" || c.Method() == "PUT" || c.Method() == "PATCH" {
|
||||
// Re-read the body since it might have been consumed
|
||||
body := c.Body()
|
||||
if len(body) > 0 {
|
||||
var prettyJSON bytes.Buffer
|
||||
if err := json.Indent(&prettyJSON, body, "", " "); err == nil {
|
||||
log.Error("Failed request body",
|
||||
zap.String("body", prettyJSON.String()),
|
||||
)
|
||||
} else {
|
||||
bodyStr := string(body)
|
||||
if len(bodyStr) > 2000 {
|
||||
bodyStr = bodyStr[:2000] + "... (truncated)"
|
||||
}
|
||||
log.Error("Failed request body",
|
||||
zap.String("body", bodyStr),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Return detailed error response in debug mode
|
||||
return c.Status(code).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
"code": code,
|
||||
"path": c.Path(),
|
||||
"method": c.Method(),
|
||||
})
|
||||
}
|
||||
|
||||
// DebugResponseLogger logs response status for debugging
|
||||
func DebugResponseLogger() fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
err := c.Next()
|
||||
|
||||
log := logger.Get()
|
||||
status := c.Response().StatusCode()
|
||||
|
||||
// Log non-2xx responses with more detail
|
||||
if status >= 400 {
|
||||
log.Debug("Response",
|
||||
zap.String("method", c.Method()),
|
||||
zap.String("path", c.Path()),
|
||||
zap.Int("status", status),
|
||||
zap.Int("body_size", len(c.Response().Body())),
|
||||
)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Ensure body can be re-read for logging (use before DebugRequestBody)
|
||||
func BodyReusable() fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
// Store original body so it can be read multiple times
|
||||
body := c.Body()
|
||||
c.Request().SetBody(body)
|
||||
|
||||
// Also set body reader for handlers that use io.Reader
|
||||
c.Request().SetBodyStream(io.NopCloser(bytes.NewReader(body)), len(body))
|
||||
|
||||
return c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||
)
|
||||
|
||||
var (
|
||||
httpRequestsTotal = promauto.NewCounterVec(prometheus.CounterOpts{
|
||||
Name: "osmedeus_http_requests_total",
|
||||
Help: "Total HTTP requests",
|
||||
}, []string{"method", "path", "status"})
|
||||
|
||||
httpRequestDuration = promauto.NewHistogramVec(prometheus.HistogramOpts{
|
||||
Name: "osmedeus_http_request_duration_seconds",
|
||||
Help: "HTTP request duration in seconds",
|
||||
Buckets: prometheus.DefBuckets,
|
||||
}, []string{"method", "path"})
|
||||
|
||||
httpRequestsInFlight = promauto.NewGauge(prometheus.GaugeOpts{
|
||||
Name: "osmedeus_http_requests_in_flight",
|
||||
Help: "Current number of HTTP requests being processed",
|
||||
})
|
||||
)
|
||||
|
||||
// PrometheusMetrics returns a Fiber middleware that records HTTP metrics
|
||||
func PrometheusMetrics() fiber.Handler {
|
||||
return func(c *fiber.Ctx) error {
|
||||
httpRequestsInFlight.Inc()
|
||||
start := time.Now()
|
||||
|
||||
err := c.Next()
|
||||
|
||||
duration := time.Since(start).Seconds()
|
||||
httpRequestsInFlight.Dec()
|
||||
|
||||
// Normalize path to avoid high cardinality (replace IDs with :id)
|
||||
path := normalizePath(c.Route().Path)
|
||||
|
||||
httpRequestDuration.WithLabelValues(c.Method(), path).Observe(duration)
|
||||
httpRequestsTotal.WithLabelValues(c.Method(), path, strconv.Itoa(c.Response().StatusCode())).Inc()
|
||||
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// normalizePath normalizes the path to avoid high cardinality metrics
|
||||
func normalizePath(path string) string {
|
||||
if path == "" {
|
||||
return "/"
|
||||
}
|
||||
return path
|
||||
}
|
||||
@@ -0,0 +1,369 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/gofiber/adaptor/v2"
|
||||
"github.com/gofiber/fiber/v2"
|
||||
"github.com/gofiber/fiber/v2/middleware/cors"
|
||||
"github.com/gofiber/fiber/v2/middleware/favicon"
|
||||
"github.com/gofiber/fiber/v2/middleware/filesystem"
|
||||
"github.com/gofiber/fiber/v2/middleware/logger"
|
||||
"github.com/gofiber/fiber/v2/middleware/recover"
|
||||
"github.com/gofiber/swagger"
|
||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/config"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/core"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/database"
|
||||
"github.com/j3ssie/osmedeus/v5/internal/distributed"
|
||||
oslogger "github.com/j3ssie/osmedeus/v5/internal/logger"
|
||||
"github.com/j3ssie/osmedeus/v5/pkg/server/handlers"
|
||||
"github.com/j3ssie/osmedeus/v5/pkg/server/middleware"
|
||||
"github.com/j3ssie/osmedeus/v5/public"
|
||||
"go.uber.org/zap"
|
||||
|
||||
_ "github.com/j3ssie/osmedeus/v5/docs/api-swagger" // swagger docs
|
||||
)
|
||||
|
||||
// Options contains server configuration options
|
||||
type Options struct {
|
||||
NoAuth bool // Disable authentication when true
|
||||
Master *distributed.Master // Master node for distributed mode (nil if not in master mode)
|
||||
Debug bool // Enable debug mode (log request bodies, detailed errors)
|
||||
}
|
||||
|
||||
// cachedServerInfo holds server info read once at startup to avoid reading config on every request
|
||||
var cachedServerInfo struct {
|
||||
License string
|
||||
Version string
|
||||
Binary string
|
||||
Repo string
|
||||
Author string
|
||||
Docs string
|
||||
}
|
||||
|
||||
// Server represents the web server
|
||||
type Server struct {
|
||||
app *fiber.App
|
||||
config *config.Config
|
||||
options *Options
|
||||
}
|
||||
|
||||
// New creates a new server instance
|
||||
func New(cfg *config.Config, opts *Options) (*Server, error) {
|
||||
if opts == nil {
|
||||
opts = &Options{}
|
||||
}
|
||||
|
||||
// Initialize database connection once at startup
|
||||
_, err := database.Connect(cfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to connect to database: %w", err)
|
||||
}
|
||||
|
||||
// Run migrations once at startup
|
||||
ctx := context.Background()
|
||||
if err := database.Migrate(ctx); err != nil {
|
||||
return nil, fmt.Errorf("failed to run database migrations: %w", err)
|
||||
}
|
||||
oslogger.Get().Info("Database initialized", zap.String("engine", cfg.Database.DBEngine))
|
||||
|
||||
// Index workflows from filesystem to database at startup
|
||||
if cfg.WorkflowsPath != "" {
|
||||
result, err := database.IndexWorkflowsFromFilesystem(ctx, cfg.WorkflowsPath, false)
|
||||
if err != nil {
|
||||
oslogger.Get().Warn("Failed to index workflows at startup", zap.Error(err))
|
||||
} else {
|
||||
oslogger.Get().Info("Workflows indexed at startup",
|
||||
zap.Int("added", result.Added),
|
||||
zap.Int("updated", result.Updated),
|
||||
zap.Int("removed", result.Removed),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// Cache server info at startup (read once, not on every request)
|
||||
license := cfg.Server.License
|
||||
if license == "" {
|
||||
license = core.LICENSE // fallback to constant
|
||||
}
|
||||
cachedServerInfo.License = license
|
||||
cachedServerInfo.Version = core.VERSION
|
||||
cachedServerInfo.Binary = core.BINARY
|
||||
cachedServerInfo.Repo = core.REPO_URL
|
||||
cachedServerInfo.Author = core.AUTHOR
|
||||
cachedServerInfo.Docs = core.DOCS
|
||||
|
||||
// Set cached info for handlers to use
|
||||
handlers.SetServerInfo(&handlers.ServerInfoData{
|
||||
License: cachedServerInfo.License,
|
||||
Version: cachedServerInfo.Version,
|
||||
Binary: cachedServerInfo.Binary,
|
||||
Repo: cachedServerInfo.Repo,
|
||||
Author: cachedServerInfo.Author,
|
||||
Docs: cachedServerInfo.Docs,
|
||||
})
|
||||
|
||||
// Select error handler based on debug mode
|
||||
var errHandler fiber.ErrorHandler
|
||||
if opts.Debug {
|
||||
errHandler = middleware.DebugErrorHandler
|
||||
} else {
|
||||
errHandler = errorHandler
|
||||
}
|
||||
|
||||
app := fiber.New(fiber.Config{
|
||||
AppName: "Osmedeus API Server",
|
||||
ServerHeader: fmt.Sprintf("%s %s (%s)", core.BINARY, core.VERSION, license),
|
||||
ErrorHandler: errHandler,
|
||||
})
|
||||
|
||||
// Apply middleware
|
||||
app.Use(recover.New())
|
||||
app.Use(logger.New())
|
||||
app.Use(cors.New(cors.Config{
|
||||
AllowOrigins: "*",
|
||||
AllowMethods: "GET,POST,PUT,DELETE,OPTIONS,HEAD",
|
||||
AllowHeaders: "Origin,Content-Type,Accept,Authorization",
|
||||
}))
|
||||
|
||||
// Apply Prometheus metrics middleware
|
||||
app.Use(middleware.PrometheusMetrics())
|
||||
|
||||
// Apply debug middleware if debug mode is enabled
|
||||
if opts.Debug {
|
||||
app.Use(middleware.BodyReusable())
|
||||
app.Use(middleware.DebugRequestBody())
|
||||
app.Use(middleware.DebugResponseLogger())
|
||||
}
|
||||
|
||||
s := &Server{
|
||||
app: app,
|
||||
config: cfg,
|
||||
options: opts,
|
||||
}
|
||||
|
||||
// Setup routes
|
||||
s.setupRoutes()
|
||||
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// setupRoutes configures all API routes
|
||||
func (s *Server) setupRoutes() {
|
||||
// Favicon from embedded filesystem
|
||||
s.app.Use(favicon.New(favicon.Config{
|
||||
File: "favicon.ico",
|
||||
FileSystem: http.FS(public.EmbedFS),
|
||||
}))
|
||||
|
||||
// Health checks
|
||||
s.app.Get("/health", handlers.HealthCheck)
|
||||
s.app.Get("/health/ready", handlers.ReadinessCheck)
|
||||
|
||||
// Prometheus metrics endpoint
|
||||
s.app.Get("/metrics", adaptor.HTTPHandler(promhttp.Handler()))
|
||||
|
||||
// Server info (JSON version info)
|
||||
s.app.Get("/server-info", handlers.ServerInfo(s.config))
|
||||
s.app.Get("/status", handlers.ServerInfo(s.config)) // Alias for /server-info
|
||||
s.app.Get("/api/info", handlers.ServerInfo(s.config)) // Alternative endpoint for server info
|
||||
|
||||
// Swagger documentation
|
||||
s.app.Get("/swagger/*", swagger.HandlerDefault)
|
||||
|
||||
// API routes
|
||||
api := s.app.Group("/osm/api")
|
||||
|
||||
// Login - always accessible
|
||||
api.Post("/login", handlers.Login(s.config, s.options.NoAuth))
|
||||
|
||||
// Apply auth middleware conditionally
|
||||
if s.config.Server.EnabledAuthAPI {
|
||||
api.Use(middleware.APIKeyAuth(s.config))
|
||||
} else if !s.options.NoAuth {
|
||||
api.Use(middleware.JWTAuth(s.config))
|
||||
}
|
||||
|
||||
// Workflows
|
||||
api.Get("/workflows", handlers.ListWorkflowsVerbose(s.config))
|
||||
api.Get("/workflows/tags", handlers.GetAllWorkflowTags(s.config))
|
||||
api.Post("/workflows/refresh", handlers.RefreshWorkflowIndex(s.config))
|
||||
api.Get("/workflows/:name", handlers.GetWorkflowVerbose(s.config))
|
||||
|
||||
// Runs
|
||||
api.Post("/runs", handlers.CreateRun(s.config))
|
||||
api.Get("/runs", handlers.ListRuns(s.config))
|
||||
api.Get("/runs/:id", handlers.GetRun(s.config))
|
||||
api.Delete("/runs/:id", handlers.CancelRun(s.config))
|
||||
api.Get("/runs/:id/steps", handlers.GetRunSteps)
|
||||
api.Get("/runs/:id/artifacts", handlers.GetRunArtifacts)
|
||||
|
||||
// Jobs (group of runs from same request)
|
||||
api.Get("/jobs/:id", handlers.GetJobStatus(s.config))
|
||||
|
||||
// File uploads
|
||||
api.Post("/upload-file", handlers.UploadFile(s.config))
|
||||
api.Post("/workflow-upload", handlers.UploadWorkflow(s.config))
|
||||
|
||||
// Snapshots (legacy endpoint for backward compatibility)
|
||||
api.Get("/snapshot-download/:workspace_name", handlers.SnapshotDownload(s.config))
|
||||
|
||||
// Snapshots (new endpoints)
|
||||
api.Get("/snapshots", handlers.ListSnapshots(s.config))
|
||||
api.Post("/snapshots/export", handlers.SnapshotExport(s.config))
|
||||
api.Post("/snapshots/import", handlers.SnapshotImport(s.config))
|
||||
api.Delete("/snapshots/:name", handlers.DeleteSnapshot(s.config))
|
||||
|
||||
// Workspaces
|
||||
api.Get("/workspaces", handlers.ListWorkspaces(s.config))
|
||||
api.Get("/workspace-names", handlers.ListWorkspaceNames(s.config))
|
||||
|
||||
// Artifacts
|
||||
api.Get("/artifacts", handlers.ListArtifacts(s.config))
|
||||
api.Get("/artifacts/:workspace_name", handlers.DownloadWorkspaceArtifact(s.config))
|
||||
|
||||
// Assets
|
||||
api.Get("/assets", handlers.ListAssets(s.config))
|
||||
|
||||
// Vulnerabilities
|
||||
api.Get("/vulnerabilities", handlers.ListVulnerabilities(s.config))
|
||||
api.Get("/vulnerabilities/summary", handlers.GetVulnerabilitySummary(s.config))
|
||||
api.Get("/vulnerabilities/:id", handlers.GetVulnerability(s.config))
|
||||
api.Post("/vulnerabilities", handlers.CreateVulnerability(s.config))
|
||||
api.Delete("/vulnerabilities/:id", handlers.DeleteVulnerability(s.config))
|
||||
|
||||
// Stats
|
||||
api.Get("/stats", handlers.GetSystemStats(s.config))
|
||||
|
||||
// Install - registry info and installation endpoints
|
||||
api.Get("/registry-info", handlers.GetRegistryInfo(s.config))
|
||||
api.Post("/registry-install", handlers.RegistryInstall(s.config))
|
||||
|
||||
// Schedules
|
||||
api.Get("/schedules", handlers.ListSchedules(s.config))
|
||||
api.Post("/schedules", handlers.CreateSchedule(s.config))
|
||||
api.Get("/schedules/:id", handlers.GetSchedule(s.config))
|
||||
api.Put("/schedules/:id", handlers.UpdateSchedule(s.config))
|
||||
api.Delete("/schedules/:id", handlers.DeleteSchedule(s.config))
|
||||
api.Post("/schedules/:id/enable", handlers.EnableSchedule(s.config))
|
||||
api.Post("/schedules/:id/disable", handlers.DisableSchedule(s.config))
|
||||
api.Post("/schedules/:id/trigger", handlers.TriggerSchedule(s.config))
|
||||
|
||||
// Event logs
|
||||
api.Get("/event-logs", handlers.ListEventLogs(s.config))
|
||||
|
||||
// Functions
|
||||
api.Post("/functions/eval", handlers.FunctionEval(s.config))
|
||||
api.Get("/functions/list", handlers.FunctionList(s.config))
|
||||
|
||||
// Settings API
|
||||
api.Get("/settings/yaml", handlers.GetSettingsYAML(s.config))
|
||||
api.Get("/settings/yaml/", handlers.GetSettingsYAML(s.config))
|
||||
|
||||
// LLM endpoints (OpenAI-compatible)
|
||||
api.Post("/llm/v1/chat/completions", handlers.LLMChat(s.config))
|
||||
api.Post("/llm/v1/embeddings", handlers.LLMEmbedding(s.config))
|
||||
|
||||
// Distributed endpoints (only available when running in master mode)
|
||||
if s.options.Master != nil {
|
||||
api.Get("/workers", handlers.ListWorkers(s.options.Master))
|
||||
api.Get("/workers/:id", handlers.GetWorker(s.options.Master))
|
||||
api.Get("/tasks", handlers.ListTasks(s.options.Master))
|
||||
api.Get("/tasks/:id", handlers.GetTask(s.options.Master))
|
||||
api.Post("/tasks", handlers.SubmitTask(s.options.Master))
|
||||
}
|
||||
|
||||
// Serve workspace files under /ws/{workspace_prefix_key}/
|
||||
// Allows direct access to run outputs in workspaces directory (no auth required)
|
||||
if s.config.Server.WorkspacePrefixKey != "" && s.config.WorkspacesPath != "" {
|
||||
wsPath := fmt.Sprintf("/ws/%s", s.config.Server.WorkspacePrefixKey)
|
||||
s.app.Static(wsPath, s.config.WorkspacesPath, fiber.Static{
|
||||
Browse: true, // Enable directory listing
|
||||
})
|
||||
}
|
||||
|
||||
// Handle HEAD requests for UI routes (filesystem middleware only handles GET)
|
||||
// This middleware intercepts HEAD requests to non-API paths and returns 200 OK
|
||||
s.app.Use(func(c *fiber.Ctx) error {
|
||||
if c.Method() == fiber.MethodHead {
|
||||
// For HEAD requests to non-API paths, return 200 OK
|
||||
// This handles UI routes that the filesystem middleware would skip
|
||||
if !strings.HasPrefix(c.Path(), "/osm/api") &&
|
||||
!strings.HasPrefix(c.Path(), "/health") &&
|
||||
!strings.HasPrefix(c.Path(), "/metrics") &&
|
||||
!strings.HasPrefix(c.Path(), "/swagger") {
|
||||
c.Set("Content-Type", "text/html; charset=utf-8")
|
||||
return c.SendStatus(fiber.StatusOK)
|
||||
}
|
||||
}
|
||||
return c.Next()
|
||||
})
|
||||
|
||||
// Serve UI at root /
|
||||
// Priority: external UI path > embedded UI
|
||||
if s.config.UIPath != "" {
|
||||
if _, err := os.Stat(s.config.UIPath); err == nil {
|
||||
s.app.Use("/", filesystem.New(filesystem.Config{
|
||||
Root: http.Dir(s.config.UIPath),
|
||||
Index: "index.html",
|
||||
Browse: false,
|
||||
NotFoundFile: "index.html", // SPA fallback: serve index.html for client-side routes
|
||||
}))
|
||||
} else {
|
||||
// Fallback to embedded UI if external path doesn't exist
|
||||
s.serveEmbeddedUI()
|
||||
}
|
||||
} else {
|
||||
// Use embedded UI when no external path configured
|
||||
s.serveEmbeddedUI()
|
||||
}
|
||||
}
|
||||
|
||||
// serveEmbeddedUI serves the embedded UI files at root /
|
||||
func (s *Server) serveEmbeddedUI() {
|
||||
uiFS, err := public.GetUIFS()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
s.app.Use("/", filesystem.New(filesystem.Config{
|
||||
Root: http.FS(uiFS),
|
||||
Index: "index.html",
|
||||
Browse: false,
|
||||
PathPrefix: "",
|
||||
NotFoundFile: "index.html", // SPA fallback: serve index.html for client-side routes
|
||||
}))
|
||||
}
|
||||
|
||||
// Start starts the server
|
||||
func (s *Server) Start(addr string) error {
|
||||
return s.app.Listen(addr)
|
||||
}
|
||||
|
||||
// Shutdown gracefully shuts down the server
|
||||
func (s *Server) Shutdown() error {
|
||||
return s.app.Shutdown()
|
||||
}
|
||||
|
||||
// ShutdownWithContext gracefully shuts down the server with a context for timeout
|
||||
func (s *Server) ShutdownWithContext(ctx context.Context) error {
|
||||
return s.app.ShutdownWithContext(ctx)
|
||||
}
|
||||
|
||||
// errorHandler handles errors globally
|
||||
func errorHandler(c *fiber.Ctx, err error) error {
|
||||
code := fiber.StatusInternalServerError
|
||||
|
||||
if e, ok := err.(*fiber.Error); ok {
|
||||
code = e.Code
|
||||
}
|
||||
|
||||
return c.Status(code).JSON(fiber.Map{
|
||||
"error": true,
|
||||
"message": err.Error(),
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user