package source_test

import (
	"fmt"
	"os"
	"path/filepath"
	"runtime"
	"strconv"
	"strings"
	"testing"

	"github.com/gkampitakis/go-snaps/snaps"
	"github.com/stretchr/testify/require"
	"go.yaml.in/yaml/v3"

	"github.com/cloudflare/pint/internal/parser"
	"github.com/cloudflare/pint/internal/parser/source"

	promParser "github.com/prometheus/prometheus/promql/parser"
)

func TestMain(t *testing.M) {
	v := t.Run()
	if _, err := snaps.Clean(t, snaps.CleanOpts{Sort: true}); err != nil {
		fmt.Printf("snaps.Clean() returned an error: %s", err)
		os.Exit(100)
	}
	os.Exit(v)
}

var testCases = []string{
	"1",
	"1 / 5",
	"(2 ^ 5) == bool 5",
	"(2 ^ 5 + 11) % 5 <= bool 2",
	"(2 ^ 5 + 11) % 5 >= bool 20",
	"(2 ^ 5 + 11) % 5 <= bool 3",
	"(2 ^ 5 + 11) % 5 < bool 1",
	"20 - 15 < bool 1",
	"2 * 5",
	"(foo or bar) * 5",
	"(foo or vector(2)) * 5",
	"(foo or vector(5)) * (vector(2) or bar)",
	`1 > bool 0`,
	`20 > bool 10`,
	`"test"`,
	"foo",
	"(foo > 1) > bool 1",
	"foo > bool 5",
	"foo > bool 5 == 1",
	"foo > bool bar",
	"(foo > bool bar) == 0",
	"foo > bool on(instance) bar",
	"(foo > bool on(instance) bar) == 1",
	"foo > bool on(instance) group_left(version) bar",
	"bar > bool on(instance) group_right(version) foo",
	"foo and bar > bool 0",
	"foo offset 5m",
	`foo{job="bar"}`,
	`foo{job=""}`,
	`foo{job="bar"} or bar{job="foo"}`,
	`foo{a="bar"} or bar{b="foo"}`,
	"foo[5m]",
	"prometheus_build_info[2m:1m]",
	"deriv(rate(distance_covered_meters_total[1m])[5m:1m])",
	"foo - 1",
	"foo / 5",
	"-foo",
	`sum(foo{job="myjob"})`,
	`sum(count(foo{job="myjob"}) by(instance))`,
	`sum(foo{job="myjob"}) > 20`,
	`sum(foo{job="myjob"}) without(job)`,
	`sum(foo) by(job)`,
	`sum(foo{job="myjob"}) by(job)`,
	`abs(foo{job="myjob"} offset 5m)`,
	`abs(foo{job="myjob"} or bar{cluster="dev"})`,
	`sum(foo{job="myjob"} or bar{cluster="dev"}) without(instance)`,
	`sum(foo{job="myjob"}) without(instance)`,
	`min(foo{job="myjob"}) / max(foo{job="myjob"})`,
	`max(foo{job="myjob"}) / min(foo{job="myjob"})`,
	`avg(foo{job="myjob"}) by(job)`,
	`group(foo) by(job)`,
	`stddev(rate(foo[5m]))`,
	`stdvar(rate(foo[5m]))`,
	`stddev_over_time(foo[5m])`,
	`stdvar_over_time(foo[5m])`,
	`quantile(0.9, rate(foo[5m]))`,
	`count_values("version", build_version)`,
	`count_values("version", build_version) without(job)`,
	`count_values("version", build_version{job="foo"}) without(job)`,
	`count_values("version", build_version) by(job)`,
	`topk(10, foo{job="myjob"}) > 10`,
	`topk(10, foo or bar)`,
	`rate(foo[10m])`,
	`sum(rate(foo[10m])) without(instance)`,
	`foo{job="foo"} / bar`,
	`foo{job="foo"} * on(instance) bar`,
	`foo{job="foo"} * on(instance) group_left(bar) bar`,
	`foo{job="foo"} * on(instance) group_left(cluster) bar{cluster="bar", ignored="true"}`,
	`foo{job="foo", ignored="true"} * on(instance) group_right(job) bar{cluster="bar"}`,
	`count(foo / bar)`,
	`count(up{job="a"} / on () up{job="b"})`,
	`count(up{job="a"} / on (env) up{job="b"})`,
	`foo{job="foo", instance="1"} and bar`,
	`foo{job="foo", instance="1"} and on(cluster) bar`,
	`topk(10, foo)`,
	`topk(10, foo) without(cluster)`,
	`topk(10, foo) by(cluster)`,
	`bottomk(10, sum(rate(foo[5m])) without(job))`,
	`foo or bar`,
	`foo or bar or baz`,
	`(foo or bar) or baz`,
	`foo unless bar`,
	`foo unless bar > 5`,
	`foo unless bar unless baz`,
	`count(sum(up{job="foo", cluster="dev"}) by(job, cluster) == 0) without(job, cluster)`,
	"year()",
	"year(foo)",
	`label_join(up{job="api-server",src1="a",src2="b",src3="c"}, "foo", ",", "src1", "src2", "src3")`,
	`
(
	sum(foo:sum > 0) without(notify)
	* on(job) group_left(notify)
	job:notify
)
and on(job)
sum(foo:count) by(job) > 20`,
	`container_file_descriptors / on (instance, app_name) container_ulimits_soft{ulimit="max_open_files"}`,
	`container_file_descriptors / on (instance, app_name) group_left() container_ulimits_soft{ulimit="max_open_files"}`,
	`absent(foo{job="bar"})`,
	`absent(foo{job="bar", cluster!="dev", instance=~".+", env="prod"})`,
	`absent(sum(foo) by(job, instance))`,
	`absent(foo{job="prometheus", xxx="1"}) AND on(job) prometheus_build_info`,
	`1 + sum(foo) by(notjob)`,
	`count(node_exporter_build_info) by (instance, version) != ignoring(package,version) group_left(foo) count(deb_package_version) by (instance, version, package)`,
	`absent(foo) or absent(bar)`,
	`absent(vector(1))`,
	`absent_over_time(foo[5m]) or absent(bar)`,
	`bar * on() group_right(cluster, env) absent(foo{job="xxx"})`,
	`bar * on() group_right() absent(foo{job="xxx"})`,
	"vector(1)",
	"vector(scalar(foo))",
	"vector(0.0  >= bool 0.5) == 1",
	`sum_over_time(foo{job="myjob"}[5m])`,
	`days_in_month()`,
	`days_in_month(foo{job="foo"})`,
	`label_replace(up{job="api-server",service="a:c"}, "foo", "$1", "service", "(.*):.*")`,
	`label_replace(sum by (pod) (pod_status) > 0, "cluster", "$1", "pod", "(.*)")`,
	`(time() - my_metric) > 5*3600`,
	`up{instance="a", job="prometheus"} * ignoring(job) up{instance="a", job="pint"}`,
	`
avg without(router, colo_id, instance) (router_anycast_prefix_enabled{cidr_use_case!~".*offpeak.*"})
< 0.5 > 0
or sum without(router, colo_id, instance) (router_anycast_prefix_enabled{cidr_use_case=~".*tier1.*"})
< on() count(colo_router_tier:disabled_pops:max{tier="1",router=~"edge.*"}) * 0.4 > 0
or avg without(router, colo_id, instance) (router_anycast_prefix_enabled{cidr_use_case=~".*regional.*"})
< 0.1 > 0
`,
	`label_replace(sum(foo) without(instance), "instance", "none", "", "")`,
	`
sum by (region, target, colo_name) (
    sum_over_time(probe_success{job="abc"}[5m])
	or
	vector(1)
) == 0`,
	`vector(1) or foo`,
	`vector(0) > 0`,
	`vector(0) > vector(1)`,
	`sum(foo or vector(0)) > 0`,
	`(sum(foo or vector(1)) > 0) == 2`,
	`(sum(foo or vector(1)) > 0) != 2`,
	`(sum(foo or vector(2)) > 0) != 2`,
	`(sum(sometimes{foo!="bar"} or vector(0)))
or
((bob > 10) or sum(foo) or vector(1))`,
	`
(
	sum(sometimes{foo!="bar"})
	or
	vector(1)
) and (
	((bob > 10) or sum(bar))
	or
	notfound > 0
)`,
	"foo offset 5m > 5",
	`
(rate(metric2[5m]) or vector(0)) +
(rate(metric1[5m]) or vector(1)) +
(rate(metric3{log_name="samplerd"}[5m]) or vector(2)) > 0
`,
	`label_replace(vector(1), "nexthop_tag", "$1", "nexthop", "(.+)")`,
	`(sum(foo{job="myjob"}))`,
	`(-foo{job="myjob"})`,
	"\n((( group(vector(0)) ))) > 0",
	"1 > bool 5",
	`prometheus_ready{job="prometheus"} unless vector(0)`,
	`prometheus_ready{job="prometheus"} unless on() vector(0)`,
	`prometheus_ready{job="prometheus"} unless on(job) vector(0)`,
	`
max by (instance, cluster) (cf_node_role{kubernetes_role="master",role="kubernetes"})
unless
	sum by (instance, cluster) (time() - node_systemd_timer_last_trigger_seconds{name=~"etcd-defrag-.*.timer"})
  	* on (instance) group_left (cluster)
    cf_node_role{kubernetes_role="master",role="kubernetes"}
`,
	`foo{a="1"} * on() bar{b="2"}`,
	`foo{a="1"} * on(instance) group_left(c,d) bar{b="2"}`,
	`foo{a="1"} * on(instance) group_right(c,d) bar{b="2"}`,
	`foo{a="1"} * on(instance) sum(bar{b="2"})`,
	`foo{a="1"} * on(instance) group_left(c,d) sum(bar{b="2"})`,
	`sum(foo{a="1"}) * on(instance) group_right(c,d) bar{b="2"}`,
	`foo{a="1"} * on(instance) group_left(c,d) sum(bar{b="2"}) without(instance)`,
	`sum(foo{a="1"}) without(instance) * on(instance) group_right(c,d) bar{b="2"}`,
	`
 max without (source_instance) (
   increase(kernel_device_io_errors_total{device!~"loop.+"}[120m]) > 3 unless on(instance, device) (
     increase(kernel_device_io_soft_errors_total{device!~"loop.+"}[125m])*2 > increase(kernel_device_io_errors_total[120m])
   )
   and on(device, instance) absent(node_disk_info)
 ) * on(instance) group_left(group) label_replace(salt_highstate_runner_configured_minions, "instance", "$1", "minion", "(.+)")
`,
	`sum(foo{a="1"}) by(job) * on() bar{b="2"}`,
	`sum(sum(foo) without(job)) by(job)`,
	`
prometheus:scrape_series_added:since_gc:sum
* on(prometheus) group_left()
label_replace(
  max(max_over_time(go_memstats_alloc_bytes{job="prometheus"}[2h])) by(instance)
  /
  max(max_over_time(prometheus_tsdb_head_series[2h])) by(instance),
  "prometheus", "$1",
  "instance", "(.+)"
)
`,
	`(day_of_week() == 6 and hour() < 1) or vector(1)`,
	`
sum by (foo, bar) (
    rate(errors_total[5m])
  * on (instance) group_left (bob, alice)
    server_errors_total
)`,
	`1 - (foo or vector(0)) < 0.999`,
	`
(
  vector(1) and month() == 2
) or vector(0)
`,
	`count by (region) (stddev by (colo_name, region) (error_total))`,
	`
(
  avg(
    rate(foo_rejections[6h])
    or
    vector(0)
  ) by (colo_name)
  /
  (
    avg(
      rate(foo_total[6h])
	  or
	  vector(1)
    ) by (colo_name)
  )
) > 5
*
(
  avg(
    rate(foo_rejections[6h] offset 1d)
	or
	vector(0)
  ) by (colo_name)
  /
  avg(
    rate(foo_total[6h] offset 1d)
	or
	vector(1)
  ) by (colo_name)
) and on (colo_name) (colo_job:foo_total:rate2m or vector(0)) > 80
  and on (colo_name) (colo_job:foo_total:rate2m offset 1d or vector(0)) > 80
`,
	`sum(selector) / sum(selector offset 30m) > 5`,
	`
count by (dc) (
  max(0 < (token_expiration - time()) < (6*60*60)) by (instance)
  * on (instance) group_right label_replace(
    configured_minions, "instance", "$1", "minion", "(.+)")
  ) > 5`,
	`topk(10, prometheus_build_info*prometheus_ready)`,
	`bottomk(10, prometheus_build_info*prometheus_ready)`,
	`histogram_fraction(0, 0.1, metric)`,
	`foo * foo `,
	`foo + on(__name__, job) foo `,
	`foo + on(__name__, job) group_left foo `,
	`foo + on(__name__, job) group_right foo `,
	`group by (env, cluster) (
      up{env="prod", job="foo"} and on (instance) (services_enabled == 999)
	)`,
	`group by (env, cluster) (
      up{env="prod", job="foo"} * on (instance) (services_enabled == 999)
	)`,
	`foo / on(instance) sum(bar)`,
	`foo / on(instance) group_left(cluster) sum(bar)`,
	`sum(bar) / on(instance) group_right(cluster) foo`,
	`sum(bar) * on(cluster) sum(foo)`,
	`
group by (colo_name, instance, tier, animal, brand, sliver, pop_name) (
  up{node_status="v", job="node_exporter"}
  and on (instance) (metal_services_enabled == 999)
  * on (colo_name) group_left(tier, animal, brand, pop_name) colo_metadata{colo_status="v"}
  * on (instance) group_left (sliver) sliver_metadata{node_status="v"}
)`,
	`
up{node_status="v", job="node_exporter"}
* on (colo_name) group_left(tier) colo_metadata{colo_status="v"}
* on (instance) group_left (sliver) sliver_metadata{node_status="v"}
`, // all group_left() labels are joined to the left
	`
up{node_status="v", job="node_exporter"}
and on (colo_name) colo_metadata{colo_status="v"}
* on (instance) group_left (sliver) sliver_metadata{node_status="v"}
`, // all group_left() labels are NOT joined to the left
	`
colo_metadata{colo_status="v"} * on (colo_name) group_right(tier, animal, brand, pop_name)
sliver_metadata{node_status="v"} * on (instance) group_right (sliver)
(metal_services_enabled == 999) * on (instance)
up{node_status="v", job="node_exporter"}
`, // only instance label will be present
	`
colo_metadata{colo_status="v"} * on (colo_name) group_right(tier, animal, brand, pop_name)
sliver_metadata{node_status="v"} * on (instance) group_right (sliver)
(metal_services_enabled == 999) * on (instance) group_right()
up{node_status="v", job="node_exporter"}
`, // no labels are joined to the right
	`
sliver_metadata{node_status="v"} * on (instance) group_right (sliver)
(metal_services_enabled == 999) * on (colo_name) group_left(tier, animal, brand, pop_name)
colo_metadata{colo_status="v"}
`, // labels from both group_left and group_right are joined
	`
colo_metadata * on (colo_name) group_right(tier, animal, brand, pop_name)
sliver_metadata * on (instance) group_right (sliver)
metal_services_enabled
`, // only sliver and tier are joined to the right
	`
colo_metadata * on (colo_name) group_right(tier, animal, brand, pop_name)
(
    sliver_metadata * on (instance) group_right (sliver)
    metal_services_enabled
)
`, // all labels are joined to the right
	`
up{node_status="v", job="node_exporter"}
* on(instance) group_left(node_status) sliver_metadata
`, // group_left on a label already guaranteed on the left
	`services_enabled{job=""}`,
	`
group by (cluster, namespace, workload, workload_type, pod) (
  label_join(
    label_join(
      group by (cluster, namespace, job_name, pod) (
        label_join(
          kube_pod_owner{job="kube-state-metrics", owner_kind="Job"}
        , "job_name", "", "owner_name")
      )
      * on (cluster, namespace, job_name) group_left(owner_kind, owner_name)
      group by (cluster, namespace, job_name, owner_kind, owner_name) (
        kube_job_owner{job="kube-state-metrics", owner_kind!="Pod", owner_kind!=""}
      )
    , "workload", "", "owner_name")
  , "workload_type", "", "owner_kind")
)
`,
	`foo{job="xxx"} + on(job) group_right(instance) bar{}`,
	`foo{job="xxx"} + ignoring(job) group_right(instance) bar{job="zzz"}`,
	`foo or ignoring(job) bar`,
	`foo or on(job) bar`,
	`
(
  sum(rate(panics_total{module_name=~".+"}[5m])) by (colo_name, module_name)
  /
  ignoring(module_name) group_left
  sum(colo:requests:rate5m) by (colo_name)
) > 0.01
`,
	`foo atan2 bar`,
	`foo offset -5m`,
	`foo @ 1609459200`,
	`foo @ start()`,
	`foo @ end()`,
	`foo[5m] @ 1609459200`,
	`rate(foo[5m] @ start())`,
	`histogram_quantile(0.9, rate(foo[5m]))`,
	`predict_linear(foo[5m], 3600)`,
	`clamp(foo, 0, 100)`,
	`irate(foo[5m])`,
	`delta(foo[5m])`,
	`idelta(foo[5m])`,
	`increase(foo[5m])`,
	`changes(foo[5m])`,
	`resets(foo[5m])`,
	`timestamp(foo)`,
	`sort(foo)`,
	`sort_desc(foo)`,
	`pi()`,
	`sgn(foo)`,
	`clamp_min(foo, 0)`,
	`clamp_max(foo, 100)`,
	`round(foo, 0.1)`,
	`exp(foo)`,
	`ln(foo)`,
	`log2(foo)`,
	`log10(foo)`,
	`sqrt(foo)`,
	`ceil(foo)`,
	`floor(foo)`,
	`avg_over_time(foo[5m])`,
	`min_over_time(foo[5m])`,
	`max_over_time(foo[5m])`,
	`count_over_time(foo[5m])`,
	`last_over_time(foo[5m])`,
	`present_over_time(foo[5m])`,
	`quantile_over_time(0.9, foo[5m])`,
	`absent_over_time(foo[5m])`,
	`deg(foo)`,
	`rad(foo)`,
	`acos(foo)`,
	`asin(foo)`,
	`atan(foo)`,
	`cos(foo)`,
	`sin(foo)`,
	`tan(foo)`,
	`acosh(foo)`,
	`asinh(foo)`,
	`atanh(foo)`,
	`cosh(foo)`,
	`sinh(foo)`,
	`tanh(foo)`,
	`histogram_count(foo)`,
	`histogram_sum(foo)`,
	`histogram_avg(foo)`,
	`histogram_fraction(0, 0.1, foo)`,
	`histogram_stddev(foo)`,
	`histogram_stdvar(foo)`,
	`min(foo) by(job)`,
	`max(foo) without(job)`,
	`sum(foo * on(job) group_left(cluster) bar) by(cluster)`,
	`sum(foo * on(job) group_left(cluster) bar) without(instance)`,
	`foo atan2 on(job) bar`,
	`foo atan2 on(job) group_left(cluster) bar`,
	`foo != bool bar`,
	`foo == on(job) bar`,
	`foo > on(job) group_left() bar`,
	`foo < on(job) group_right() bar`,
	`foo >= on(job) group_left(cluster) bar`,
	`foo <= on(job) group_right(cluster) bar`,
	`foo * ignoring(job) group_left() bar`,
	`foo / ignoring(job) group_right() bar`,
	`foo % bar`,
	`foo ^ bar`,
	`sum(foo * on(job) group_left(cluster) bar) by(job) * on(job) group_left(cluster) baz`,
	`sum by (colo) (metric_a) / scalar(sum(metric_b))`, // {colo} / <number> to avoid on()
	`clamp(foo, scalar(bar{job=~"test"}), 10)`,
	`histogram_quantile(scalar(threshold), rate(foo[5m]))`,
	`round(foo, scalar(precision_metric))`,
	`predict_linear(foo[5m], scalar(horizon))`,
	`foo + scalar(bar)`,
	`
(
  clamp_min(
    sum(foo{status_class="5xx"})
    -
    (sum(bar{status_class="5xx"}) or vector(0)),
    0
  )
  /
  sum(foo)
) > 0.05`,
	`abs(sum(foo{job="bar"}))`,
	`ceil(sum(foo{job="bar"}))`,
	`hour(sum(foo{job="bar"}))`,
	`ln(sum(foo{job="bar"}))`,
	`histogram_quantile(0.9, sum(foo{job="bar"}))`,
	`timestamp(sum(foo{job="bar"}))`,
}

