//go:build go1.27 // Copyright 2020 The OPA Authors. All rights reserved. // Use of this source code is governed by an Apache2 // license that can be found in the LICENSE file. package cmd import ( "bytes" "context" "crypto/tls" "encoding/json" "fmt" "path/filepath" "slices" "strings" "testing" "time" "github.com/google/go-cmp/cmp" internal_logging "github.com/open-policy-agent/opa/internal/logging" "github.com/open-policy-agent/opa/v1/logging" "github.com/open-policy-agent/opa/v1/repl" "github.com/open-policy-agent/opa/v1/storage/inmem" "github.com/open-policy-agent/opa/v1/test/e2e" "github.com/open-policy-agent/opa/v1/util/test" "github.com/spf13/cobra" ) func TestREPLJSONOutputBytes(t *testing.T) { store := inmem.New() var buf bytes.Buffer r := repl.New(store, "", &buf, "json", 0, "") ctx := context.Background() if err := r.OneShot(ctx, "1 == 1"); err != nil { t.Fatalf("Unexpected error: %v", err) } expected := `{ "result": [ { "expressions": [ { "value": true, "text": "1 == 1", "location": { "row": 1, "col": 1 } } ] } ] } ` if diff := cmp.Diff(expected, buf.String()); diff != "" { t.Errorf("unexpected result (-want, +got):\n%s", diff) } } func TestRunServerBase(t *testing.T) { params := newTestRunParams() ctx, cancel := context.WithCancel(t.Context()) rt, err := initRuntime(ctx, params, nil, false) if err != nil { t.Fatalf("Unexpected error: %v", err) } testRuntime := e2e.WrapRuntime(ctx, cancel, rt) done := make(chan bool) go func() { err := rt.Serve(ctx) if err != nil { t.Errorf("Unexpected error: %s", err) } done <- true }() err = testRuntime.WaitForServer() if err != nil { t.Fatalf("Unexpected error: %s", err) } validateBasicServe(t, testRuntime) cancel() <-done } func TestRunServerBaseListenOnLocalhost(t *testing.T) { params := newTestRunParams() params.rt.V1Compatible = true ctx, cancel := context.WithCancel(t.Context()) rt, err := initRuntime(ctx, params, nil, true) if err != nil { t.Fatalf("Unexpected error: %v", err) } testRuntime := e2e.WrapRuntime(ctx, cancel, rt) done := make(chan bool) go func() { err := rt.Serve(ctx) if err != nil { t.Errorf("Unexpected error: %s", err) } done <- true }() err = testRuntime.WaitForServer() if err != nil { t.Fatalf("Unexpected error: %s", err) } validateBasicServe(t, testRuntime) if len(rt.Addrs()) != 1 { t.Fatalf("Expected 1 listening address but got %v", len(rt.Addrs())) } expected := "127.0.0.1:8181" if rt.Addrs()[0] != expected { t.Fatalf("Expected listening address %v but got %v", expected, rt.Addrs()[0]) } cancel() <-done } func TestRunServerWithDiagnosticAddr(t *testing.T) { params := newTestRunParams() params.rt.DiagnosticAddrs = &[]string{"localhost:0"} ctx, cancel := context.WithCancel(t.Context()) rt, err := initRuntime(ctx, params, nil, false) if err != nil { t.Fatalf("Unexpected error: %v", err) } testRuntime := e2e.WrapRuntime(ctx, cancel, rt) done := make(chan bool) go func() { err := rt.Serve(ctx) if err != nil { t.Errorf("Unexpected error: %s", err) } done <- true }() err = testRuntime.WaitForServer() if err != nil { t.Fatalf("Unexpected error: %s", err) } validateBasicServe(t, testRuntime) diagURL, err := testRuntime.AddrToURL(rt.DiagnosticAddrs()[0]) if err != nil { t.Fatalf("Unexpected error: %s", err) } if err := testRuntime.HealthCheck(diagURL); err != nil { t.Error(err) } cancel() <-done } func TestInitRuntimeVerifyNonBundle(t *testing.T) { params := newTestRunParams() params.pubKey = "secret" params.serverMode = false _, err := initRuntime(t.Context(), params, nil, false) if err == nil { t.Fatal("Expected error but got nil") } exp := "enable bundle mode (ie. --bundle) to verify bundle files or directories" if err.Error() != exp { t.Fatalf("expected error message %v but got %v", exp, err.Error()) } } func TestInitRuntimeCipherSuites(t *testing.T) { testCases := []struct { name string cipherSuites []string expErr bool expCipherSuites []uint16 }{ {"no cipher suites", []string{}, false, []uint16{}}, {"secure and insecure cipher suites", []string{"TLS_RSA_WITH_AES_128_CBC_SHA", "TLS_ECDHE_ECDSA_WITH_AES_128_CBC_SHA", "TLS_RSA_WITH_RC4_128_SHA"}, false, []uint16{tls.TLS_RSA_WITH_AES_128_CBC_SHA, tls.TLS_ECDHE_ECDSA_WITH_AES_128_CBC_SHA, tls.TLS_RSA_WITH_RC4_128_SHA}}, {"invalid cipher suites", []string{"foo"}, true, []uint16{}}, {"tls 1.3 cipher suite", []string{"TLS_AES_128_GCM_SHA256"}, true, []uint16{}}, {"tls 1.2-1.3 cipher suite", []string{"TLS_RSA_WITH_AES_128_GCM_SHA256", "TLS_AES_128_GCM_SHA256"}, true, []uint16{}}, } for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { params := newTestRunParams() if len(tc.cipherSuites) != 0 { params.cipherSuites = tc.cipherSuites } rt, err := initRuntime(t.Context(), params, nil, false) fmt.Println(err) if !tc.expErr && err != nil { t.Fatal("Unexpected error occurred:", err) } else if tc.expErr && err == nil { t.Fatal("Expected error but got nil") } else if err == nil { if len(tc.expCipherSuites) > 0 { if !slices.Equal(*rt.Params.CipherSuites, tc.expCipherSuites) { t.Fatalf("expected cipher suites %v but got %v", tc.expCipherSuites, *rt.Params.CipherSuites) } } else { if rt.Params.CipherSuites != nil { t.Fatal("expected no value defined for cipher suites") } } } }) } } func TestInitRuntimeSkipKnownSchemaCheck(t *testing.T) { fs := map[string]string{ "test/authz.rego": `package system.authz import rego.v1 default allow := false allow if { input.identty = "foo" # this is a typo }`, } test.WithTempFS(fs, func(rootDir string) { rootDir = filepath.Join(rootDir, "test") params := newTestRunParams() err := params.authorization.Set("basic") if err != nil { t.Fatal(err) } _, err = initRuntime(t.Context(), params, []string{rootDir}, false) if err == nil { t.Fatal("Expected error but got nil") } if !strings.Contains(err.Error(), "undefined ref: input.identty") { t.Errorf("Expected error \"%v\" not found", "undefined ref: input.identty") } // skip type checking for known input schemas params.skipKnownSchemaCheck = true _, err = initRuntime(t.Context(), params, []string{rootDir}, false) if err != nil { t.Fatal(err) } }) } func TestRunServerUploadPolicy(t *testing.T) { v0Policy := `package test p { q["a"] } q[x] { x = "a" }` v1Policy := `package test p if { q["a"] } q contains x if { x = "a" }` tests := []struct { note string v0Compatible bool module string expErr bool }{ { note: "v0-compatible, v0 policy", v0Compatible: true, module: v0Policy, }, { note: "v0-compatible, v1 policy", v0Compatible: true, module: v1Policy, expErr: true, }, { note: "v1, v0 policy", v0Compatible: false, module: v0Policy, expErr: true, }, { note: "v1, v1 policy", v0Compatible: false, module: v1Policy, }, } for i, tc := range tests { t.Run(tc.note, func(t *testing.T) { ctx, cancel := context.WithCancel(t.Context()) params := newTestRunParams() params.rt.V0Compatible = tc.v0Compatible rt, err := initRuntime(ctx, params, nil, false) if err != nil { t.Fatalf("Unexpected error: %v", err) } testRuntime := e2e.WrapRuntime(ctx, cancel, rt) done := make(chan bool) go func() { err := rt.Serve(ctx) if err != nil { t.Errorf("Unexpected error: %s", err) } done <- true }() err = testRuntime.WaitForServer() if err != nil { t.Fatalf("Unexpected error: %s", err) } // upload policy err = testRuntime.UploadPolicy(fmt.Sprintf("mod%d", i), bytes.NewBufferString(tc.module)) if tc.expErr { if err == nil { t.Fatalf("Expected error but got nil") } } else { if err != nil { t.Fatalf("Unexpected error: %s", err) } } cancel() <-done }) } } func TestRunServerCheckLogTimestampFormat(t *testing.T) { for _, format := range []string{time.Kitchen, time.RFC3339Nano} { t.Run(format, func(t *testing.T) { t.Run("parameter", func(t *testing.T) { params := newTestRunParams() params.logTimestampFormat = format params.rt.Addrs = &[]string{"localhost:0"} checkLogTimeStampFormat(t, params, format) }) t.Run("environment variable", func(t *testing.T) { t.Setenv("OPA_LOG_TIMESTAMP_FORMAT", format) params := newTestRunParams() params.rt.Addrs = &[]string{"localhost:0"} checkLogTimeStampFormat(t, params, format) }) }) } } func checkLogTimeStampFormat(t *testing.T, params runCmdParams, format string) { ctx, cancel := context.WithCancel(t.Context()) // Pass a pre-configured StandardLogger to bypass BufferedLogger and capture logs directly. var buf bytes.Buffer stdLogger := logging.New() stdLogger.SetFormatter(internal_logging.GetFormatter(params.logFormat.String(), format)) stdLogger.SetOutput(&buf) params.rt.Logger = stdLogger rt, err := initRuntime(ctx, params, nil, false) if err != nil { t.Fatalf("Unexpected error: %v", err) } testRuntime := e2e.WrapRuntime(ctx, cancel, rt) done := make(chan bool) go func() { err := rt.Serve(ctx) if err != nil { t.Errorf("Unexpected error: %s", err) } done <- true }() err = testRuntime.WaitForServer() if err != nil { t.Fatalf("Unexpected error: %s", err) } validateBasicServe(t, testRuntime) cancel() <-done for line := range strings.SplitSeq(buf.String(), "\n") { line = strings.TrimSpace(line) if line == "" { continue } var rec struct { Time string `json:"time"` } if err := json.Unmarshal([]byte(line), &rec); err != nil { t.Fatalf("incorrect log message %s: %v", line, err) } if rec.Time == "" { t.Fatalf("the time field is empty in log message: %s", line) } if _, err := time.Parse(format, rec.Time); err != nil { t.Fatalf("incorrect timestamp format %q: %v", rec.Time, err) } } } func TestInitRuntimeAddrSetByUser(t *testing.T) { testCases := []struct { name string addrValue string addrFlagSet bool }{ {"AddrSetByUser_True", "localhost:8181", true}, {"AddrSetByUser_False", "", false}, } for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { cmd := &cobra.Command{} cmd.Flags().String("addr", "", "set address") if tc.addrFlagSet { if err := cmd.Flags().Set("addr", tc.addrValue); err != nil { t.Fatalf("Failed to set addr flag: %v", err) } } params := newTestRunParams() params.rt.Addrs = &[]string{"localhost:0"} ctx, cancel := context.WithCancel(t.Context()) rt, err := initRuntime(ctx, params, []string{}, cmd.Flags().Changed("addr")) if err != nil { t.Fatalf("Unexpected error: %v", err) } if rt.Params.AddrSetByUser != tc.addrFlagSet { t.Errorf("Expected AddrSetByUser to be %v, but got %v", tc.addrFlagSet, rt.Params.AddrSetByUser) } cancel() }) } } func newTestRunParams() runCmdParams { params := newRunParams() params.rt.GracefulShutdownPeriod = 1 params.rt.Addrs = &[]string{"localhost:8181"} params.rt.DiagnosticAddrs = &[]string{} params.serverMode = true return params } func validateBasicServe(t *testing.T, runtime *e2e.TestRuntime) { t.Helper() err := runtime.UploadData(bytes.NewBufferString(`{"x": 1}`)) if err != nil { t.Fatalf("Unexpected error: %s", err) } resp := struct { Result int `json:"result"` }{} err = runtime.GetDataWithInputTyped("x", nil, &resp) if err != nil { t.Fatalf("Unexpected error: %s", err) } if resp.Result != 1 { t.Fatalf("Expected x to be 1, got %v", resp) } }