From 28aba2c07c6a41026acc586b6361d26611d4187e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Gabriel=20Mu=C3=B1oz?= Date: Thu, 1 Oct 2026 16:16:26 -0700 Subject: [PATCH 1/2] Prevent SQL injection in FlameDB queries Parameterize request and RQL filter values passed to ClickHouse instead of interpolating them into SQL strings. Allowlist dynamic query fields and use QueryContext for query execution. Fixes #152 --- src/gprofiler_flamedb_rest/common/params.go | 5 + src/gprofiler_flamedb_rest/db/clickhouse.go | 359 +++++++++++------- src/gprofiler_flamedb_rest/handlers/common.go | 66 ++-- .../handlers/common_test.go | 50 ++- .../handlers/handlers.go | 42 +- 5 files changed, 320 insertions(+), 202 deletions(-) diff --git a/src/gprofiler_flamedb_rest/common/params.go b/src/gprofiler_flamedb_rest/common/params.go index 1af963f8..0fa110e4 100644 --- a/src/gprofiler_flamedb_rest/common/params.go +++ b/src/gprofiler_flamedb_rest/common/params.go @@ -41,6 +41,11 @@ type AllFiltersParams struct { K8SObject []string `form:"k8s_obj"` } +type QueryFilter struct { + Clause string + Args []any +} + type FiltersParams struct { ContainerName string `rql:"column=ContainerName,filter"` HostName string `rql:"column=HostName,filter"` diff --git a/src/gprofiler_flamedb_rest/db/clickhouse.go b/src/gprofiler_flamedb_rest/db/clickhouse.go index f641531a..efd4e23d 100644 --- a/src/gprofiler_flamedb_rest/db/clickhouse.go +++ b/src/gprofiler_flamedb_rest/db/clickhouse.go @@ -24,14 +24,15 @@ import ( "log" "math" "regexp" - "restflamedb/common" - "restflamedb/config" "sort" "strings" "sync" "time" "github.com/ClickHouse/clickhouse-go" + + "restflamedb/common" + "restflamedb/config" ) const ( @@ -69,6 +70,12 @@ var ( "day": "1day", "raw": "raw", } + queryFields = map[string]struct{}{ + "ContainerName": {}, + "HostName": {}, + "InstanceType": {}, + "ContainerEnvName": {}, + } ) type ClickHouseClient struct { @@ -211,23 +218,23 @@ func GetTimeRanges(start time.Time, end time.Time, resolution string) map[string delta := end.Sub(start) fullRange := makeTimeRange(start, trimEndTime(end)) now := time.Now().UTC() - + // Retention thresholds from configurable settings rawRetentionInterval := time.Hour * 24 * time.Duration(config.RawRetentionDays) // Raw data TTL hourlyRetentionInterval := time.Hour * 24 * time.Duration(config.HourlyRetentionDays) // Hourly data TTL dailyThreshold := time.Hour * 24 * time.Duration(config.HourlyRetentionDays) // Switch to daily aggregation - + // Debug logging for table selection dataAge := now.Sub(start) - log.Printf("📊 Table Selection: range=%v, delta=%v, age=%v, resolution=%s", + log.Printf("📊 Table Selection: range=%v, delta=%v, age=%v, resolution=%s", start.Format("2006-01-02"), delta, dataAge, resolution) - + // For very old data (>90 days), use daily aggregation with day boundaries if now.Sub(start) >= dailyThreshold { result["1day_historical"] = append(result["1day_historical"], makeTimeRange(makeStartOfDay(start), makeEndOfDay(end))) return result } - + // For 7-90 day old data, raw is expired but hourly is available if now.Sub(start) >= rawRetentionInterval && now.Sub(start) < hourlyRetentionInterval { switch resolution { @@ -281,40 +288,50 @@ func GetTimeRanges(start time.Time, end time.Time, resolution string) map[string } func BuildConditions(ContainerName []string, HostName []string, InstanceType []string, K8SObject []string, - filterQuery string) (string, string) { - if filterQuery != "" { + filterQuery common.QueryFilter) (string, common.QueryFilter) { + if filterQuery.Clause != "" { return "", filterQuery } - var conditions string + conditions := common.QueryFilter{} tablePrefix := "_all" if len(ContainerName) > 0 { - containerNamesHash := make([]string, 0) + containerNamesHash := make([]uint32, 0, len(ContainerName)) for _, singleContainerName := range ContainerName { - containerNamesHash = append(containerNamesHash, fmt.Sprint(common.GetHash32AsInt(singleContainerName))) + containerNamesHash = append(containerNamesHash, common.GetHash32AsInt(singleContainerName)) } - conditions += fmt.Sprintf(" AND (ContainerNameHash IN (%s))", strings.Join(containerNamesHash, ",")) + conditions.Clause += " AND (ContainerNameHash IN (@container_names))" + conditions.Args = append(conditions.Args, sql.Named("container_names", containerNamesHash)) tablePrefix = "" } if len(HostName) > 0 { - hostNamesHash := make([]string, 0) + hostNamesHash := make([]uint32, 0, len(HostName)) for _, singleHostName := range HostName { - hostNamesHash = append(hostNamesHash, fmt.Sprint(common.GetHash32AsInt(singleHostName))) + hostNamesHash = append(hostNamesHash, common.GetHash32AsInt(singleHostName)) } - conditions += fmt.Sprintf(" AND (HostNameHash IN (%s))", strings.Join(hostNamesHash, ",")) + conditions.Clause += " AND (HostNameHash IN (@host_names))" + conditions.Args = append(conditions.Args, sql.Named("host_names", hostNamesHash)) tablePrefix = "" } if len(InstanceType) > 0 { tablePrefix = "" - conditions += fmt.Sprintf(" AND (InstanceType IN ('%s'))", strings.Join(InstanceType, "','")) + conditions.Clause += " AND (InstanceType IN (@instance_types))" + conditions.Args = append(conditions.Args, sql.Named("instance_types", InstanceType)) } if len(K8SObject) > 0 { tablePrefix = "" - conditions += fmt.Sprintf(" AND (ContainerEnvName IN ('%s'))", strings.Join(K8SObject, "','")) + conditions.Clause += " AND (ContainerEnvName IN (@k8s_objects))" + conditions.Args = append(conditions.Args, sql.Named("k8s_objects", K8SObject)) } return tablePrefix, conditions } +func withFilterArgs(filterQuery common.QueryFilter, args ...any) []any { + queryArgs := make([]any, 0, len(filterQuery.Args)+len(args)) + queryArgs = append(queryArgs, filterQuery.Args...) + return append(queryArgs, args...) +} + func NewClickHouseClient(addr string) *ClickHouseClient { db, err := sql.Open("clickhouse", "tcp://"+addr) if err != nil { @@ -378,7 +395,7 @@ func scanFrames(rows *sql.Rows, frames map[uint64]Frame, collapsed bool) (int, i } func (c *ClickHouseClient) GetTopFrames(ctx context.Context, params common.FlameGraphParams, - filterQuery string) (Graph, error) { + filterQuery common.QueryFilter) (Graph, error) { var wg sync.WaitGroup var mutex sync.Mutex queryErrors := make([]error, 0) @@ -398,16 +415,21 @@ func (c *ClickHouseClient) GetTopFrames(ctx context.Context, params common.Flame query := fmt.Sprintf(` SELECT CallStackHash, any(CallStackName), any(CallStackParent), sum(NumSamples) AS SumNumSamples FROM %s - WHERE ServiceId == '%d' AND (Timestamp BETWEEN '%s' AND '%s') %s - GROUP BY CallStackHash - ORDER BY SumNumSamples DESC - LIMIT %d`, sTable, params.ServiceId, sStart, sEnd, conditions, params.StacksNum) - + WHERE ServiceId == @service_id AND (Timestamp BETWEEN @start_time AND @end_time) %s + GROUP BY CallStackHash + ORDER BY SumNumSamples DESC + LIMIT @stacks_num`, sTable, conditions.Clause) + queryStart := time.Now() log.Printf("🚀 Service %d: Starting query on %s (%s to %s)", params.ServiceId, sTable, sStart, sEnd) - rows, err := c.client.Query(query) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", sStart), + sql.Named("end_time", sEnd), + sql.Named("stacks_num", params.StacksNum), + )...) queryDuration := time.Since(queryStart) - + if err == nil { defer rows.Close() frames := make(map[uint64]Frame) @@ -447,19 +469,22 @@ func (c *ClickHouseClient) GetTopFrames(ctx context.Context, params common.Flame } func (c *ClickHouseClient) FetchInstanceTypeCount(ctx context.Context, params common.QueryParams, - filterQuery string) []common.InstanceTypeCount { + filterQuery common.QueryFilter) []common.InstanceTypeCount { var selectQuery string result := make([]common.InstanceTypeCount, 0) _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) selectQuery = ` SELECT InstanceType, COUNT(DISTINCT HostName) as InstanceCount - FROM flamedb.samples_1min where ServiceId = '%d' AND (Timestamp BETWEEN '%s' AND '%s' ) + FROM flamedb.samples_1min where ServiceId = @service_id + AND (Timestamp BETWEEN @start_time AND @end_time) %s GROUP BY InstanceType ORDER BY InstanceCount DESC` - query := fmt.Sprintf(selectQuery, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), conditions) - - rows, err := c.client.Query(query) + query := fmt.Sprintf(selectQuery, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -480,17 +505,24 @@ func (c *ClickHouseClient) FetchInstanceTypeCount(ctx context.Context, params co } func (c *ClickHouseClient) FetchFieldValueSample(ctx context.Context, field string, params common.QueryParams, - filterQuery string) []common.FilterData { + filterQuery common.QueryFilter) []common.FilterData { var selectQuery string result := make([]common.FilterData, 0) + if _, ok := queryFields[field]; !ok { + log.Printf("unsupported query field %q", field) + return result + } _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) selectQuery = ` - SELECT %s, SUM(NumSamples) as samples from flamedb.samples_1min WHERE ServiceId == '%d' AND - (Timestamp BETWEEN '%s' AND '%s') %s GROUP BY %s ORDER BY samples DESC;` - query := fmt.Sprintf(selectQuery, field, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), conditions, field) - - rows, err := c.client.Query(query) + SELECT %s, SUM(NumSamples) as samples from flamedb.samples_1min WHERE ServiceId == @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s GROUP BY %s ORDER BY samples DESC;` + query := fmt.Sprintf(selectQuery, field, conditions.Clause, field) + + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -512,16 +544,23 @@ func (c *ClickHouseClient) FetchFieldValueSample(ctx context.Context, field stri } func (c *ClickHouseClient) FetchFieldValues(ctx context.Context, field string, params common.QueryParams, - filterQuery string) []common.FilterData { + filterQuery common.QueryFilter) []common.FilterData { result := make([]common.FilterData, 0) + if _, ok := queryFields[field]; !ok { + log.Printf("unsupported query field %q", field) + return result + } _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) query := fmt.Sprintf(` - SELECT %s from flamedb.samples_1min WHERE ServiceId == '%d' AND - (Timestamp BETWEEN '%s' AND '%s') %s GROUP BY %s; - `, field, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), conditions, field) - - rows, err := c.client.Query(query) + SELECT %s from flamedb.samples_1min WHERE ServiceId == @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s GROUP BY %s; + `, field, conditions.Clause, field) + + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -542,19 +581,23 @@ func (c *ClickHouseClient) FetchFieldValues(ctx context.Context, field string, p } func (c *ClickHouseClient) FetchSampleCount(ctx context.Context, params common.QueryParams, - filterQuery string) []common.Sample { + filterQuery common.QueryFilter) []common.Sample { _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) interval := getInterval(params.StartDateTime, params.EndDateTime, params.Interval) result := make([]common.Sample, 0) query := fmt.Sprintf(` - SELECT toStartOfInterval(Timestamp, INTERVAL '%s') as Datetime, SUM(NumSamples) + SELECT toStartOfInterval(Timestamp, INTERVAL @interval) as Datetime, SUM(NumSamples) FROM flamedb.samples_1min - WHERE ServiceId == '%d' AND (Timestamp BETWEEN '%s' AND '%s') %s + WHERE ServiceId == @service_id AND (Timestamp BETWEEN @start_time AND @end_time) %s GROUP BY Datetime ORDER BY Datetime DESC; - `, interval, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), conditions) - rows, err := c.client.Query(query) + `, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("interval", interval), + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -575,7 +618,7 @@ func (c *ClickHouseClient) FetchSampleCount(ctx context.Context, params common.Q } func (c *ClickHouseClient) FetchSampleCountByFunction(ctx context.Context, params common.QueryParams, - filterQuery string) []common.SamplesCountByFunction { + filterQuery common.QueryFilter) []common.SamplesCountByFunction { _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) interval := getInterval(params.StartDateTime, params.EndDateTime, "") if interval == "15 second" || interval == "30 second" { @@ -584,15 +627,16 @@ func (c *ClickHouseClient) FetchSampleCountByFunction(ctx context.Context, param result := make([]common.SamplesCountByFunction, 0) query := fmt.Sprintf(` WITH all_samples as( - SELECT toStartOfInterval(Timestamp, INTERVAL '%s') AS Datetime, SUM(NumSamples) AS sum_cpu + SELECT toStartOfInterval(Timestamp, INTERVAL @interval) AS Datetime, SUM(NumSamples) AS sum_cpu FROM flamedb.samples_1min - WHERE ServiceId == '%d' AND (Timestamp BETWEEN '%s' AND '%s') %s + WHERE ServiceId == @service_id AND (Timestamp BETWEEN @start_time AND @end_time) %s GROUP BY Datetime ORDER BY Datetime DESC ), function_samples AS ( - SELECT toStartOfInterval(Timestamp, INTERVAL '%s') AS Datetime, SUM(NumSamples) AS sum_cpu + SELECT toStartOfInterval(Timestamp, INTERVAL @interval) AS Datetime, SUM(NumSamples) AS sum_cpu FROM flamedb.samples - WHERE ServiceId == '%d' AND (Timestamp BETWEEN '%s' AND '%s') AND (CallStackName = '%s') %s + WHERE ServiceId == @service_id AND (Timestamp BETWEEN @start_time AND @end_time) + AND (CallStackName = @function_name) %s GROUP BY Datetime ORDER BY Datetime DESC ) @@ -600,11 +644,15 @@ func (c *ClickHouseClient) FetchSampleCountByFunction(ctx context.Context, param SELECT (function_samples.sum_cpu/all_samples.sum_cpu) AS Samples , all_samples.Datetime AS Datetime FROM all_samples LEFT JOIN function_samples ON function_samples.Datetime = all_samples.Datetime; - `, interval, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), conditions, interval, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), params.FunctionName, conditions) - - rows, err := c.client.Query(query) + `, conditions.Clause, conditions.Clause) + + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("interval", interval), + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + sql.Named("function_name", params.FunctionName), + )...) if err == nil { defer rows.Close() @@ -638,7 +686,8 @@ func (c *ClickHouseClient) FetchSampleCountByFunction(ctx context.Context, param return result } -func (c *ClickHouseClient) FetchTimes(ctx context.Context, params common.QueryParams, filterQuery string) []string { +func (c *ClickHouseClient) FetchTimes(ctx context.Context, params common.QueryParams, + filterQuery common.QueryFilter) []string { var interval string result := make([]string, 0) _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) @@ -654,13 +703,17 @@ func (c *ClickHouseClient) FetchTimes(ctx context.Context, params common.QueryPa interval = getInterval(params.StartDateTime, params.EndDateTime, "") } query := fmt.Sprintf(` - SELECT toStartOfInterval(Timestamp, INTERVAL '%s') as Datetime - from flamedb.samples_1min WHERE ServiceId == '%d' AND - (Timestamp BETWEEN '%s' AND '%s') %s + SELECT toStartOfInterval(Timestamp, INTERVAL @interval) as Datetime + from flamedb.samples_1min WHERE ServiceId == @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s GROUP BY Datetime - ORDER BY Datetime DESC;`, interval, params.ServiceId, - common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), conditions) - rows, err := c.client.Query(query) + ORDER BY Datetime DESC;`, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("interval", interval), + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -679,17 +732,21 @@ func (c *ClickHouseClient) FetchTimes(ctx context.Context, params common.QueryPa return result } -func (c *ClickHouseClient) FetchTimeRange(ctx context.Context, params common.QueryParams, filterQuery string) []string { +func (c *ClickHouseClient) FetchTimeRange(ctx context.Context, params common.QueryParams, + filterQuery common.QueryFilter) []string { result := make([]string, 0) _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) query := fmt.Sprintf(` SELECT min(Timestamp), max(Timestamp) from flamedb.samples_1min WHERE - ServiceId == '%d' AND - (Timestamp BETWEEN '%s' AND '%s') %s;`, params.ServiceId, - common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), conditions) - rows, err := c.client.Query(query) + ServiceId == @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s;`, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -714,23 +771,27 @@ func (c *ClickHouseClient) FetchTimeRange(ctx context.Context, params common.Que } func (c *ClickHouseClient) FetchMetricsSummary(ctx context.Context, params common.MetricsSummaryParams, - filterQuery string) (common.MetricsSummary, error) { + filterQuery common.QueryFilter) (common.MetricsSummary, error) { defaultEmptyList := make([]string, 0) _, conditions := BuildConditions(defaultEmptyList, params.HostName, params.InstanceType, defaultEmptyList, filterQuery) percentile := float64(params.Percentile) / 100.0 query := fmt.Sprintf(` SELECT arrayAvg(flatten(groupArray(CPUArray))), MAX(MaxCPU), - AVG(MaxMemory), MAX(MaxMemory), quantile(%f)(MaxMemory), count() FROM + AVG(MaxMemory), MAX(MaxMemory), quantile(@percentile)(MaxMemory), count() FROM (SELECT MAX(MemoryAverageUsedPercent) AS MaxMemory, MAX(CPUAverageUsedPercent) as MaxCPU, groupArray(CPUAverageUsedPercent) as CPUArray FROM %s - WHERE ServiceId = %d AND (Timestamp BETWEEN '%s' AND '%s') %s - GROUP BY HostName)`, percentile, config.ClickHouseMetricsTable, params.ServiceId, - common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), conditions) - rows, err := c.client.Query(query) + WHERE ServiceId = @service_id AND (Timestamp BETWEEN @start_time AND @end_time) %s + GROUP BY HostName)`, config.ClickHouseMetricsTable, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("percentile", percentile), + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer func(rows *sql.Rows) { err := rows.Close() @@ -771,7 +832,6 @@ func (c *ClickHouseClient) FetchMetricsSummary(ctx context.Context, params commo func (c *ClickHouseClient) FetchMetricsServicesListSummary(ctx context.Context, params common.MetricsServicesListSummaryParams) ([]common.MetricsServicesListSummary, error) { - formattedServicesList := joinIntSlice(params.ServicesList, ",") percentile := float64(params.Percentile) / 100.0 query := fmt.Sprintf(` @@ -779,7 +839,7 @@ func (c *ClickHouseClient) FetchMetricsServicesListSummary(ctx context.Context, SELECT ServiceId as s_id, max(Timestamp) as last_seen FROM %s - WHERE ServiceId in (%s) AND (Timestamp BETWEEN '%s' AND '%s') + WHERE ServiceId in (@service_ids) AND (Timestamp BETWEEN @start_time AND @end_time) GROUP BY ServiceId ), GroupedMetrics AS ( SELECT @@ -787,18 +847,21 @@ func (c *ClickHouseClient) FetchMetricsServicesListSummary(ctx context.Context, max(MemoryAverageUsedPercent) AS MaxMemory, max(CPUAverageUsedPercent) as MaxCPU, groupArray(CPUAverageUsedPercent) as CPUArray - FROM %s - GLOBAL JOIN LatestServices ON ServiceId = s_id - WHERE ServiceId in (%s) AND (Timestamp BETWEEN last_seen - toIntervalHour(24) AND last_seen) - GROUP BY HostName, ServiceId + FROM %s + GLOBAL JOIN LatestServices ON ServiceId = s_id + WHERE ServiceId in (@service_ids) AND (Timestamp BETWEEN last_seen - toIntervalHour(24) AND last_seen) + GROUP BY HostName, ServiceId + ) + SELECT arrayAvg(flatten(groupArray(CPUArray))), max(MaxCPU), ServiceId, + avg(MaxMemory), max(MaxMemory), quantile(@percentile)(MaxMemory), count() + FROM GroupedMetrics + GROUP BY ServiceId`, config.ClickHouseMetricsTable, config.ClickHouseMetricsTable) + rows, err := c.client.QueryContext(ctx, query, + sql.Named("service_ids", params.ServicesList), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + sql.Named("percentile", percentile), ) - SELECT arrayAvg(flatten(groupArray(CPUArray))), max(MaxCPU), ServiceId, - avg(MaxMemory), max(MaxMemory), quantile(%f)(MaxMemory), count() - FROM GroupedMetrics - GROUP BY ServiceId`, config.ClickHouseMetricsTable, formattedServicesList, - common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), - config.ClickHouseMetricsTable, formattedServicesList, percentile) - rows, err := c.client.Query(query) var results []common.MetricsServicesListSummary @@ -845,7 +908,7 @@ func (c *ClickHouseClient) FetchMetricsServicesListSummary(ctx context.Context, } func (c *ClickHouseClient) FetchMetricsGraph(ctx context.Context, params common.MetricsSummaryParams, - filterQuery string) ([]common.MetricsSummary, error) { + filterQuery common.QueryFilter) ([]common.MetricsSummary, error) { defaultEmptyList := make([]string, 0) _, conditions := BuildConditions(defaultEmptyList, params.HostName, params.InstanceType, defaultEmptyList, filterQuery) @@ -855,24 +918,29 @@ func (c *ClickHouseClient) FetchMetricsGraph(ctx context.Context, params common. percentile := float64(params.Percentile) / 100.0 groupBy := "" - if params.GroupBy != "none" { - groupBy = fmt.Sprintf(", %s", params.GroupBy) + if params.GroupBy == "instance_type" { + groupBy = ", instance_type" } query := fmt.Sprintf(` SELECT Datetime %s, arrayAvg(flatten(groupArray(CPUArray))), MAX(MaxCPU), - AVG(MaxMemory), MAX(MaxMemory), quantile(%f)(MaxMemory) FROM - (SELECT toStartOfInterval(Timestamp, INTERVAL '%s') as + AVG(MaxMemory), MAX(MaxMemory), quantile(@percentile)(MaxMemory) FROM + (SELECT toStartOfInterval(Timestamp, INTERVAL @interval) as Datetime %s, HostName, MAX(MemoryAverageUsedPercent) AS MaxMemory, MAX(CPUAverageUsedPercent) as MaxCPU, groupArray(CPUAverageUsedPercent) as CPUArray FROM %s - WHERE ServiceId = %d AND (Datetime BETWEEN '%s' AND '%s') %s + WHERE ServiceId = @service_id AND (Datetime BETWEEN @start_time AND @end_time) %s GROUP BY Datetime %s, HostName) GROUP BY Datetime %s ORDER BY Datetime DESC; - `, groupBy, percentile, interval, groupBy, config.ClickHouseMetricsTable, params.ServiceId, - common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), conditions, groupBy, groupBy) - rows, err := c.client.Query(query) + `, groupBy, groupBy, config.ClickHouseMetricsTable, conditions.Clause, groupBy, groupBy) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("percentile", percentile), + sql.Named("interval", interval), + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer func(rows *sql.Rows) { err := rows.Close() @@ -916,7 +984,7 @@ func (c *ClickHouseClient) FetchMetricsGraph(ctx context.Context, params common. } func (c *ClickHouseClient) FetchMetricsCpuTrend(ctx context.Context, params common.MetricsCpuTrendParams, - filterQuery string) (common.MetricsCpuTrend, error) { + filterQuery common.QueryFilter) (common.MetricsCpuTrend, error) { defaultEmptyList := make([]string, 0) finalResult := common.MetricsCpuTrend{} _, conditions := BuildConditions(defaultEmptyList, params.HostName, params.InstanceType, defaultEmptyList, filterQuery) @@ -932,12 +1000,13 @@ func (c *ClickHouseClient) FetchMetricsCpuTrend(ctx context.Context, params comm FROM (SELECT MAX(MemoryAverageUsedPercent) AS MaxMemory, - MAX(CPUAverageUsedPercent) as MaxCPU, - groupArray(CPUAverageUsedPercent) as CPUArray - FROM %s - WHERE ServiceId = %d AND (Timestamp BETWEEN '%s' AND '%s') %s - GROUP BY HostName) - ), + MAX(CPUAverageUsedPercent) as MaxCPU, + groupArray(CPUAverageUsedPercent) as CPUArray + FROM %s + WHERE ServiceId = @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s + GROUP BY HostName) + ), PREVIOUS_CONSUMPTION AS ( SELECT arrayAvg(flatten(groupArray(CPUArray))) AS avg_cpu, @@ -948,24 +1017,29 @@ func (c *ClickHouseClient) FetchMetricsCpuTrend(ctx context.Context, params comm FROM (SELECT MAX(MemoryAverageUsedPercent) AS MaxMemory, - MAX(CPUAverageUsedPercent) as MaxCPU, - groupArray(CPUAverageUsedPercent) as CPUArray - FROM %s - WHERE ServiceId = %d AND (Timestamp BETWEEN '%s' AND '%s') %s - GROUP BY HostName) - ) + MAX(CPUAverageUsedPercent) as MaxCPU, + groupArray(CPUAverageUsedPercent) as CPUArray + FROM %s + WHERE ServiceId = @service_id AND + (Timestamp BETWEEN @compared_start_time AND @compared_end_time) %s + GROUP BY HostName) + ) SELECT avg_cpu, max_cpu, avg_memory, max_memory FROM ( - SELECT * FROM CURRENT_CONSUMPTION - UNION ALL - SELECT * FROM PREVIOUS_CONSUMPTION) - order by SortOrder`, config.ClickHouseMetricsTable, params.ServiceId, - common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), conditions, - config.ClickHouseMetricsTable, params.ServiceId, - common.FormatTime(params.ComparedStartDateTime), common.FormatTime(params.ComparedEndDateTime), conditions) + SELECT * FROM CURRENT_CONSUMPTION + UNION ALL + SELECT * FROM PREVIOUS_CONSUMPTION) + order by SortOrder`, config.ClickHouseMetricsTable, conditions.Clause, + config.ClickHouseMetricsTable, conditions.Clause) first := true - rows, err := c.client.Query(query) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + sql.Named("compared_start_time", params.ComparedStartDateTime), + sql.Named("compared_end_time", params.ComparedEndDateTime), + )...) if err == nil { defer func(rows *sql.Rows) { err := rows.Close() @@ -1019,9 +1093,13 @@ func (c *ClickHouseClient) FetchServices(ctx context.Context, params common.Serv groupByExpr = "GROUP BY (ServiceId,ContainerEnvName)" } query := fmt.Sprintf(` - SELECT %s from flamedb.samples_1min WHERE (Timestamp BETWEEN '%s' AND '%s') %s; - `, expr, common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), groupByExpr) - rows, err := c.client.Query(query) + SELECT %s from flamedb.samples_1min + WHERE (Timestamp BETWEEN @start_time AND @end_time) %s; + `, expr, groupByExpr) + rows, err := c.client.QueryContext(ctx, query, + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + ) result := make([]SrvResp, 0) if err == nil { defer rows.Close() @@ -1048,15 +1126,18 @@ func (c *ClickHouseClient) FetchServices(ctx context.Context, params common.Serv } func (c *ClickHouseClient) FetchSessionsCount(ctx context.Context, params common.SessionsCountParams, - filterQuery string) (int, error) { + filterQuery common.QueryFilter) (int, error) { _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) query := fmt.Sprintf(` - SELECT uniq(HostName,Timestamp) FROM flamedb.samples_1min WHERE ServiceId = %d AND - (Timestamp BETWEEN '%s' AND '%s') %s; - `, params.ServiceId, common.FormatTime(params.StartDateTime), - common.FormatTime(params.EndDateTime), conditions) - rows, err := c.client.Query(query) + SELECT uniq(HostName,Timestamp) FROM flamedb.samples_1min WHERE ServiceId = @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s; + `, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { @@ -1075,13 +1156,17 @@ func (c *ClickHouseClient) FetchSessionsCount(ctx context.Context, params common } func (c *ClickHouseClient) FetchLastHTML(ctx context.Context, params common.MetricsLastHTMLParams, - filterQuery string) (string, error) { + filterQuery common.QueryFilter) (string, error) { _, conditions := BuildConditions(params.ContainerName, params.HostName, params.InstanceType, params.K8SObject, filterQuery) query := fmt.Sprintf(` - SELECT argMax(HTMLPath,Timestamp) FROM flamedb.metrics WHERE ServiceId = %d AND - (Timestamp BETWEEN '%s' AND '%s') %s; - `, params.ServiceId, common.FormatTime(params.StartDateTime), common.FormatTime(params.EndDateTime), conditions) - rows, err := c.client.Query(query) + SELECT argMax(HTMLPath,Timestamp) FROM flamedb.metrics WHERE ServiceId = @service_id AND + (Timestamp BETWEEN @start_time AND @end_time) %s; + `, conditions.Clause) + rows, err := c.client.QueryContext(ctx, query, withFilterArgs(conditions, + sql.Named("service_id", params.ServiceId), + sql.Named("start_time", params.StartDateTime), + sql.Named("end_time", params.EndDateTime), + )...) if err == nil { defer rows.Close() for rows.Next() { diff --git a/src/gprofiler_flamedb_rest/handlers/common.go b/src/gprofiler_flamedb_rest/handlers/common.go index c28e2ba7..f96a0798 100644 --- a/src/gprofiler_flamedb_rest/handlers/common.go +++ b/src/gprofiler_flamedb_rest/handlers/common.go @@ -17,15 +17,18 @@ package handlers import ( + "database/sql" "fmt" - "github.com/a8m/rql" - "github.com/gin-gonic/gin" - "log" "net/http" "reflect" "strconv" "strings" "time" + + "github.com/a8m/rql" + "github.com/gin-gonic/gin" + + "restflamedb/common" ) func StartTime() gin.HandlerFunc { @@ -35,12 +38,12 @@ func StartTime() gin.HandlerFunc { } } -func parseParams[T any](params T, parser *rql.Parser, c *gin.Context) (T, string, error) { - var query string +func parseParams[T any](params T, parser *rql.Parser, c *gin.Context) (T, common.QueryFilter, error) { + var filterQuery common.QueryFilter var err error if err = c.ShouldBindQuery(¶ms); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) - return params, query, err + return params, filterQuery, err } metaValue := reflect.ValueOf(¶ms).Elem() @@ -48,10 +51,10 @@ func parseParams[T any](params T, parser *rql.Parser, c *gin.Context) (T, string if filter.IsValid() { rawFilterData := []byte(filter.String()) if len(rawFilterData) > 0 && parser != nil { // filter parameter was passed - query, err = buildQuery(parser, rawFilterData) + filterQuery, err = buildQuery(parser, rawFilterData) if err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) - return params, query, err + return params, filterQuery, err } } } @@ -61,39 +64,34 @@ func parseParams[T any](params T, parser *rql.Parser, c *gin.Context) (T, string fn.Call(nil) } - return params, query, nil + return params, filterQuery, nil } -func buildQuery(parser *rql.Parser, rawFilterData []byte) (string, error) { - var query string - var expressions []string - var args []interface{} +func buildQuery(parser *rql.Parser, rawFilterData []byte) (common.QueryFilter, error) { + var query common.QueryFilter filters, err := parser.Parse(rawFilterData) if err != nil { - return "", err + return query, err } - if filters != nil { - expressions = strings.Split(filters.FilterExp, "?") - args = filters.FilterArgs + if filters == nil || filters.FilterExp == "" { + return query, nil } - for idx, expr := range expressions { - v := "" - if idx < len(args) { - switch args[idx].(type) { - case int: - v = strconv.Itoa(args[idx].(int)) - case string: - v = fmt.Sprintf(`'%s'`, args[idx].(string)) - default: - log.Printf("Error: unable to cast and handle arg %v", args[idx]) - continue - } - } - query += fmt.Sprintf("%s%s", expr, v) + + expressions := strings.Split(filters.FilterExp, "?") + if len(expressions) != len(filters.FilterArgs)+1 { + return query, fmt.Errorf("filter placeholder count does not match argument count") } - // if query is nonempty add prefix AND for using inside SQL query after time range block - if len(query) > 0 { - query = "AND " + query + + var clause strings.Builder + clause.WriteString("AND ") + for idx, expression := range expressions { + clause.WriteString(expression) + if idx < len(filters.FilterArgs) { + name := "filter_" + strconv.Itoa(idx) + clause.WriteString("@" + name) + query.Args = append(query.Args, sql.Named(name, filters.FilterArgs[idx])) + } } + query.Clause = clause.String() return query, nil } diff --git a/src/gprofiler_flamedb_rest/handlers/common_test.go b/src/gprofiler_flamedb_rest/handlers/common_test.go index f0d12adb..194cdaa7 100644 --- a/src/gprofiler_flamedb_rest/handlers/common_test.go +++ b/src/gprofiler_flamedb_rest/handlers/common_test.go @@ -17,15 +17,20 @@ package handlers import ( - "github.com/a8m/rql" + "database/sql" + "reflect" + "strings" "testing" + + "github.com/a8m/rql" ) func TestBuildQuery(t *testing.T) { tests := []struct { - arg string - output string - parser *rql.Parser + arg string + output string + outputArgs []any + parser *rql.Parser }{ { arg: ` @@ -41,8 +46,14 @@ func TestBuildQuery(t *testing.T) { } `, parser: QueryParser, - output: "AND ContainerEnvName <> 'order-router-ar' AND (HostName = 'i-052b60b314570ca6c' " + - "OR HostName = 'i-0dc8c3917b36b7bcb' OR HostName = 'i-000a551704de2f0ab')", + output: "AND ContainerEnvName <> @filter_0 AND (HostName = @filter_1 " + + "OR HostName = @filter_2 OR HostName = @filter_3)", + outputArgs: []any{ + sql.Named("filter_0", "order-router-ar"), + sql.Named("filter_1", "i-052b60b314570ca6c"), + sql.Named("filter_2", "i-0dc8c3917b36b7bcb"), + sql.Named("filter_3", "i-000a551704de2f0ab"), + }, }, { arg: "{}", @@ -53,10 +64,31 @@ func TestBuildQuery(t *testing.T) { for _, test := range tests { query, err := buildQuery(test.parser, []byte(test.arg)) if err != nil { - t.Errorf(err.Error()) + t.Error(err) } - if query != test.output { - t.Errorf("%v != %v", query, test.output) + if query.Clause != test.output { + t.Errorf("%v != %v", query.Clause, test.output) } + if !reflect.DeepEqual(query.Args, test.outputArgs) { + t.Errorf("%v != %v", query.Args, test.outputArgs) + } + } +} + +func TestBuildQueryKeepsInjectionPayloadInParameter(t *testing.T) { + payload := "x')) UNION ALL SELECT name FROM system.tables--" + rawFilter := []byte(`{"filter":{"InstanceType":"` + payload + `"}}`) + + query, err := buildQuery(QueryParser, rawFilter) + if err != nil { + t.Fatal(err) + } + if strings.Contains(query.Clause, payload) || strings.Contains(query.Clause, "UNION ALL") { + t.Fatalf("injection payload was embedded in SQL: %s", query.Clause) + } + + expectedArgs := []any{sql.Named("filter_0", payload)} + if !reflect.DeepEqual(query.Args, expectedArgs) { + t.Fatalf("payload was not preserved as a query parameter: %v", query.Args) } } diff --git a/src/gprofiler_flamedb_rest/handlers/handlers.go b/src/gprofiler_flamedb_rest/handlers/handlers.go index 7657b50e..656f99e6 100644 --- a/src/gprofiler_flamedb_rest/handlers/handlers.go +++ b/src/gprofiler_flamedb_rest/handlers/handlers.go @@ -17,7 +17,6 @@ package handlers import ( - "fmt" "io" "log" "net/http" @@ -48,13 +47,13 @@ var MetricsQueryParser = rql.MustNewParser(rql.Config{ }) func (h Handlers) GetFlamegraph(c *gin.Context) { - params, query, err := parseParams(common.FlameGraphParams{}, QueryParser, c) + params, filterQuery, err := parseParams(common.FlameGraphParams{}, QueryParser, c) if err != nil { return } start := c.GetTime("requestStartTime") - graph, err := h.ChClient.GetTopFrames(c.Request.Context(), params, query) + graph, err := h.ChClient.GetTopFrames(c.Request.Context(), params, filterQuery) if err != nil { c.String(http.StatusInternalServerError, err.Error()) return @@ -111,7 +110,7 @@ func (h Handlers) GetFlamegraph(c *gin.Context) { func (h Handlers) QueryMeta(c *gin.Context) { var response ExecTimeInterface - params, query, err := parseParams(common.QueryParams{}, QueryParser, c) + params, filterQuery, err := parseParams(common.QueryParams{}, QueryParser, c) if err != nil { return } @@ -131,34 +130,34 @@ func (h Handlers) QueryMeta(c *gin.Context) { switch params.LookupFor { case "HostName", "hostname", "InstanceType", "instance_type": response = &FieldValueSampleResponse{ - Result: h.ChClient.FetchFieldValues(ctx, mapping[params.LookupFor], params, query), + Result: h.ChClient.FetchFieldValues(ctx, mapping[params.LookupFor], params, filterQuery), } case "ContainerEnvName", "k8s_obj", "ContainerName", "container": response = &FieldValueSampleResponse{ - Result: h.ChClient.FetchFieldValueSample(ctx, mapping[params.LookupFor], params, query), + Result: h.ChClient.FetchFieldValueSample(ctx, mapping[params.LookupFor], params, filterQuery), } case "instance_type_count": response = &InstanceTypeCountResponse{ - Result: h.ChClient.FetchInstanceTypeCount(ctx, params, query), + Result: h.ChClient.FetchInstanceTypeCount(ctx, params, filterQuery), } case "time": response = &QueryResponse{ - Result: h.ChClient.FetchTimes(ctx, params, query), + Result: h.ChClient.FetchTimes(ctx, params, filterQuery), } case "time_range": response = &QueryResponse{ - Result: h.ChClient.FetchTimeRange(ctx, params, query), + Result: h.ChClient.FetchTimeRange(ctx, params, filterQuery), } case "samples": response = &SampleCountResponse{ - Result: h.ChClient.FetchSampleCount(ctx, params, query), + Result: h.ChClient.FetchSampleCount(ctx, params, filterQuery), } case "samples_count_by_function": if len(params.FunctionName) > 0 { response = &SampleCountByFunctionResponse{ - Result: h.ChClient.FetchSampleCountByFunction(ctx, params, query), + Result: h.ChClient.FetchSampleCountByFunction(ctx, params, filterQuery), } } else { c.JSON(http.StatusBadRequest, "missing function name") @@ -187,13 +186,13 @@ func (h Handlers) QueryServices(c *gin.Context) { } func (h Handlers) QuerySessionsCount(c *gin.Context) { - params, query, err := parseParams(common.SessionsCountParams{}, QueryParser, c) + params, filterQuery, err := parseParams(common.SessionsCountParams{}, QueryParser, c) if err != nil { return } ctx := c.Request.Context() response := SessionsResponse{} - result, err := h.ChClient.FetchSessionsCount(ctx, params, query) + result, err := h.ChClient.FetchSessionsCount(ctx, params, filterQuery) response.SetExecTime(c.GetTime("requestStartTime")) if err == nil { response.Result = result @@ -204,13 +203,13 @@ func (h Handlers) QuerySessionsCount(c *gin.Context) { } func (h Handlers) GetMetricsSummary(c *gin.Context) { - params, query, err := parseParams(common.MetricsSummaryParams{}, MetricsQueryParser, c) + params, filterQuery, err := parseParams(common.MetricsSummaryParams{}, MetricsQueryParser, c) if err != nil { return } ctx := c.Request.Context() - if fetchResponse, err := h.ChClient.FetchMetricsSummary(ctx, params, query); err != nil { + if fetchResponse, err := h.ChClient.FetchMetricsSummary(ctx, params, filterQuery); err != nil { log.Print(err) c.Status(http.StatusNoContent) return @@ -251,13 +250,13 @@ func (h Handlers) GetMetricsServicesListSummary(c *gin.Context) { } func (h Handlers) GetMetricsGraph(c *gin.Context) { - params, query, err := parseParams(common.MetricsSummaryParams{}, MetricsQueryParser, c) + params, filterQuery, err := parseParams(common.MetricsSummaryParams{}, MetricsQueryParser, c) if err != nil { return } ctx := c.Request.Context() - if fetchResponse, err := h.ChClient.FetchMetricsGraph(ctx, params, query); err != nil { + if fetchResponse, err := h.ChClient.FetchMetricsGraph(ctx, params, filterQuery); err != nil { log.Print(err) c.Status(http.StatusNoContent) return @@ -271,13 +270,13 @@ func (h Handlers) GetMetricsGraph(c *gin.Context) { } func (h Handlers) GetMetricsCpuTrends(c *gin.Context) { - params, query, err := parseParams(common.MetricsCpuTrendParams{}, MetricsQueryParser, c) + params, filterQuery, err := parseParams(common.MetricsCpuTrendParams{}, MetricsQueryParser, c) if err != nil { return } ctx := c.Request.Context() - if fetchResponse, err := h.ChClient.FetchMetricsCpuTrend(ctx, params, query); err != nil { + if fetchResponse, err := h.ChClient.FetchMetricsCpuTrend(ctx, params, filterQuery); err != nil { log.Print(err) c.Status(http.StatusNoContent) return @@ -291,13 +290,12 @@ func (h Handlers) GetMetricsCpuTrends(c *gin.Context) { } func (h Handlers) GetLastHTML(c *gin.Context) { - params, query, err := parseParams(common.MetricsLastHTMLParams{}, MetricsQueryParser, c) + params, filterQuery, err := parseParams(common.MetricsLastHTMLParams{}, MetricsQueryParser, c) if err != nil { return } - fmt.Println(params, query) ctx := c.Request.Context() - htmlPath, err := h.ChClient.FetchLastHTML(ctx, params, query) + htmlPath, err := h.ChClient.FetchLastHTML(ctx, params, filterQuery) if err != nil { return } From 4057cb52defa0d1b24df6d5758e420f58a33e841 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Gabriel=20Mu=C3=B1oz?= Date: Thu, 1 Oct 2026 18:26:11 -0700 Subject: [PATCH 2/2] Fix flame graph download URL to match route Remove the trailing slash from the SVG download link to match the registered FastAPI endpoint. This avoids a 307 redirect that can downgrade HTTPS to HTTP behind a load balancer and cause browsers to block downloads. Verified downloads with development and minified production frontends. Fixes #44 --- .../frontend/src/components/profiles/header/ProfilesActions.jsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/gprofiler/frontend/src/components/profiles/header/ProfilesActions.jsx b/src/gprofiler/frontend/src/components/profiles/header/ProfilesActions.jsx index a0ce2aa8..b9ac5fbd 100644 --- a/src/gprofiler/frontend/src/components/profiles/header/ProfilesActions.jsx +++ b/src/gprofiler/frontend/src/components/profiles/header/ProfilesActions.jsx @@ -41,7 +41,7 @@ const DownloadLink = ({ serviceName, timeSelection, activeFilters, fileType, dow return (