func TestLabelsSource(t *testing.T) {
	type Snapshot struct {
		Expr   string
		Output []*source.Source
	}

	_, file, _, ok := runtime.Caller(0)
	require.True(t, ok, "can't get caller function")
	file = strings.TrimSuffix(filepath.Base(file), ".go")

	done := map[string]struct{}{}
	for i, expr := range testCases {
		t.Run(strconv.Itoa(i+1), func(t *testing.T) {
			if _, ok := done[expr]; ok {
				t.Fatalf("Duplicated query: %s", expr)
			}
			done[expr] = struct{}{}

			t.Log(expr)
			n, err := parser.DecodeExpr(expr)
			if err != nil {
				t.Error(err)
				t.FailNow()
			}
			output := source.LabelsSource(expr, n.Expr)

			for _, src := range output {
				src.WalkSources(func(s *source.Source, _ *source.Join, _ *source.Unless) {
					require.Positive(t, s.Position.End, "empty position %+v", s)
					if s.DeadInfo != nil {
						require.Positive(t, s.DeadInfo.Fragment.End, "empty dead position %+v", s)
					}
				})
			}

			snap := Snapshot{
				Expr:   expr,
				Output: output,
			}
			d, err := yaml.Marshal(snap)
			require.NoError(t, err, "failed to YAML encode snapshots")
			snaps.WithConfig(snaps.Dir("."), snaps.Filename(file)).MatchSnapshot(t, string(d))
		})
	}
}

func TestLabelsSourceCallCoverage(t *testing.T) {
	for name, def := range promParser.Functions {
		t.Run(name, func(t *testing.T) {
			if def.Experimental {
				t.SkipNow()
			}

			var b strings.Builder
			b.WriteString(name)
			b.WriteRune('(')
			for i, at := range def.ArgTypes {
				if i > 0 {
					b.WriteString(", ")
				}
				switch at {
				case promParser.ValueTypeNone:
				case promParser.ValueTypeScalar:
					b.WriteRune('1')
				case promParser.ValueTypeVector:
					b.WriteString("http_requests_total")
				case promParser.ValueTypeMatrix:
					b.WriteString("http_requests_total[2m]")
				case promParser.ValueTypeString:
					b.WriteString(`"foo"`)
				}
			}
			b.WriteRune(')')

			n, err := parser.DecodeExpr(b.String())
			if err != nil {
				t.Error(err)
				t.FailNow()
			}
			output := source.LabelsSource(b.String(), n.Expr)
			require.Len(t, output, 1)
			require.NotEmpty(t, output[0].Operations)
			call, ok := source.MostOuterOperation[*promParser.Call](output[0])
			require.True(t, ok, "no call found in operations for: %q ~> %+v", b.String(), output)
			require.NotNil(t, call, "no call detected in: %q ~> %+v", b.String(), output)
			require.Equal(t, name, output[0].Operation())
			require.Equal(t, def.ReturnType, output[0].Returns, "incorrect return type on Source{}")
		})
	}
}

// Verifies that experimental functions can be parsed and that
// LabelsSource produces output.
func TestLabelsSourceCallCoverageExperimental(t *testing.T) {
	// info() second arg must be a bare label selector, not a named metric.
	overrides := map[string]string{
		"info": `info(http_requests_total, {job="foo"})`,
	}

	for name, def := range promParser.Functions {
		t.Run(name, func(t *testing.T) {
			if !def.Experimental {
				t.SkipNow()
			}

			var b strings.Builder
			if override, ok := overrides[name]; ok {
				b.WriteString(override)
			} else {
				b.WriteString(name)
				b.WriteRune('(')
				for i, at := range def.ArgTypes {
					if i > 0 {
						b.WriteString(", ")
					}
					switch at {
					case promParser.ValueTypeNone:
					case promParser.ValueTypeScalar:
						b.WriteRune('1')
					case promParser.ValueTypeVector:
						b.WriteString("http_requests_total")
					case promParser.ValueTypeMatrix:
						b.WriteString("http_requests_total[2m]")
					case promParser.ValueTypeString:
						b.WriteString(`"foo"`)
					}
				}
				b.WriteRune(')')
			}

			n, err := parser.DecodeExpr(b.String())
			require.NoError(t, err, "unexpected parse error for: %s", b.String())

			output := source.LabelsSource(b.String(), n.Expr)
			require.NotEmpty(t, output, "LabelsSource returned no sources for: %s", b.String())
		})
	}
}

func TestLabelsSourceCallCoverageFail(t *testing.T) {
	n := &parser.PromQLNode{
		Expr: &promParser.Call{
			Func: &promParser.Function{
				Name: "fake_call",
			},
		},
	}
	output := source.LabelsSource("fake_call()", n.Expr)
	require.Len(t, output, 1)
	call, ok := source.MostOuterOperation[*promParser.Call](output[0])
	require.False(t, ok, "no call should have been detected in fake function, got: %v", ok)
	require.Nil(t, call, "no call should have been detected in fake function, got: %+v", call)
}

func TestVectorOperation(t *testing.T) {
	n := &parser.PromQLNode{
		Expr: &promParser.NumberLiteral{
			Val: 1,
		},
	}
	output := source.LabelsSource("1", n.Expr)
	require.Len(t, output, 1)
	require.Empty(t, output[0].Operation())
}

func TestDeadLabelKindUnknownString(t *testing.T) {
	var dl source.DeadLabelKind = 100
	require.Equal(t, "unknown", dl.String())
}

// Verifies that RangeSelectorMode.MarshalYAML returns "default" for the zero value.
func TestRangeSelectorModeDefaultMarshalYAML(t *testing.T) {
	var rsm source.RangeSelectorMode
	val, err := rsm.MarshalYAML()
	require.NoError(t, err)
	require.Equal(t, "default", val)
}

// Verifies that LabelsSource correctly processes expressions that require
// experimental parser features.
func TestLabelsSourceWithFeatures(t *testing.T) {
	type Snapshot struct {
		Expr   string
		Output []*source.Source
	}

	_, file, _, ok := runtime.Caller(0)
	require.True(t, ok, "can't get caller function")
	file = strings.TrimSuffix(filepath.Base(file), ".go")

	type testCaseT struct {
		description string
		expr        string
	}

	testCases := []testCaseT{
		// Experimental function: mad_over_time.
		{
			description: "mad_over_time",
			expr:        `mad_over_time(foo[5m])`,
		},
		// Experimental aggregator: limitk.
		{
			description: "limitk",
			expr:        `limitk(5, foo)`,
		},
		// Experimental function: sort_by_label.
		{
			description: "sort_by_label",
			expr:        `sort_by_label(foo, "job")`,
		},
		// Duration expression in matrix selector.
		{
			description: "duration expression in matrix selector",
			expr:        `foo[11s+10s]`,
		},
		// Duration expression inside rate().
		{
			description: "duration expression inside rate",
			expr:        `rate(foo[5m+1m])`,
		},
		// Anchored modifier on vector selector.
		{
			description: "anchored vector selector",
			expr:        `foo anchored`,
		},
		// Smoothed modifier on matrix selector.
		{
			description: "smoothed matrix selector inside rate",
			expr:        `rate(foo[5m] smoothed)`,
		},
		// Fill modifier on binop.
		{
			description: "binop fill modifier",
			expr:        `foo + on(job) fill(0) bar`,
		},
		// fill_left modifier on binop.
		{
			description: "binop fill_left modifier",
			expr:        `foo + on(job) fill_left(0) bar`,
		},
		// Experimental function with fill modifier.
		{
			description: "experimental function with fill modifier",
			expr:        `mad_over_time(foo[5m]) + on(job) fill(0) bar`,
		},
		// Experimental function with smoothed modifier.
		{
			description: "experimental function with smoothed modifier",
			expr:        `mad_over_time(foo[5m] smoothed)`,
		},
		// Duration expression with smoothed modifier.
		{
			description: "duration expression with smoothed modifier",
			expr:        `rate(foo[5m+1m] smoothed)`,
		},
		// fill_left modifier on binop.
		{
			description: "binop fill_left modifier only",
			expr:        `foo + on(job) fill_left(0) bar`,
		},
		// fill_right modifier on binop.
		{
			description: "binop fill_right modifier only",
			expr:        `foo + on(job) fill_right(0) bar`,
		},
		// Duration expression in subquery range.
		{
			description: "duration expression in subquery range",
			expr:        `max_over_time(rate(foo[5m])[1h+10m:5m])`,
		},
		// Duration expression in subquery step.
		{
			description: "duration expression in subquery step",
			expr:        `max_over_time(foo[1h:5m+1m])`,
		},
		// Duration expression in both subquery range and step.
		{
			description: "duration expression in subquery range and step",
			expr:        `max_over_time(rate(foo[5m])[1h+10m:5m+1m])`,
		},
		// Duration expression in vector selector offset.
		{
			description: "duration expression in vector offset",
			expr:        `foo offset 5m+1m`,
		},
		// Duration expression in vector selector offset using parenthesized
		// duration to trigger OriginalOffsetExpr on VectorSelector.
		{
			description: "duration expression in vector selector offset expr",
			expr:        `foo offset (5m+1m)`,
		},
		// Duration expression in subquery offset using parenthesized duration
		// to trigger OriginalOffsetExpr on SubqueryExpr.
		{
			description: "duration expression in subquery offset expr",
			expr:        `max_over_time(foo[1h:5m] offset (5m+1m))`,
		},
		// Fill modifier with group_right (CardOneToMany).
		{
			description: "fill with group_right",
			expr:        `foo + on(job) group_right() fill(0) bar`,
		},
		// Fill modifier with group_left (CardManyToOne).
		{
			description: "fill with group_left",
			expr:        `foo + on(job) group_left() fill(0) bar`,
		},
		// All four features combined in a single query.
		{
			description: "all four features combined",
			expr:        `mad_over_time(foo[5m+1m] smoothed) + on(job) fill(0) bar`,
		},
		// Experimental aggregator: limitk.
		{
			description: "limitk aggregator",
			expr:        `limitk(5, foo) by(job)`,
		},
		// Experimental aggregator: limit_ratio.
		{
			description: "limit_ratio aggregator",
			expr:        `limit_ratio(0.5, foo) by(job)`,
		},
		// Experimental aggregator: limitk without grouping.
		{
			description: "limitk without grouping",
			expr:        `limitk(5, foo)`,
		},
		// Experimental aggregator: limit_ratio with without().
		{
			description: "limit_ratio with without",
			expr:        `limit_ratio(0.5, foo) without(job)`,
		},
		// Experimental function: start() and end() used to compute query range.
		{
			description: "start and end functions",
			expr:        `foo / (end() - start())`,
		},
		// Experimental function: range() used as matrix selector duration.
		{
			description: "range function in rate",
			expr:        `rate(foo_total[range()])`,
		},
		// Experimental function: step() used as matrix selector duration.
		{
			description: "step function in rate",
			expr:        `rate(foo_total[step() * 4])`,
		},
		// @ start() on a vector selector requires start feature.
		{
			description: "vector selector with @ start()",
			expr:        `foo @ start()`,
		},
		// @ end() on a vector selector requires end feature.
		{
			description: "vector selector with @ end()",
			expr:        `foo @ end()`,
		},
		// @ start() on a matrix selector via rate().
		{
			description: "rate with @ start()",
			expr:        `rate(foo[5m] @ start())`,
		},
		// range() used inside a duration expression requires both features.
		{
			description: "range() inside duration expression",
			expr:        `foo[5m+range()]`,
		},
		// Smoothed vector selector with @ modifier in binary operation.
		{
			description: "smoothed with @ modifier in binop",
			expr:        `rate(foo[5m] smoothed @ 1609459200) + bar`,
		},
	}

	for _, tc := range testCases {
		t.Run(tc.description, func(t *testing.T) {
			n, err := parser.DecodeExpr(tc.expr)
			require.NoError(t, err, "unexpected parse error for: %s", tc.expr)

			output := source.LabelsSource(tc.expr, n.Expr)
			require.NotEmpty(t, output, "LabelsSource returned no sources for: %s", tc.expr)

			snap := Snapshot{
				Expr:   tc.expr,
				Output: output,
			}
			d, err := yaml.Marshal(snap)
			require.NoError(t, err, "failed to YAML encode snapshots")
			snaps.WithConfig(snaps.Dir("."), snaps.Filename(file)).MatchSnapshot(t, string(d))
		})
	}
}

func BenchmarkLabelsSource(b *testing.B) {
	type testCase struct {
		node promParser.Node
		expr string
	}
	queries := make([]testCase, 0, len(testCases))
	for _, expr := range testCases {
		n, err := parser.DecodeExpr(expr)
		require.NoError(b, err)
		queries = append(queries, testCase{
			expr: expr,
			node: n.Expr,
		})
	}

	for b.Loop() {
		for _, tc := range queries {
			source.LabelsSource(tc.expr, tc.node)
		}
	}
}
